Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 5 additions & 1 deletion api/oss/src/apis/fastapi/sessions/models.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,7 +29,7 @@
from oss.src.core.sessions.mounts.dtos import SessionMount, SessionMountQuery
from oss.src.core.sessions.turns.dtos import HarnessKind, SessionTurn, SessionTurnQuery
from oss.src.core.sessions.types import SessionReference
from oss.src.core.sessions.inputs.dtos import PendingInput
from oss.src.core.sessions.inputs.dtos import PendingInputUpdate, PendingInput
from oss.src.core.shared.dtos import OTelSpanId, Windowing
from oss.src.dbs.postgres.sessions.streams.dao import MAX_SESSION_QUERY_LIMIT

Expand Down Expand Up @@ -607,3 +607,7 @@ class SessionControlOutcomeResponse(BaseModel):

class SessionContinuationResumeResponse(BaseModel):
resumed: bool


class PendingInputUpdateRequest(PendingInputUpdate):
pass
68 changes: 68 additions & 0 deletions api/oss/src/apis/fastapi/sessions/router.py
Original file line number Diff line number Diff line change
Expand Up @@ -110,6 +110,8 @@
SessionInputIdempotencyConflict,
SessionInputNotFound,
SessionInputNotRemovable,
SessionInputNotEditable,
SessionInputContentInvalid,
SessionInputRemoved,
)
from oss.src.core.sessions.inputs.dtos import PendingInputState
Expand Down Expand Up @@ -204,6 +206,7 @@
SessionResponse,
SessionsResponse,
PendingInputResponse,
PendingInputUpdateRequest,
PendingInputAdmissionRequest,
PendingInputAdmissionResponse,
SessionCapabilities,
Expand Down Expand Up @@ -2130,6 +2133,14 @@ def __init__(
if inputs_service is not None:
# The snapshot itself is `get_session_snapshot`, registered below: one route serves
# both the reconnect watermark and the durable queue.
self.router.add_api_route(
"/sessions/{session_id}/inputs/{input_id}",
self.update_pending_input,
methods=["PATCH"],
operation_id="update_pending_session_input",
response_model=PendingInputResponse,
tags=["Sessions"],
)
self.router.add_api_route(
"/sessions/{session_id}/inputs/{input_id}",
self.remove_pending_input,
Expand Down Expand Up @@ -2303,6 +2314,63 @@ async def query_sessions(
windowing=response_windowing,
)

@intercept_exceptions()
async def update_pending_input(
self,
request: Request,
session_id: str,
input_id: UUID,
payload: PendingInputUpdateRequest,
) -> PendingInputResponse:
_validate_session_id_http(session_id)
project_id = UUID(str(request.state.project_id))
user_id = request.state.user_id
if not await check_action_access(
user_uid=str(user_id),
project_id=str(project_id),
permission=Permission.RUN_SESSIONS,
):
raise FORBIDDEN_EXCEPTION
try:
item = await self.inputs_service.update(
project_id=project_id,
user_id=UUID(str(user_id)) if user_id else None,
session_id=session_id,
input_id=input_id,
update=payload,
)
except SessionInputNotFound as error:
raise HTTPException(
status_code=404,
detail={
"code": "pending_input_not_found",
"message": str(error),
"retryable": False,
"details": {"input_id": str(input_id)},
},
) from error
except SessionInputNotEditable as error:
raise HTTPException(
status_code=409,
detail={
"code": "pending_input_not_editable",
"message": str(error),
"retryable": False,
"details": {"input_id": str(input_id)},
},
) from error
except SessionInputContentInvalid as error:
raise HTTPException(
status_code=422,
detail={
"code": "pending_input_content_invalid",
"message": str(error),
"retryable": False,
"details": {"input_id": str(input_id)},
},
) from error
return PendingInputResponse(input=item)

@intercept_exceptions()
async def remove_pending_input(
self, request: Request, session_id: str, input_id: UUID
Expand Down
18 changes: 16 additions & 2 deletions api/oss/src/core/sessions/inputs/dtos.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,9 @@
from datetime import datetime
from enum import Enum
from typing import Any, Dict, Literal, Optional
from typing import Any, Dict, List, Literal, Optional
from uuid import UUID

from pydantic import BaseModel
from pydantic import BaseModel, ConfigDict, Field

from oss.src.core.shared.dtos import Identifier, Lifecycle

Expand Down Expand Up @@ -45,3 +45,17 @@ class PendingInputPromotion(BaseModel):
input: PendingInput
execution_id: str
created_at: datetime


class PendingInputAttachment(BaseModel):
model_config = ConfigDict(extra="forbid")
uri: str = Field(min_length=1)
mime_type: str = Field(min_length=1)
filename: Optional[str] = None
attachment_id: Optional[str] = Field(default=None, min_length=1)


class PendingInputUpdate(BaseModel):
model_config = ConfigDict(extra="forbid")
text: str
attachments: List[PendingInputAttachment] = Field(default_factory=list)
21 changes: 20 additions & 1 deletion api/oss/src/core/sessions/inputs/interfaces.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
from abc import ABC, abstractmethod
from typing import Any, AsyncContextManager, List, Optional
from typing import Any, AsyncContextManager, Dict, List, Optional
from uuid import UUID

from oss.src.core.sessions.inputs.dtos import PendingInput, PendingInputCreate
Expand Down Expand Up @@ -93,3 +93,22 @@ async def promote_next(
transaction: Optional[Any] = None,
) -> Optional[PendingInput]:
pass

@abstractmethod
async def lock_pending_for_edit(
self, *, project_id: UUID, session_id: str, input_id: UUID, transaction: Any
) -> Optional[PendingInput]:
pass

@abstractmethod
async def update_content(
self,
*,
project_id: UUID,
session_id: str,
input_id: UUID,
content: Dict[str, Any],
user_id: Optional[UUID],
transaction: Any,
) -> PendingInput:
pass
146 changes: 146 additions & 0 deletions api/oss/src/core/sessions/inputs/service.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,5 @@
from copy import deepcopy

import hashlib
import json
from typing import Any, Awaitable, Callable, Dict, List, Optional
Expand All @@ -8,13 +10,15 @@
PendingInputAdmission,
PendingInputCreate,
PendingInputState,
PendingInputUpdate,
)
from oss.src.core.sessions.inputs.interfaces import SessionInputsDAOInterface
from oss.src.core.sessions.inputs.types import (
SessionInputBusy,
SessionInputIdempotencyConflict,
SessionInputNotFound,
SessionInputNotRemovable,
SessionInputContentInvalid,
)
from oss.src.core.sessions.interactions.dtos import SessionInteractionStatus
from oss.src.core.sessions.interactions.interfaces import (
Expand All @@ -35,6 +39,120 @@ def input_fingerprint(*, content: Dict[str, Any], policy: str) -> str:
return hashlib.sha256(canonical).hexdigest()


def edit_pending_input_content(
content: Dict[str, Any], update: PendingInputUpdate
) -> Dict[str, Any]:
edited = deepcopy(content)
data = edited.get("data")
inputs = data.get("inputs") if isinstance(data, dict) else None
messages = inputs.get("messages") if isinstance(inputs, dict) else None
if not isinstance(messages, list):
raise SessionInputContentInvalid(
"The queued input has no editable user message."
)
message = next(
(
item
for item in reversed(messages)
if isinstance(item, dict) and item.get("role") == "user"
),
None,
)
if message is None:
raise SessionInputContentInvalid(
"The queued input has no editable user message."
)
original = message.get("content")
field = "content"
if isinstance(original, str):
if not update.attachments:
message[field] = update.text
return edited
blocks = [{"type": "text", "text": original}]
elif isinstance(original, list):
blocks = original
elif isinstance(message.get("parts"), list):
field = "parts"
blocks = message[field]
else:
raise SessionInputContentInvalid(
"The queued user message uses an unsupported content format."
)
kept = []
wrote_text = False
for block in blocks:
if isinstance(block, dict) and block.get("type") == "text":
if not wrote_text:
kept.append({**block, "text": update.text})
wrote_text = True
else:
kept.append(block)
if not wrote_text and update.text:
kept.insert(0, {"type": "text", "text": update.text})
uris = {
block.get("uri", block.get("url"))
for block in kept
if isinstance(block, dict)
and isinstance(block.get("uri", block.get("url")), str)
}
attachment_ids = set()
for block in kept:
if not isinstance(block, dict):
continue
attachment_id = block.get("attachmentId", block.get("attachment_id"))
provider_metadata = block.get("providerMetadata")
agenta_metadata = (
provider_metadata.get("agenta")
if isinstance(provider_metadata, dict)
else None
)
if not attachment_id and isinstance(agenta_metadata, dict):
attachment_id = agenta_metadata.get("attachmentId")
if isinstance(attachment_id, str) and attachment_id:
attachment_ids.add(attachment_id)
for attachment in update.attachments:
if attachment.uri in uris or (
attachment.attachment_id and attachment.attachment_id in attachment_ids
):
continue
if field == "parts":
block = {
"type": "file",
"url": attachment.uri,
"mediaType": attachment.mime_type,
}
if attachment.attachment_id is not None:
block["providerMetadata"] = {
"agenta": {"attachmentId": attachment.attachment_id}
}
if attachment.filename is not None:
block["filename"] = attachment.filename
elif attachment.attachment_id is not None:
block = {
"type": "attachment",
"attachmentId": attachment.attachment_id,
"mimeType": attachment.mime_type,
}
if attachment.filename is not None:
block["filename"] = attachment.filename
else:
block = {
"type": "image"
if attachment.mime_type.startswith("image/")
else "resource",
"uri": attachment.uri,
"mimeType": attachment.mime_type,
}
if attachment.filename is not None:
block["filename"] = attachment.filename
kept.append(block)
uris.add(attachment.uri)
if attachment.attachment_id:
attachment_ids.add(attachment.attachment_id)
message[field] = kept
return edited


class SessionInputsService:
def __init__(
self,
Expand Down Expand Up @@ -275,3 +393,31 @@ async def remove(
if existing is not None and existing.state != PendingInputState.pending:
raise SessionInputNotRemovable(str(input_id))
raise SessionInputNotFound(str(input_id))

async def update(
self,
*,
project_id: UUID,
session_id: str,
input_id: UUID,
user_id: Optional[UUID],
update: PendingInputUpdate,
) -> PendingInput:
async with self._dao.transaction() as transaction:
item = await self._dao.lock_pending_for_edit(
project_id=project_id,
session_id=session_id,
input_id=input_id,
transaction=transaction,
)
if item is None:
raise SessionInputNotFound(str(input_id))
content = edit_pending_input_content(item.content, update)
return await self._dao.update_content(
project_id=project_id,
session_id=session_id,
input_id=input_id,
content=content,
user_id=user_id,
transaction=transaction,
)
12 changes: 12 additions & 0 deletions api/oss/src/core/sessions/inputs/types.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,3 +32,15 @@ class SessionInputRemoved(SessionInputError):
def __init__(self, input_id: str):
self.input_id = input_id
super().__init__("The queued input was removed and cannot be sent.")


class SessionInputNotEditable(SessionInputError):
def __init__(self, input_id: str):
self.input_id = input_id
super().__init__(
"The queued input is no longer editable because it was removed, promoted, or selected to run next."
)


class SessionInputContentInvalid(SessionInputError):
pass
Loading
Loading