diff --git a/sdk/python/src/codex_app_server/async_client.py b/sdk/python/src/codex_app_server/async_client.py index 6ca0c42a78..bb767fa103 100644 --- a/sdk/python/src/codex_app_server/async_client.py +++ b/sdk/python/src/codex_app_server/async_client.py @@ -6,7 +6,13 @@ from typing import AsyncIterator, Callable, Iterable, ParamSpec, TypeVar from pydantic import BaseModel -from .client import AppServerClient, AppServerConfig +from .client import ( + AppServerClient, + AppServerConfig, + ApprovalHandler, + ModelCompactHandler, + ModelRequestHandler, +) from .generated.v2_all import ( AgentMessageDeltaNotification, ModelListResponse, @@ -39,8 +45,19 @@ ReturnT = TypeVar("ReturnT") class AsyncAppServerClient: """Async wrapper around AppServerClient using thread offloading.""" - def __init__(self, config: AppServerConfig | None = None) -> None: - self._sync = AppServerClient(config=config) + def __init__( + self, + config: AppServerConfig | None = None, + approval_handler: ApprovalHandler | None = None, + model_request_handler: ModelRequestHandler | None = None, + model_compact_handler: ModelCompactHandler | None = None, + ) -> None: + self._sync = AppServerClient( + config=config, + approval_handler=approval_handler, + model_request_handler=model_request_handler, + model_compact_handler=model_compact_handler, + ) # Single stdio transport cannot be read safely from multiple threads. self._transport_lock = asyncio.Lock() @@ -99,7 +116,57 @@ class AsyncAppServerClient: response_model=response_model, ) - async def thread_start(self, params: V2ThreadStartParams | JsonObject | None = None) -> ThreadStartResponse: + async def send_model_stream_metadata( + self, + *, + thread_id: str, + turn_id: str, + request_id: str, + metadata: dict[str, str], + ) -> None: + await self._call_sync( + self._sync.send_model_stream_metadata, + thread_id=thread_id, + turn_id=turn_id, + request_id=request_id, + metadata=metadata, + ) + + async def send_model_stream_event( + self, + *, + thread_id: str, + turn_id: str, + request_id: str, + event: JsonObject, + ) -> None: + await self._call_sync( + self._sync.send_model_stream_event, + thread_id=thread_id, + turn_id=turn_id, + request_id=request_id, + event=event, + ) + + async def send_model_request_failed( + self, + *, + thread_id: str, + turn_id: str, + request_id: str, + error: JsonObject, + ) -> None: + await self._call_sync( + self._sync.send_model_request_failed, + thread_id=thread_id, + turn_id=turn_id, + request_id=request_id, + error=error, + ) + + async def thread_start( + self, params: V2ThreadStartParams | JsonObject | None = None + ) -> ThreadStartResponse: return await self._call_sync(self._sync.thread_start, params) async def thread_resume( @@ -109,10 +176,14 @@ class AsyncAppServerClient: ) -> ThreadResumeResponse: return await self._call_sync(self._sync.thread_resume, thread_id, params) - async def thread_list(self, params: V2ThreadListParams | JsonObject | None = None) -> ThreadListResponse: + async def thread_list( + self, params: V2ThreadListParams | JsonObject | None = None + ) -> ThreadListResponse: return await self._call_sync(self._sync.thread_list, params) - async def thread_read(self, thread_id: str, include_turns: bool = False) -> ThreadReadResponse: + async def thread_read( + self, thread_id: str, include_turns: bool = False + ) -> ThreadReadResponse: return await self._call_sync(self._sync.thread_read, thread_id, include_turns) async def thread_fork( @@ -140,9 +211,13 @@ class AsyncAppServerClient: input_items: list[JsonObject] | JsonObject | str, params: V2TurnStartParams | JsonObject | None = None, ) -> TurnStartResponse: - return await self._call_sync(self._sync.turn_start, thread_id, input_items, params) + return await self._call_sync( + self._sync.turn_start, thread_id, input_items, params + ) - async def turn_interrupt(self, thread_id: str, turn_id: str) -> TurnInterruptResponse: + async def turn_interrupt( + self, thread_id: str, turn_id: str + ) -> TurnInterruptResponse: return await self._call_sync(self._sync.turn_interrupt, thread_id, turn_id) async def turn_steer( @@ -187,7 +262,9 @@ class AsyncAppServerClient: async def wait_for_turn_completed(self, turn_id: str) -> TurnCompletedNotification: return await self._call_sync(self._sync.wait_for_turn_completed, turn_id) - async def stream_until_methods(self, methods: Iterable[str] | str) -> list[Notification]: + async def stream_until_methods( + self, methods: Iterable[str] | str + ) -> list[Notification]: return await self._call_sync(self._sync.stream_until_methods, methods) async def stream_text( diff --git a/sdk/python/src/codex_app_server/client.py b/sdk/python/src/codex_app_server/client.py index aa7b574a3f..ae0f3546e5 100644 --- a/sdk/python/src/codex_app_server/client.py +++ b/sdk/python/src/codex_app_server/client.py @@ -12,7 +12,7 @@ from typing import Callable, Iterable, Iterator, TypeVar from pydantic import BaseModel -from .errors import AppServerError, TransportClosedError, map_jsonrpc_error +from .errors import AppServerError, JsonRpcError, TransportClosedError, map_jsonrpc_error from .generated.notification_registry import NOTIFICATION_MODELS from .generated.v2_all import ( AgentMessageDeltaNotification, @@ -47,6 +47,8 @@ from .retry import retry_on_overload ModelT = TypeVar("ModelT", bound=BaseModel) ApprovalHandler = Callable[[str, JsonObject | None], JsonObject] +ModelRequestHandler = Callable[[JsonObject | None], JsonObject] +ModelCompactHandler = Callable[[JsonObject | None], JsonObject] RUNTIME_PKG_NAME = "codex-cli-bin" @@ -74,7 +76,9 @@ def _params_dict( return dumped if isinstance(params, dict): return params - raise TypeError(f"Expected generated params model or dict, got {type(params).__name__}") + raise TypeError( + f"Expected generated params model or dict, got {type(params).__name__}" + ) def _installed_codex_path() -> Path: @@ -140,9 +144,17 @@ class AppServerClient: self, config: AppServerConfig | None = None, approval_handler: ApprovalHandler | None = None, + model_request_handler: ModelRequestHandler | None = None, + model_compact_handler: ModelCompactHandler | None = None, ) -> None: self.config = config or AppServerConfig() self._approval_handler = approval_handler or self._default_approval_handler + self._model_request_handler = ( + model_request_handler or self._default_model_request_handler + ) + self._model_compact_handler = ( + model_compact_handler or self._default_model_compact_handler + ) self._proc: subprocess.Popen[str] | None = None self._lock = threading.Lock() self._turn_consumer_lock = threading.Lock() @@ -238,14 +250,15 @@ class AppServerClient: def _request_raw(self, method: str, params: JsonObject | None = None) -> JsonValue: request_id = str(uuid.uuid4()) - self._write_message({"id": request_id, "method": method, "params": params or {}}) + self._write_message( + {"id": request_id, "method": method, "params": params or {}} + ) while True: msg = self._read_message() if "method" in msg and "id" in msg: - response = self._handle_server_request(msg) - self._write_message({"id": msg["id"], "result": response}) + self._respond_to_server_request(msg) continue if "method" in msg and "id" not in msg: @@ -272,6 +285,60 @@ class AppServerClient: def notify(self, method: str, params: JsonObject | None = None) -> None: self._write_message({"method": method, "params": params or {}}) + def send_model_stream_metadata( + self, + *, + thread_id: str, + turn_id: str, + request_id: str, + metadata: dict[str, str], + ) -> None: + self.notify( + "model/streamMetadata", + { + "threadId": thread_id, + "turnId": turn_id, + "requestId": request_id, + "metadata": metadata, + }, + ) + + def send_model_stream_event( + self, + *, + thread_id: str, + turn_id: str, + request_id: str, + event: JsonObject, + ) -> None: + self.notify( + "model/streamEvent", + { + "threadId": thread_id, + "turnId": turn_id, + "requestId": request_id, + "event": event, + }, + ) + + def send_model_request_failed( + self, + *, + thread_id: str, + turn_id: str, + request_id: str, + error: JsonObject, + ) -> None: + self.notify( + "model/requestFailed", + { + "threadId": thread_id, + "turnId": turn_id, + "requestId": request_id, + "error": error, + }, + ) + def next_notification(self) -> Notification: if self._pending_notifications: return self._pending_notifications.popleft() @@ -279,8 +346,7 @@ class AppServerClient: while True: msg = self._read_message() if "method" in msg and "id" in msg: - response = self._handle_server_request(msg) - self._write_message({"id": msg["id"], "result": response}) + self._respond_to_server_request(msg) continue if "method" in msg and "id" not in msg: return self._coerce_notification(msg["method"], msg.get("params")) @@ -300,8 +366,12 @@ class AppServerClient: if self._active_turn_consumer == turn_id: self._active_turn_consumer = None - def thread_start(self, params: V2ThreadStartParams | JsonObject | None = None) -> ThreadStartResponse: - return self.request("thread/start", _params_dict(params), response_model=ThreadStartResponse) + def thread_start( + self, params: V2ThreadStartParams | JsonObject | None = None + ) -> ThreadStartResponse: + return self.request( + "thread/start", _params_dict(params), response_model=ThreadStartResponse + ) def thread_resume( self, @@ -309,12 +379,20 @@ class AppServerClient: params: V2ThreadResumeParams | JsonObject | None = None, ) -> ThreadResumeResponse: payload = {"threadId": thread_id, **_params_dict(params)} - return self.request("thread/resume", payload, response_model=ThreadResumeResponse) + return self.request( + "thread/resume", payload, response_model=ThreadResumeResponse + ) - def thread_list(self, params: V2ThreadListParams | JsonObject | None = None) -> ThreadListResponse: - return self.request("thread/list", _params_dict(params), response_model=ThreadListResponse) + def thread_list( + self, params: V2ThreadListParams | JsonObject | None = None + ) -> ThreadListResponse: + return self.request( + "thread/list", _params_dict(params), response_model=ThreadListResponse + ) - def thread_read(self, thread_id: str, include_turns: bool = False) -> ThreadReadResponse: + def thread_read( + self, thread_id: str, include_turns: bool = False + ) -> ThreadReadResponse: return self.request( "thread/read", {"threadId": thread_id, "includeTurns": include_turns}, @@ -330,10 +408,18 @@ class AppServerClient: return self.request("thread/fork", payload, response_model=ThreadForkResponse) def thread_archive(self, thread_id: str) -> ThreadArchiveResponse: - return self.request("thread/archive", {"threadId": thread_id}, response_model=ThreadArchiveResponse) + return self.request( + "thread/archive", + {"threadId": thread_id}, + response_model=ThreadArchiveResponse, + ) def thread_unarchive(self, thread_id: str) -> ThreadUnarchiveResponse: - return self.request("thread/unarchive", {"threadId": thread_id}, response_model=ThreadUnarchiveResponse) + return self.request( + "thread/unarchive", + {"threadId": thread_id}, + response_model=ThreadUnarchiveResponse, + ) def thread_set_name(self, thread_id: str, name: str) -> ThreadSetNameResponse: return self.request( @@ -457,12 +543,16 @@ class AppServerClient: model = NOTIFICATION_MODELS.get(method) if model is None: - return Notification(method=method, payload=UnknownNotification(params=params_dict)) + return Notification( + method=method, payload=UnknownNotification(params=params_dict) + ) try: payload = model.model_validate(params_dict) except Exception: # noqa: BLE001 - return Notification(method=method, payload=UnknownNotification(params=params_dict)) + return Notification( + method=method, payload=UnknownNotification(params=params_dict) + ) return Notification(method=method, payload=payload) def _normalize_input_items( @@ -475,13 +565,28 @@ class AppServerClient: return [input_items] return input_items - def _default_approval_handler(self, method: str, params: JsonObject | None) -> JsonObject: + def _default_approval_handler( + self, method: str, params: JsonObject | None + ) -> JsonObject: if method == "item/commandExecution/requestApproval": return {"decision": "accept"} if method == "item/fileChange/requestApproval": return {"decision": "accept"} return {} + def _default_model_request_handler(self, _params: JsonObject | None) -> JsonObject: + return { + "accepted": False, + "rejectionReason": ( + "No model_request_handler is installed on this AppServerClient" + ), + } + + def _default_model_compact_handler(self, _params: JsonObject | None) -> JsonObject: + raise AppServerError( + "No model_compact_handler is installed on this AppServerClient" + ) + def _start_stderr_drain_thread(self) -> None: if self._proc is None or self._proc.stderr is None: return @@ -504,11 +609,49 @@ class AppServerClient: params = msg.get("params") if not isinstance(method, str): return {} + if method == "model/request": + return self._model_request_handler( + params if isinstance(params, dict) else None, + ) + if method == "model/compact": + return self._model_compact_handler( + params if isinstance(params, dict) else None, + ) return self._approval_handler( method, params if isinstance(params, dict) else None, ) + def _respond_to_server_request(self, msg: dict[str, JsonValue]) -> None: + request_id = msg.get("id") + try: + response = self._handle_server_request(msg) + except JsonRpcError as exc: + self._write_message( + { + "id": request_id, + "error": { + "code": exc.code, + "message": exc.message, + "data": exc.data, + }, + } + ) + return + except AppServerError as exc: + self._write_message( + { + "id": request_id, + "error": { + "code": -32603, + "message": str(exc), + }, + } + ) + return + + self._write_message({"id": request_id, "result": response}) + def _write_message(self, payload: JsonObject) -> None: if self._proc is None or self._proc.stdin is None: raise TransportClosedError("app-server is not running") diff --git a/sdk/python/src/codex_app_server/generated/v2_all.py b/sdk/python/src/codex_app_server/generated/v2_all.py index 0ff2c5897d..2fd1bb04e3 100644 --- a/sdk/python/src/codex_app_server/generated/v2_all.py +++ b/sdk/python/src/codex_app_server/generated/v2_all.py @@ -1501,6 +1501,26 @@ class ModelListParams(BaseModel): ] = None +class ModelRequestError(BaseModel): + model_config = ConfigDict( + populate_by_name=True, + ) + code: str | None = None + message: str + param: str | None = None + type: str + + +class ModelRequestFailedNotification(BaseModel): + model_config = ConfigDict( + populate_by_name=True, + ) + error: ModelRequestError + request_id: Annotated[str, Field(alias="requestId")] + thread_id: Annotated[str, Field(alias="threadId")] + turn_id: Annotated[str, Field(alias="turnId")] + + class ModelRerouteReason(RootModel[Literal["highRiskCyberActivity"]]): model_config = ConfigDict( populate_by_name=True, @@ -1519,6 +1539,26 @@ class ModelReroutedNotification(BaseModel): turn_id: Annotated[str, Field(alias="turnId")] +class ModelStreamEventNotification(BaseModel): + model_config = ConfigDict( + populate_by_name=True, + ) + event: Any + request_id: Annotated[str, Field(alias="requestId")] + thread_id: Annotated[str, Field(alias="threadId")] + turn_id: Annotated[str, Field(alias="turnId")] + + +class ModelStreamMetadataNotification(BaseModel): + model_config = ConfigDict( + populate_by_name=True, + ) + metadata: dict[str, str] + request_id: Annotated[str, Field(alias="requestId")] + thread_id: Annotated[str, Field(alias="threadId")] + turn_id: Annotated[str, Field(alias="turnId")] + + class ModelUpgradeInfo(BaseModel): model_config = ConfigDict( populate_by_name=True, @@ -1534,6 +1574,11 @@ class NetworkAccess(Enum): enabled = "enabled" +class NetworkDelegationMode(Enum): + enabled = "enabled" + disabled = "disabled" + + class NetworkRequirements(BaseModel): model_config = ConfigDict( populate_by_name=True, @@ -1799,6 +1844,11 @@ class ReadOnlyAccess(RootModel[RestrictedReadOnlyAccess | FullAccessReadOnlyAcce root: RestrictedReadOnlyAccess | FullAccessReadOnlyAccess +class RealtimeConversationVersion(Enum): + v1 = "v1" + v2 = "v2" + + class ReasoningEffort(Enum): none = "none" minimal = "minimal" @@ -3058,6 +3108,7 @@ class ThreadRealtimeStartedNotification(BaseModel): ) session_id: Annotated[str | None, Field(alias="sessionId")] = None thread_id: Annotated[str, Field(alias="threadId")] + version: RealtimeConversationVersion class ThreadResumeParams(BaseModel): @@ -4593,6 +4644,16 @@ class ModelListResponse(BaseModel): ] = None +class NetworkDelegationConfig(BaseModel): + model_config = ConfigDict( + populate_by_name=True, + ) + mode: NetworkDelegationMode + stream_idle_timeout_ms: Annotated[ + int | None, Field(alias="streamIdleTimeoutMs", ge=0) + ] = None + + class OverriddenMetadata(BaseModel): model_config = ConfigDict( populate_by_name=True,