diff --git a/python/omp-rpc/src/omp_rpc/__init__.py b/python/omp-rpc/src/omp_rpc/__init__.py index 014da5814..3576f18bc 100644 --- a/python/omp-rpc/src/omp_rpc/__init__.py +++ b/python/omp-rpc/src/omp_rpc/__init__.py @@ -16,7 +16,13 @@ from .client import ( ProtocolErrorListener, UiRequestListener, ) -from .host_tools import HostTool, HostToolContext, HostToolResultPayload, HostToolResultValue, host_tool +from .host_tools import ( + HostTool, + HostToolContext, + HostToolResultPayload, + HostToolResultValue, + host_tool, +) from .host_uris import ( HostUri, HostUriContentType, diff --git a/python/omp-rpc/src/omp_rpc/client.py b/python/omp-rpc/src/omp_rpc/client.py index 04e7e5b8f..570c6fb2f 100644 --- a/python/omp-rpc/src/omp_rpc/client.py +++ b/python/omp-rpc/src/omp_rpc/client.py @@ -322,8 +322,12 @@ class RpcClient: self._extra_args = tuple(extra_args) self._startup_timeout = startup_timeout self._request_timeout = request_timeout - self._max_event_history = self._validate_history_limit("max_event_history", max_event_history) - self._max_stderr_chunks = self._validate_history_limit("max_stderr_chunks", max_stderr_chunks) + self._max_event_history = self._validate_history_limit( + "max_event_history", max_event_history + ) + self._max_stderr_chunks = self._validate_history_limit( + "max_stderr_chunks", max_stderr_chunks + ) self._process: subprocess.Popen[str] | None = None self._stdout_thread: threading.Thread | None = None @@ -337,7 +341,9 @@ class RpcClient: self._pending_host_uri_requests: dict[str, _PendingHostUriRequest] = {} self._request_id = 0 self._events = _BoundedHistory[JsonObject](self._max_event_history) - self._async_errors = _BoundedHistory[BaseException](_DEFAULT_ERROR_HISTORY_LIMIT) + self._async_errors = _BoundedHistory[BaseException]( + _DEFAULT_ERROR_HISTORY_LIMIT + ) self._scheduled_agent_runs = 0 self._completed_agent_runs = 0 self._last_schedule_async_error_index = 0 @@ -346,8 +352,12 @@ class RpcClient: self._closed_error: BaseException | None = None self._stopping = False self._ready_received = False - self._protocol_errors = _BoundedHistory[RpcProtocolError](_DEFAULT_ERROR_HISTORY_LIMIT) - self._listener_errors = _BoundedHistory[ListenerErrorEvent](_DEFAULT_ERROR_HISTORY_LIMIT) + self._protocol_errors = _BoundedHistory[RpcProtocolError]( + _DEFAULT_ERROR_HISTORY_LIMIT + ) + self._listener_errors = _BoundedHistory[ListenerErrorEvent]( + _DEFAULT_ERROR_HISTORY_LIMIT + ) self._prompt_lifecycle = _PromptLifecycleCoordinator() self._notification_listeners: list[NotificationListener] = [] @@ -422,15 +432,21 @@ class RpcClient: ) self._process = process - self._stdout_thread = threading.Thread(target=self._read_stdout_loop, name="omp-rpc-stdout", daemon=True) - self._stderr_thread = threading.Thread(target=self._read_stderr_loop, name="omp-rpc-stderr", daemon=True) + self._stdout_thread = threading.Thread( + target=self._read_stdout_loop, name="omp-rpc-stdout", daemon=True + ) + self._stderr_thread = threading.Thread( + target=self._read_stderr_loop, name="omp-rpc-stderr", daemon=True + ) self._stdout_thread.start() self._stderr_thread.start() if not self._ready.wait(self._startup_timeout): stderr = self.stderr self.stop() - raise RpcTimeoutError(f"Timed out waiting for RPC ready signal. Stderr: {stderr}") + raise RpcTimeoutError( + f"Timed out waiting for RPC ready signal. Stderr: {stderr}" + ) if not self._ready_received: error = self._closed_error @@ -439,8 +455,12 @@ class RpcClient: if isinstance(error, RpcError): raise error if error is not None: - raise RpcProcessExitError(f"RPC process stopped before ready: {error}. Stderr: {stderr}") from error - raise RpcTimeoutError(f"Timed out waiting for RPC ready signal. Stderr: {stderr}") + raise RpcProcessExitError( + f"RPC process stopped before ready: {error}. Stderr: {stderr}" + ) from error + raise RpcTimeoutError( + f"Timed out waiting for RPC ready signal. Stderr: {stderr}" + ) if self._custom_tools: self.set_custom_tools(self._custom_tools) @@ -536,31 +556,47 @@ class RpcClient: def on_message_end(self, listener: MessageEndListener) -> Callable[[], None]: return self._add_typed_event_listener("message_end", listener) - def on_tool_execution_start(self, listener: ToolExecutionStartListener) -> Callable[[], None]: + def on_tool_execution_start( + self, listener: ToolExecutionStartListener + ) -> Callable[[], None]: return self._add_typed_event_listener("tool_execution_start", listener) - def on_tool_execution_update(self, listener: ToolExecutionUpdateListener) -> Callable[[], None]: + def on_tool_execution_update( + self, listener: ToolExecutionUpdateListener + ) -> Callable[[], None]: return self._add_typed_event_listener("tool_execution_update", listener) - def on_tool_execution_end(self, listener: ToolExecutionEndListener) -> Callable[[], None]: + def on_tool_execution_end( + self, listener: ToolExecutionEndListener + ) -> Callable[[], None]: return self._add_typed_event_listener("tool_execution_end", listener) - def on_auto_compaction_start(self, listener: AutoCompactionStartListener) -> Callable[[], None]: + def on_auto_compaction_start( + self, listener: AutoCompactionStartListener + ) -> Callable[[], None]: return self._add_typed_event_listener("auto_compaction_start", listener) - def on_auto_compaction_end(self, listener: AutoCompactionEndListener) -> Callable[[], None]: + def on_auto_compaction_end( + self, listener: AutoCompactionEndListener + ) -> Callable[[], None]: return self._add_typed_event_listener("auto_compaction_end", listener) - def on_auto_retry_start(self, listener: AutoRetryStartListener) -> Callable[[], None]: + def on_auto_retry_start( + self, listener: AutoRetryStartListener + ) -> Callable[[], None]: return self._add_typed_event_listener("auto_retry_start", listener) def on_auto_retry_end(self, listener: AutoRetryEndListener) -> Callable[[], None]: return self._add_typed_event_listener("auto_retry_end", listener) - def on_retry_fallback_applied(self, listener: RetryFallbackAppliedListener) -> Callable[[], None]: + def on_retry_fallback_applied( + self, listener: RetryFallbackAppliedListener + ) -> Callable[[], None]: return self._add_typed_event_listener("retry_fallback_applied", listener) - def on_retry_fallback_succeeded(self, listener: RetryFallbackSucceededListener) -> Callable[[], None]: + def on_retry_fallback_succeeded( + self, listener: RetryFallbackSucceededListener + ) -> Callable[[], None]: return self._add_typed_event_listener("retry_fallback_succeeded", listener) def on_ttsr_triggered(self, listener: TtsrTriggeredListener) -> Callable[[], None]: @@ -576,7 +612,9 @@ class RpcClient: self._ui_request_listeners.append(listener) return lambda: self._remove_listener(self._ui_request_listeners, listener) - def on_extension_error(self, listener: ExtensionErrorListener) -> Callable[[], None]: + def on_extension_error( + self, listener: ExtensionErrorListener + ) -> Callable[[], None]: self._extension_error_listeners.append(listener) return lambda: self._remove_listener(self._extension_error_listeners, listener) @@ -588,9 +626,13 @@ class RpcClient: self._listener_error_listeners.append(listener) return lambda: self._remove_listener(self._listener_error_listeners, listener) - def on_unknown_notification(self, listener: UnknownNotificationListener) -> Callable[[], None]: + def on_unknown_notification( + self, listener: UnknownNotificationListener + ) -> Callable[[], None]: self._unknown_notification_listeners.append(listener) - return lambda: self._remove_listener(self._unknown_notification_listeners, listener) + return lambda: self._remove_listener( + self._unknown_notification_listeners, listener + ) def install_headless_ui( self, @@ -651,16 +693,26 @@ class RpcClient: try: return self._ui_requests.get(timeout=timeout) except queue.Empty as exc: - raise RpcTimeoutError("Timed out waiting for an extension UI request") from exc + raise RpcTimeoutError( + "Timed out waiting for an extension UI request" + ) from exc def send_ui_value(self, request_id: str, value: str) -> None: - self._send_notification({"type": "extension_ui_response", "id": request_id, "value": value}) + self._send_notification( + {"type": "extension_ui_response", "id": request_id, "value": value} + ) def send_ui_confirmation(self, request_id: str, confirmed: bool) -> None: - self._send_notification({"type": "extension_ui_response", "id": request_id, "confirmed": confirmed}) + self._send_notification( + {"type": "extension_ui_response", "id": request_id, "confirmed": confirmed} + ) def cancel_ui_request(self, request_id: str, *, timed_out: bool = False) -> None: - payload: JsonObject = {"type": "extension_ui_response", "id": request_id, "cancelled": True} + payload: JsonObject = { + "type": "extension_ui_response", + "id": request_id, + "cancelled": True, + } if timed_out: payload["timedOut"] = True self._send_notification(payload) @@ -724,14 +776,21 @@ class RpcClient: return parse_session_stats(payload) def export_html(self, output_path: str | Path | None = None) -> Path: - payload = self._request("export_html", outputPath=str(output_path) if output_path is not None else None) + payload = self._request( + "export_html", + outputPath=str(output_path) if output_path is not None else None, + ) return Path(str(payload["path"])) def new_session(self, parent_session: str | None = None) -> CancellationResult: - return parse_cancellation_result(self._request("new_session", parentSession=parent_session)) + return parse_cancellation_result( + self._request("new_session", parentSession=parent_session) + ) def switch_session(self, session_path: str | Path) -> CancellationResult: - return parse_cancellation_result(self._request("switch_session", sessionPath=str(session_path))) + return parse_cancellation_result( + self._request("switch_session", sessionPath=str(session_path)) + ) def branch(self, entry_id: str) -> BranchResult: return parse_branch_result(self._request("branch", entryId=entry_id)) @@ -750,7 +809,9 @@ class RpcClient: def get_todos(self) -> tuple[TodoPhase, ...]: return self.get_state().todo_phases - def set_todos(self, todos: Sequence[TodoSeed | TodoPhaseSeed]) -> tuple[TodoPhase, ...]: + def set_todos( + self, todos: Sequence[TodoSeed | TodoPhaseSeed] + ) -> tuple[TodoPhase, ...]: phases = self._normalize_todo_phases(todos) payload = self._request("set_todos", phases=cast(JsonValue, phases)) return parse_todo_phases(payload.get("todoPhases")) @@ -795,7 +856,11 @@ class RpcClient: schemes_payload: list[JsonObject] = [] for uri in self._host_uris: - entry: JsonObject = {"scheme": uri.scheme, "writable": uri.writable, "immutable": uri.immutable} + entry: JsonObject = { + "scheme": uri.scheme, + "writable": uri.writable, + "immutable": uri.immutable, + } if uri.description is not None: entry["description"] = uri.description schemes_payload.append(entry) @@ -824,17 +889,35 @@ class RpcClient: ) self._mark_agent_run_scheduled() - def steer(self, message: str, *, images: Sequence[ImageContent] | None = None) -> None: - self._request("steer", message=message, images=list(images) if images is not None else None) + def steer( + self, message: str, *, images: Sequence[ImageContent] | None = None + ) -> None: + self._request( + "steer", + message=message, + images=list(images) if images is not None else None, + ) - def follow_up(self, message: str, *, images: Sequence[ImageContent] | None = None) -> None: - self._request("follow_up", message=message, images=list(images) if images is not None else None) + def follow_up( + self, message: str, *, images: Sequence[ImageContent] | None = None + ) -> None: + self._request( + "follow_up", + message=message, + images=list(images) if images is not None else None, + ) def abort(self) -> None: self._request("abort") - def abort_and_prompt(self, message: str, *, images: Sequence[ImageContent] | None = None) -> None: - self._request("abort_and_prompt", message=message, images=list(images) if images is not None else None) + def abort_and_prompt( + self, message: str, *, images: Sequence[ImageContent] | None = None + ) -> None: + self._request( + "abort_and_prompt", + message=message, + images=list(images) if images is not None else None, + ) self._mark_agent_run_scheduled() def prompt_and_wait( @@ -851,7 +934,9 @@ class RpcClient: start_index = self._current_event_index() start_async_error_index = self._current_async_error_index() self.prompt(message, images=images, streaming_behavior=streaming_behavior) - events = self._wait_for_agent_end(start_index, start_async_error_index, timeout=timeout) + events = self._wait_for_agent_end( + start_index, start_async_error_index, timeout=timeout + ) return self._build_prompt_turn(events) finally: self._prompt_lifecycle.release(operation) @@ -865,7 +950,9 @@ class RpcClient: return start_index = self._current_event_index() start_async_error_index = self._current_async_error_index() - self._wait_for_agent_end(start_index, start_async_error_index, timeout=timeout) + self._wait_for_agent_end( + start_index, start_async_error_index, timeout=timeout + ) finally: self._prompt_lifecycle.release(operation) @@ -875,7 +962,9 @@ class RpcClient: try: start_index = self._current_event_index() start_async_error_index = self._current_async_error_index() - return self._wait_for_agent_end(start_index, start_async_error_index, timeout=timeout) + return self._wait_for_agent_end( + start_index, start_async_error_index, timeout=timeout + ) finally: self._prompt_lifecycle.release(operation) @@ -894,6 +983,7 @@ class RpcClient: with self._event_condition: self._scheduled_agent_runs += 1 self._last_schedule_async_error_index = self._async_errors.current_index() + def _mark_agent_run_completed(self) -> None: with self._event_condition: self._completed_agent_runs += 1 @@ -905,7 +995,9 @@ class RpcClient: def _check_async_errors(self) -> None: with self._event_condition: - errors = self._async_errors.snapshot_from(self._last_schedule_async_error_index) + errors = self._async_errors.snapshot_from( + self._last_schedule_async_error_index + ) if errors: raise errors[0] @@ -934,7 +1026,9 @@ class RpcClient: events=events, messages=final_messages, assistant_message=assistant_message, - assistant_text=assistant_text(assistant_message) if assistant_message is not None else None, + assistant_text=assistant_text(assistant_message) + if assistant_message is not None + else None, ) def _wait_for_agent_end( @@ -966,13 +1060,20 @@ class RpcClient: raise async_errors[0] event_payloads = self._events.snapshot_from(start_index) - if any(payload.get("type") == "agent_end" for payload in event_payloads): - events = tuple(cast(RpcAgentEvent, parse_notification(payload)) for payload in event_payloads) + if any( + payload.get("type") == "agent_end" for payload in event_payloads + ): + events = tuple( + cast(RpcAgentEvent, parse_notification(payload)) + for payload in event_payloads + ) return events remaining = deadline - time.monotonic() if remaining <= 0: - raise RpcTimeoutError(f"Timed out waiting for agent_end. Stderr: {self.stderr}") + raise RpcTimeoutError( + f"Timed out waiting for agent_end. Stderr: {self.stderr}" + ) self._event_condition.wait(remaining) def _request(self, command_type: str, **payload: JsonValue) -> JsonObject: @@ -985,7 +1086,9 @@ class RpcClient: response_queue: queue.Queue[JsonObject | BaseException] = queue.Queue(maxsize=1) with self._state_lock: - self._pending[request_id] = _PendingRequest(command=command_type, response_queue=response_queue) + self._pending[request_id] = _PendingRequest( + command=command_type, response_queue=response_queue + ) try: self._write_json(process, envelope) @@ -999,13 +1102,18 @@ class RpcClient: except queue.Empty as exc: with self._state_lock: self._pending.pop(request_id, None) - raise RpcTimeoutError(f"Timed out waiting for response to {command_type}. Stderr: {self.stderr}") from exc + raise RpcTimeoutError( + f"Timed out waiting for response to {command_type}. Stderr: {self.stderr}" + ) from exc if isinstance(response, BaseException): raise response if not bool(response.get("success", False)): - raise RpcCommandError(command=str(response.get("command", command_type)), error=str(response.get("error", ""))) + raise RpcCommandError( + command=str(response.get("command", command_type)), + error=str(response.get("error", "")), + ) data = response.get("data") if data is None: @@ -1028,27 +1136,51 @@ class RpcClient: tool_name = payload.get("toolName") tool_call_id = payload.get("toolCallId") raw_arguments = payload.get("arguments") - if not isinstance(request_id, str) or not isinstance(tool_name, str) or not isinstance(tool_call_id, str): + if ( + not isinstance(request_id, str) + or not isinstance(tool_name, str) + or not isinstance(tool_call_id, str) + ): return if not isinstance(raw_arguments, Mapping): self._send_notification( { "type": "host_tool_result", "id": request_id, - "result": {"content": [{"type": "text", "text": "Host tool arguments must be an object"}], "details": {}}, + "result": { + "content": [ + { + "type": "text", + "text": "Host tool arguments must be an object", + } + ], + "details": {}, + }, "isError": True, } ) return - tool = next((candidate for candidate in self._custom_tools if candidate.name == tool_name), None) + tool = next( + ( + candidate + for candidate in self._custom_tools + if candidate.name == tool_name + ), + None, + ) if tool is None: self._send_notification( { "type": "host_tool_result", "id": request_id, "result": { - "content": [{"type": "text", "text": f'Host tool "{tool_name}" is not registered'}], + "content": [ + { + "type": "text", + "text": f'Host tool "{tool_name}" is not registered', + } + ], "details": {}, }, "isError": True, @@ -1066,7 +1198,11 @@ class RpcClient: tool_call_id=tool_call_id, _cancel_event=pending_call.cancel_event, _send_update=lambda result: self._send_notification( - {"type": "host_tool_update", "id": request_id, "partialResult": result} + { + "type": "host_tool_update", + "id": request_id, + "partialResult": result, + } ), ) result = tool.execute(params, context) @@ -1086,14 +1222,19 @@ class RpcClient: { "type": "host_tool_result", "id": request_id, - "result": {"content": [{"type": "text", "text": str(exc)}], "details": {}}, + "result": { + "content": [{"type": "text", "text": str(exc)}], + "details": {}, + }, "isError": True, } ) finally: self._pending_host_tool_calls.pop(request_id, None) - threading.Thread(target=run_tool, name=f"omp-rpc-host-tool:{tool_name}", daemon=True).start() + threading.Thread( + target=run_tool, name=f"omp-rpc-host-tool:{tool_name}", daemon=True + ).start() def _handle_host_tool_cancel(self, payload: JsonObject) -> None: target_id = payload.get("targetId") @@ -1117,10 +1258,16 @@ class RpcClient: request_id = payload.get("id") operation = payload.get("operation") url = payload.get("url") - if not isinstance(request_id, str) or not isinstance(operation, str) or not isinstance(url, str): + if ( + not isinstance(request_id, str) + or not isinstance(operation, str) + or not isinstance(url, str) + ): return if operation not in ("read", "write"): - self._send_host_uri_error(request_id, f"Unsupported host URI operation: {operation}") + self._send_host_uri_error( + request_id, f"Unsupported host URI operation: {operation}" + ) return try: @@ -1131,14 +1278,20 @@ class RpcClient: self._send_host_uri_error(request_id, f"Could not parse host URI: {url}") return scheme = (parsed.scheme or "").lower() - uri = next((candidate for candidate in self._host_uris if candidate.scheme == scheme), None) + uri = next( + (candidate for candidate in self._host_uris if candidate.scheme == scheme), + None, + ) if uri is None: - self._send_host_uri_error(request_id, f'Host URI scheme "{scheme}://" is not registered') + self._send_host_uri_error( + request_id, f'Host URI scheme "{scheme}://" is not registered' + ) return if operation == "write" and uri.write is None: self._send_host_uri_error( - request_id, f'Host URI scheme "{scheme}://" was not registered with a write handler' + request_id, + f'Host URI scheme "{scheme}://" was not registered with a write handler', ) return @@ -1147,7 +1300,11 @@ class RpcClient: def run() -> None: try: - context = HostUriContext(url=url, operation=cast(Any, operation), _cancel_event=pending.cancel_event) + context = HostUriContext( + url=url, + operation=cast(Any, operation), + _cancel_event=pending.cancel_event, + ) if operation == "read": value = uri.read(url, context) if pending.cancel_event.is_set(): @@ -1167,7 +1324,9 @@ class RpcClient: uri.write(url, content, context) if pending.cancel_event.is_set(): return - self._send_notification({"type": "host_uri_result", "id": request_id}) + self._send_notification( + {"type": "host_uri_result", "id": request_id} + ) except Exception as exc: if pending.cancel_event.is_set(): return @@ -1175,7 +1334,9 @@ class RpcClient: finally: self._pending_host_uri_requests.pop(request_id, None) - threading.Thread(target=run, name=f"omp-rpc-host-uri:{scheme}:{operation}", daemon=True).start() + threading.Thread( + target=run, name=f"omp-rpc-host-uri:{scheme}:{operation}", daemon=True + ).start() def _handle_host_uri_cancel(self, payload: JsonObject) -> None: target_id = payload.get("targetId") @@ -1185,14 +1346,18 @@ class RpcClient: if pending is not None: pending.cancel_event.set() - def _add_typed_event_listener(self, event_type: str, listener: TEventListener) -> Callable[[], None]: + def _add_typed_event_listener( + self, event_type: str, listener: TEventListener + ) -> Callable[[], None]: listeners = self._typed_event_listeners.setdefault(event_type, []) typed_listener = cast(AgentEventListener, listener) listeners.append(typed_listener) return lambda: self._remove_listener(listeners, typed_listener) @staticmethod - def _normalize_todo_phases(todos: Sequence[TodoSeed | TodoPhaseSeed]) -> list[JsonObject]: + def _normalize_todo_phases( + todos: Sequence[TodoSeed | TodoPhaseSeed], + ) -> list[JsonObject]: if len(todos) == 0: return [] @@ -1206,7 +1371,11 @@ class RpcClient: def normalize_todo_item(seed: TodoSeed) -> JsonObject: if isinstance(seed, str): - return {"id": next_task(), "content": seed, "status": cast(JsonValue, "pending")} + return { + "id": next_task(), + "content": seed, + "status": cast(JsonValue, "pending"), + } if isinstance(seed, TodoItem): if seed.status not in _TODO_STATUS_VALUES: @@ -1234,7 +1403,9 @@ class RpcClient: else: status = "pending" return { - "id": str(raw_id) if isinstance(raw_id, str) and raw_id else next_task(), + "id": str(raw_id) + if isinstance(raw_id, str) and raw_id + else next_task(), "content": content, "status": cast(JsonValue, status), "notes": raw_notes if isinstance(raw_notes, str) else None, @@ -1259,11 +1430,19 @@ class RpcClient: raise RpcError("Todo phases must provide a non-empty 'name' value") phase_id_value = seed.get("id") raw_tasks = seed.get("tasks") or () - if not isinstance(raw_tasks, Sequence) or isinstance(raw_tasks, (str, bytes)): + if not isinstance(raw_tasks, Sequence) or isinstance( + raw_tasks, (str, bytes) + ): raise RpcError("Todo phase 'tasks' must be a sequence") - phase_id = str(phase_id_value) if isinstance(phase_id_value, str) and phase_id_value else f"phase-{index}" + phase_id = ( + str(phase_id_value) + if isinstance(phase_id_value, str) and phase_id_value + else f"phase-{index}" + ) name = raw_name - tasks = [normalize_todo_item(cast(TodoSeed, task)) for task in raw_tasks] + tasks = [ + normalize_todo_item(cast(TodoSeed, task)) for task in raw_tasks + ] return {"id": phase_id, "name": name, "tasks": tasks} @@ -1271,11 +1450,19 @@ class RpcClient: phases: list[JsonObject] = [] for index, seed in enumerate(todos, start=1): if not is_phase_seed(seed): - raise RpcError("Cannot mix flat todo items with todo phases in one set_todos() call") + raise RpcError( + "Cannot mix flat todo items with todo phases in one set_todos() call" + ) phases.append(normalize_phase(cast(TodoPhaseSeed, seed), index)) return phases - return [{"id": "phase-1", "name": "Todos", "tasks": [normalize_todo_item(cast(TodoSeed, todo)) for todo in todos]}] + return [ + { + "id": "phase-1", + "name": "Todos", + "tasks": [normalize_todo_item(cast(TodoSeed, todo)) for todo in todos], + } + ] def _build_command(self) -> tuple[str, ...]: if self._command is not None: @@ -1305,7 +1492,9 @@ class RpcClient: command.append("--no-skills") if self._no_rules: command.append("--no-rules") - emit_no_title = self._no_title if self._no_title is not None else self._rpc_defaults + emit_no_title = ( + self._no_title if self._no_title is not None else self._rpc_defaults + ) if emit_no_title: command.append("--no-title") command.extend(self._extra_args) @@ -1330,7 +1519,9 @@ class RpcClient: process.stdin.write("\n") process.stdin.flush() except (BrokenPipeError, OSError) as exc: - raise RpcProcessExitError(f"Failed to write RPC command: {exc}") from exc + raise RpcProcessExitError( + f"Failed to write RPC command: {exc}" + ) from exc def _read_stdout_loop(self) -> None: process = self._process @@ -1382,7 +1573,12 @@ class RpcClient: if isinstance(notification, ReadyEvent): self._ready_received = True self._ready.set() - self._dispatch_listeners("ready", listener_notification.type, self._ready_listeners, listener_notification) + self._dispatch_listeners( + "ready", + listener_notification.type, + self._ready_listeners, + listener_notification, + ) continue if isinstance(notification, ExtensionUiRequest): @@ -1417,9 +1613,14 @@ class RpcClient: self._append_event(payload) if listener_event.type == "agent_end": self._mark_agent_run_completed() - self._dispatch_listeners("event", listener_event.type, self._event_listeners, listener_event) self._dispatch_listeners( - "typed_event", listener_event.type, self._typed_event_listeners.get(listener_event.type, []), listener_event + "event", listener_event.type, self._event_listeners, listener_event + ) + self._dispatch_listeners( + "typed_event", + listener_event.type, + self._typed_event_listeners.get(listener_event.type, []), + listener_event, ) except Exception as exc: self._mark_closed(exc) @@ -1430,9 +1631,17 @@ class RpcClient: try: exit_code = process.wait(timeout=1.0) except subprocess.TimeoutExpired: - self._mark_closed(RpcProcessExitError("RPC process stdout closed before the process exited")) + self._mark_closed( + RpcProcessExitError( + "RPC process stdout closed before the process exited" + ) + ) return - self._mark_closed(RpcProcessExitError(f"RPC process exited with code {exit_code}. Stderr: {self.stderr}")) + self._mark_closed( + RpcProcessExitError( + f"RPC process exited with code {exit_code}. Stderr: {self.stderr}" + ) + ) def _read_stderr_loop(self) -> None: process = self._process @@ -1478,8 +1687,13 @@ class RpcClient: if protocol_error is None: return - if protocol_error.command in _ASYNC_COMMANDS and protocol_error.remote_error is not None: - self._append_async_error(RpcCommandError(protocol_error.command, protocol_error.remote_error)) + if ( + protocol_error.command in _ASYNC_COMMANDS + and protocol_error.remote_error is not None + ): + self._append_async_error( + RpcCommandError(protocol_error.command, protocol_error.remote_error) + ) self._mark_agent_run_completed() self._record_protocol_error(protocol_error) @@ -1493,7 +1707,11 @@ class RpcClient: return False with self._state_lock: - matching_ids = [request_id for request_id, pending in self._pending.items() if pending.command == command] + matching_ids = [ + request_id + for request_id, pending in self._pending.items() + if pending.command == command + ] target_id: str | None = None if len(matching_ids) == 1: target_id = matching_ids[0] @@ -1528,7 +1746,9 @@ class RpcClient: def _record_protocol_error(self, error: RpcProtocolError) -> None: with self._state_lock: self._protocol_errors.append(error) - self._dispatch_listeners("protocol_error", error.command, self._protocol_error_listeners, error) + self._dispatch_listeners( + "protocol_error", error.command, self._protocol_error_listeners, error + ) def _record_listener_error(self, event: ListenerErrorEvent) -> None: with self._state_lock: diff --git a/python/omp-rpc/src/omp_rpc/host_uris.py b/python/omp-rpc/src/omp_rpc/host_uris.py index 87b470be5..be320f9eb 100644 --- a/python/omp-rpc/src/omp_rpc/host_uris.py +++ b/python/omp-rpc/src/omp_rpc/host_uris.py @@ -9,7 +9,9 @@ from .protocol import JsonObject TPayload = TypeVar("TPayload") -HostUriContentType: TypeAlias = Literal["text/markdown", "application/json", "text/plain"] +HostUriContentType: TypeAlias = Literal[ + "text/markdown", "application/json", "text/plain" +] class HostUriReadResult(TypedDict, total=False): @@ -99,7 +101,9 @@ def normalize_read_result(value: HostUriReadValue) -> JsonObject: if isinstance(value, str): return {"content": value} if not isinstance(value, dict): - raise TypeError("Host URI read handlers must return a string or a HostUriReadResult mapping") + raise TypeError( + "Host URI read handlers must return a string or a HostUriReadResult mapping" + ) payload: JsonObject = {} if "content" not in value: diff --git a/python/omp-rpc/src/omp_rpc/protocol.py b/python/omp-rpc/src/omp_rpc/protocol.py index 0d9b5e5aa..f29182e8c 100644 --- a/python/omp-rpc/src/omp_rpc/protocol.py +++ b/python/omp-rpc/src/omp_rpc/protocol.py @@ -31,24 +31,38 @@ ExtensionUiMethod: TypeAlias = Literal[ "setTitle", "set_editor_text", ] -InteractiveExtensionUiMethod: TypeAlias = Literal["select", "confirm", "input", "editor"] -PassiveExtensionUiMethod: TypeAlias = Literal["notify", "setStatus", "setWidget", "setTitle", "set_editor_text"] +InteractiveExtensionUiMethod: TypeAlias = Literal[ + "select", "confirm", "input", "editor" +] +PassiveExtensionUiMethod: TypeAlias = Literal[ + "notify", "setStatus", "setWidget", "setTitle", "set_editor_text" +] ValueExtensionUiMethod: TypeAlias = Literal["select", "input", "editor"] PASSIVE_EXTENSION_UI_METHODS: Final[frozenset[PassiveExtensionUiMethod]] = frozenset( {"notify", "setStatus", "setWidget", "setTitle", "set_editor_text"} ) -INTERACTIVE_EXTENSION_UI_METHODS: Final[frozenset[InteractiveExtensionUiMethod]] = frozenset( - {"select", "confirm", "input", "editor"} +INTERACTIVE_EXTENSION_UI_METHODS: Final[frozenset[InteractiveExtensionUiMethod]] = ( + frozenset({"select", "confirm", "input", "editor"}) +) +VALUE_EXTENSION_UI_METHODS: Final[frozenset[ValueExtensionUiMethod]] = frozenset( + {"select", "input", "editor"} +) +_THINKING_LEVEL_VALUES: Final[frozenset[str]] = frozenset( + {"off", "minimal", "low", "medium", "high", "xhigh"} ) -VALUE_EXTENSION_UI_METHODS: Final[frozenset[ValueExtensionUiMethod]] = frozenset({"select", "input", "editor"}) -_THINKING_LEVEL_VALUES: Final[frozenset[str]] = frozenset({"off", "minimal", "low", "medium", "high", "xhigh"}) _STEERING_MODE_VALUES: Final[frozenset[str]] = frozenset({"all", "one-at-a-time"}) _INTERRUPT_MODE_VALUES: Final[frozenset[str]] = frozenset({"immediate", "wait"}) -_STOP_REASON_VALUES: Final[frozenset[str]] = frozenset({"stop", "length", "toolUse", "error", "aborted"}) +_STOP_REASON_VALUES: Final[frozenset[str]] = frozenset( + {"stop", "length", "toolUse", "error", "aborted"} +) _NOTIFY_TYPE_VALUES: Final[frozenset[str]] = frozenset({"info", "warning", "error"}) -_WIDGET_PLACEMENT_VALUES: Final[frozenset[str]] = frozenset({"aboveEditor", "belowEditor"}) -_TODO_STATUS_VALUES: Final[frozenset[str]] = frozenset({"pending", "in_progress", "completed", "abandoned"}) +_WIDGET_PLACEMENT_VALUES: Final[frozenset[str]] = frozenset( + {"aboveEditor", "belowEditor"} +) +_TODO_STATUS_VALUES: Final[frozenset[str]] = frozenset( + {"pending", "in_progress", "completed", "abandoned"} +) _EXTENSION_UI_METHOD_VALUES: Final[frozenset[str]] = frozenset( { "select", @@ -94,10 +108,16 @@ _ASSISTANT_MESSAGE_EVENT_TYPE_VALUES: Final[frozenset[str]] = frozenset( "error", } ) -_ASSISTANT_DONE_REASON_VALUES: Final[frozenset[str]] = frozenset({"stop", "length", "toolUse"}) +_ASSISTANT_DONE_REASON_VALUES: Final[frozenset[str]] = frozenset( + {"stop", "length", "toolUse"} +) _ASSISTANT_ERROR_REASON_VALUES: Final[frozenset[str]] = frozenset({"aborted", "error"}) -_AUTO_COMPACTION_REASON_VALUES: Final[frozenset[str]] = frozenset({"threshold", "overflow", "idle"}) -_AUTO_COMPACTION_ACTION_VALUES: Final[frozenset[str]] = frozenset({"context-full", "handoff"}) +_AUTO_COMPACTION_REASON_VALUES: Final[frozenset[str]] = frozenset( + {"threshold", "overflow", "idle"} +) +_AUTO_COMPACTION_ACTION_VALUES: Final[frozenset[str]] = frozenset( + {"context-full", "handoff"} +) def _clone_json_value(value: object, *, field: str) -> JsonValue: @@ -142,7 +162,9 @@ def _require_literal(value: object, allowed: frozenset[str], *, field: str) -> s return value -def _optional_literal(value: object, allowed: frozenset[str], *, field: str) -> str | None: +def _optional_literal( + value: object, allowed: frozenset[str], *, field: str +) -> str | None: if value is None: return None return _require_literal(value, allowed, field=field) @@ -220,7 +242,9 @@ def _tuple_of_strings(values: object, *, field: str) -> tuple[str, ...] | None: def _parse_agent_message(payload: JsonObject, *, field: str) -> AgentMessage: - _require_literal(payload.get("role"), _AGENT_MESSAGE_ROLE_VALUES, field=f"{field}.role") + _require_literal( + payload.get("role"), _AGENT_MESSAGE_ROLE_VALUES, field=f"{field}.role" + ) return cast(AgentMessage, _clone_json_object(payload, field=field)) @@ -246,7 +270,12 @@ def parse_agent_messages(payload: JsonValue | None) -> tuple[AgentMessage, ...]: messages: list[AgentMessage] = [] for index, item in enumerate(payload): - messages.append(_parse_agent_message(_clone_json_object(item, field=f"messages[{index}]"), field=f"messages[{index}]")) + messages.append( + _parse_agent_message( + _clone_json_object(item, field=f"messages[{index}]"), + field=f"messages[{index}]", + ) + ) return tuple(messages) @@ -259,13 +288,17 @@ def parse_assistant_message_event(payload: JsonObject) -> AssistantMessageEvent: if event_type == "start": return AssistantMessageStartEvent( partial=_parse_assistant_message( - _clone_json_object(payload.get("partial"), field="assistantMessageEvent.partial"), + _clone_json_object( + payload.get("partial"), field="assistantMessageEvent.partial" + ), field="assistantMessageEvent.partial", ) ) if event_type in {"text_start", "thinking_start", "toolcall_start"}: partial = _parse_assistant_message( - _clone_json_object(payload.get("partial"), field="assistantMessageEvent.partial"), + _clone_json_object( + payload.get("partial"), field="assistantMessageEvent.partial" + ), field="assistantMessageEvent.partial", ) content_index = _optional_int(payload, "contentIndex") @@ -274,11 +307,15 @@ def parse_assistant_message_event(payload: JsonObject) -> AssistantMessageEvent: if event_type == "text_start": return AssistantTextStartEvent(contentIndex=content_index, partial=partial) if event_type == "thinking_start": - return AssistantThinkingStartEvent(contentIndex=content_index, partial=partial) + return AssistantThinkingStartEvent( + contentIndex=content_index, partial=partial + ) return AssistantToolCallStartEvent(contentIndex=content_index, partial=partial) if event_type in {"text_delta", "thinking_delta", "toolcall_delta"}: partial = _parse_assistant_message( - _clone_json_object(payload.get("partial"), field="assistantMessageEvent.partial"), + _clone_json_object( + payload.get("partial"), field="assistantMessageEvent.partial" + ), field="assistantMessageEvent.partial", ) content_index = _optional_int(payload, "contentIndex") @@ -288,13 +325,21 @@ def parse_assistant_message_event(payload: JsonObject) -> AssistantMessageEvent: if delta is None: raise ValueError("assistantMessageEvent.delta must be a string") if event_type == "text_delta": - return AssistantTextDeltaEvent(contentIndex=content_index, delta=delta, partial=partial) + return AssistantTextDeltaEvent( + contentIndex=content_index, delta=delta, partial=partial + ) if event_type == "thinking_delta": - return AssistantThinkingDeltaEvent(contentIndex=content_index, delta=delta, partial=partial) - return AssistantToolCallDeltaEvent(contentIndex=content_index, delta=delta, partial=partial) + return AssistantThinkingDeltaEvent( + contentIndex=content_index, delta=delta, partial=partial + ) + return AssistantToolCallDeltaEvent( + contentIndex=content_index, delta=delta, partial=partial + ) if event_type in {"text_end", "thinking_end"}: partial = _parse_assistant_message( - _clone_json_object(payload.get("partial"), field="assistantMessageEvent.partial"), + _clone_json_object( + payload.get("partial"), field="assistantMessageEvent.partial" + ), field="assistantMessageEvent.partial", ) content_index = _optional_int(payload, "contentIndex") @@ -304,36 +349,60 @@ def parse_assistant_message_event(payload: JsonObject) -> AssistantMessageEvent: if content is None: raise ValueError("assistantMessageEvent.content must be a string") if event_type == "text_end": - return AssistantTextEndEvent(contentIndex=content_index, content=content, partial=partial) - return AssistantThinkingEndEvent(contentIndex=content_index, content=content, partial=partial) + return AssistantTextEndEvent( + contentIndex=content_index, content=content, partial=partial + ) + return AssistantThinkingEndEvent( + contentIndex=content_index, content=content, partial=partial + ) if event_type == "toolcall_end": partial = _parse_assistant_message( - _clone_json_object(payload.get("partial"), field="assistantMessageEvent.partial"), + _clone_json_object( + payload.get("partial"), field="assistantMessageEvent.partial" + ), field="assistantMessageEvent.partial", ) content_index = _optional_int(payload, "contentIndex") if content_index is None: raise ValueError("assistantMessageEvent.contentIndex must be an integer") - tool_call = _clone_json_object(payload.get("toolCall"), field="assistantMessageEvent.toolCall") - return AssistantToolCallEndEvent(contentIndex=content_index, toolCall=cast(ToolCall, tool_call), partial=partial) + tool_call = _clone_json_object( + payload.get("toolCall"), field="assistantMessageEvent.toolCall" + ) + return AssistantToolCallEndEvent( + contentIndex=content_index, + toolCall=cast(ToolCall, tool_call), + partial=partial, + ) if event_type == "done": return AssistantDoneEvent( reason=cast( Literal["stop", "length", "toolUse"], - _require_literal(payload.get("reason"), _ASSISTANT_DONE_REASON_VALUES, field="assistantMessageEvent.reason"), + _require_literal( + payload.get("reason"), + _ASSISTANT_DONE_REASON_VALUES, + field="assistantMessageEvent.reason", + ), ), message=_parse_assistant_message( - _clone_json_object(payload.get("message"), field="assistantMessageEvent.message"), + _clone_json_object( + payload.get("message"), field="assistantMessageEvent.message" + ), field="assistantMessageEvent.message", ), ) return AssistantErrorEvent( reason=cast( Literal["aborted", "error"], - _require_literal(payload.get("reason"), _ASSISTANT_ERROR_REASON_VALUES, field="assistantMessageEvent.reason"), + _require_literal( + payload.get("reason"), + _ASSISTANT_ERROR_REASON_VALUES, + field="assistantMessageEvent.reason", + ), ), error=_parse_assistant_message( - _clone_json_object(payload.get("error"), field="assistantMessageEvent.error"), + _clone_json_object( + payload.get("error"), field="assistantMessageEvent.error" + ), field="assistantMessageEvent.error", ), ) @@ -984,12 +1053,22 @@ RpcAgentEvent: TypeAlias = ( | TodoAutoClearEvent ) -RpcNotification: TypeAlias = ReadyEvent | ExtensionUiRequest | ExtensionError | RpcAgentEvent | UnknownNotification +RpcNotification: TypeAlias = ( + ReadyEvent + | ExtensionUiRequest + | ExtensionError + | RpcAgentEvent + | UnknownNotification +) def image_from_path(path: str | Path, mime_type: str | None = None) -> ImageContent: file_path = Path(path) - resolved_mime_type = mime_type or mimetypes.guess_type(file_path.name)[0] or "application/octet-stream" + resolved_mime_type = ( + mime_type + or mimetypes.guess_type(file_path.name)[0] + or "application/octet-stream" + ) return { "type": "image", "mimeType": resolved_mime_type, @@ -997,9 +1076,18 @@ def image_from_path(path: str | Path, mime_type: str | None = None) -> ImageCont } -def message_text(message: AgentMessage, *, include_thinking: bool = False) -> str | None: +def message_text( + message: AgentMessage, *, include_thinking: bool = False +) -> str | None: role = message.get("role") - if role not in {"user", "developer", "assistant", "toolResult", "custom", "hookMessage"}: + if role not in { + "user", + "developer", + "assistant", + "toolResult", + "custom", + "hookMessage", + }: return None content = message.get("content") @@ -1015,7 +1103,11 @@ def message_text(message: AgentMessage, *, include_thinking: bool = False) -> st block_type = block.get("type") if block_type == "text" and isinstance(block.get("text"), str): fragments.append(cast(str, block["text"])) - elif include_thinking and block_type == "thinking" and isinstance(block.get("thinking"), str): + elif ( + include_thinking + and block_type == "thinking" + and isinstance(block.get("thinking"), str) + ): fragments.append(cast(str, block["thinking"])) return "".join(fragments) or None @@ -1024,7 +1116,9 @@ def message_text_with_thinking(message: AgentMessage) -> str | None: return message_text(message, include_thinking=True) -def assistant_text(message: AgentMessage, *, include_thinking: bool = False) -> str | None: +def assistant_text( + message: AgentMessage, *, include_thinking: bool = False +) -> str | None: if message.get("role") != "assistant": return None return message_text(message, include_thinking=include_thinking) @@ -1048,7 +1142,8 @@ def parse_model_info(payload: JsonObject | None) -> ModelInfo | None: provider=_require_str(payload, "provider"), base_url=_require_str(payload, "baseUrl"), reasoning=bool(payload.get("reasoning", False)), - input_modalities=_tuple_of_strings(payload.get("input"), field="model.input") or (), + input_modalities=_tuple_of_strings(payload.get("input"), field="model.input") + or (), cost=ModelCost( input=float(cost_payload.get("input", 0.0)), output=float(cost_payload.get("output", 0.0)), @@ -1057,22 +1152,39 @@ def parse_model_info(payload: JsonObject | None) -> ModelInfo | None: ), context_window=int(payload.get("contextWindow", 0)), max_tokens=int(payload.get("maxTokens", 0)), - headers=cast(dict[str, str] | None, _optional_json_object(headers_payload, field="model.headers")), - premium_multiplier=float(payload["premiumMultiplier"]) if "premiumMultiplier" in payload else None, - prefer_websockets=bool(payload["preferWebsockets"]) if "preferWebsockets" in payload else None, + headers=cast( + dict[str, str] | None, + _optional_json_object(headers_payload, field="model.headers"), + ), + premium_multiplier=float(payload["premiumMultiplier"]) + if "premiumMultiplier" in payload + else None, + prefer_websockets=bool(payload["preferWebsockets"]) + if "preferWebsockets" in payload + else None, context_promotion_target=( - str(payload["contextPromotionTarget"]) if "contextPromotionTarget" in payload else None + str(payload["contextPromotionTarget"]) + if "contextPromotionTarget" in payload + else None ), priority=int(payload["priority"]) if "priority" in payload else None, thinking=( ThinkingConfig( min_level=cast( ThinkingLevel, - _require_literal(thinking_payload.get("minLevel"), _THINKING_LEVEL_VALUES, field="model.thinking.minLevel"), + _require_literal( + thinking_payload.get("minLevel"), + _THINKING_LEVEL_VALUES, + field="model.thinking.minLevel", + ), ), max_level=cast( ThinkingLevel, - _require_literal(thinking_payload.get("maxLevel"), _THINKING_LEVEL_VALUES, field="model.thinking.maxLevel"), + _require_literal( + thinking_payload.get("maxLevel"), + _THINKING_LEVEL_VALUES, + field="model.thinking.maxLevel", + ), ), mode=_require_str(cast(JsonObject, thinking_payload), "mode"), ) @@ -1087,7 +1199,9 @@ def parse_tool_descriptor(payload: JsonObject) -> ToolDescriptor: return ToolDescriptor( name=_require_str(payload, "name"), description=_require_str(payload, "description"), - parameters=_clone_json_value(payload.get("parameters"), field="tool.parameters"), + parameters=_clone_json_value( + payload.get("parameters"), field="tool.parameters" + ), ) @@ -1097,7 +1211,11 @@ def parse_todo_item(payload: JsonObject) -> TodoItem: content=_require_str(payload, "content"), status=cast( TodoStatus, - _require_literal(payload.get("status", "pending"), _TODO_STATUS_VALUES, field="todo.status"), + _require_literal( + payload.get("status", "pending"), + _TODO_STATUS_VALUES, + field="todo.status", + ), ), notes=_optional_str(payload, "notes"), details=_optional_str(payload, "details"), @@ -1111,7 +1229,10 @@ def parse_todo_phase(payload: JsonObject) -> TodoPhase: else: if not isinstance(raw_tasks, list): raise ValueError("tasks must be a list") - tasks = tuple(parse_todo_item(_clone_json_object(item, field="tasks[]")) for item in raw_tasks) + tasks = tuple( + parse_todo_item(_clone_json_object(item, field="tasks[]")) + for item in raw_tasks + ) return TodoPhase( id=str(payload.get("id", "")), name=_require_str(payload, "name"), @@ -1127,27 +1248,44 @@ def parse_todo_phases(payload: JsonValue | None) -> tuple[TodoPhase, ...]: def parse_session_state(payload: JsonObject) -> SessionState: dump_tools = tuple( - parse_tool_descriptor(_clone_json_object(item, field="dumpTools[]")) for item in cast(list[Any], payload.get("dumpTools") or []) + parse_tool_descriptor(_clone_json_object(item, field="dumpTools[]")) + for item in cast(list[Any], payload.get("dumpTools") or []) ) return SessionState( model=parse_model_info(cast(JsonObject | None, payload.get("model"))), thinking_level=cast( ThinkingLevel | None, - _optional_literal(payload.get("thinkingLevel"), _THINKING_LEVEL_VALUES, field="thinkingLevel"), + _optional_literal( + payload.get("thinkingLevel"), + _THINKING_LEVEL_VALUES, + field="thinkingLevel", + ), ), is_streaming=bool(payload.get("isStreaming", False)), is_compacting=bool(payload.get("isCompacting", False)), steering_mode=cast( SteeringMode, - _require_literal(payload.get("steeringMode", "one-at-a-time"), _STEERING_MODE_VALUES, field="steeringMode"), + _require_literal( + payload.get("steeringMode", "one-at-a-time"), + _STEERING_MODE_VALUES, + field="steeringMode", + ), ), follow_up_mode=cast( SteeringMode, - _require_literal(payload.get("followUpMode", "one-at-a-time"), _STEERING_MODE_VALUES, field="followUpMode"), + _require_literal( + payload.get("followUpMode", "one-at-a-time"), + _STEERING_MODE_VALUES, + field="followUpMode", + ), ), interrupt_mode=cast( InterruptMode, - _require_literal(payload.get("interruptMode", "immediate"), _INTERRUPT_MODE_VALUES, field="interruptMode"), + _require_literal( + payload.get("interruptMode", "immediate"), + _INTERRUPT_MODE_VALUES, + field="interruptMode", + ), ), session_file=_optional_str(payload, "sessionFile"), session_id=_require_str(payload, "sessionId"), @@ -1155,7 +1293,9 @@ def parse_session_state(payload: JsonObject) -> SessionState: auto_compaction_enabled=bool(payload.get("autoCompactionEnabled", False)), message_count=int(payload.get("messageCount", 0)), queued_message_count=int(payload.get("queuedMessageCount", 0)), - todo_phases=parse_todo_phases(cast(JsonValue | None, payload.get("todoPhases"))), + todo_phases=parse_todo_phases( + cast(JsonValue | None, payload.get("todoPhases")) + ), system_prompt=_optional_str_list(payload, "systemPrompt"), dump_tools=dump_tools, ) @@ -1181,8 +1321,12 @@ def parse_compaction_result(payload: JsonObject) -> CompactionResult: short_summary=_optional_str(payload, "shortSummary"), first_kept_entry_id=str(payload.get("firstKeptEntryId", "")), tokens_before=int(payload.get("tokensBefore", 0)), - details=_clone_json_value(payload.get("details"), field="compaction.details") if "details" in payload else None, - preserve_data=_optional_json_object(payload.get("preserveData"), field="compaction.preserveData"), + details=_clone_json_value(payload.get("details"), field="compaction.details") + if "details" in payload + else None, + preserve_data=_optional_json_object( + payload.get("preserveData"), field="compaction.preserveData" + ), ) @@ -1199,7 +1343,9 @@ def parse_model_cycle_result(payload: JsonObject | None) -> ModelCycleResult | N ) -def parse_thinking_level_cycle_result(payload: JsonObject | None) -> ThinkingLevelCycleResult | None: +def parse_thinking_level_cycle_result( + payload: JsonObject | None, +) -> ThinkingLevelCycleResult | None: if payload is None or payload.get("level") is None: return None return ThinkingLevelCycleResult(level=cast(ThinkingLevel, payload["level"])) @@ -1211,7 +1357,10 @@ def parse_cancellation_result(payload: JsonObject | None) -> CancellationResult: def parse_branch_result(payload: JsonObject | None) -> BranchResult: payload = payload or {} - return BranchResult(text=str(payload.get("text", "")), cancelled=bool(payload.get("cancelled", False))) + return BranchResult( + text=str(payload.get("text", "")), + cancelled=bool(payload.get("cancelled", False)), + ) def parse_branch_messages(payload: JsonObject | None) -> tuple[BranchMessage, ...]: @@ -1220,7 +1369,9 @@ def parse_branch_messages(payload: JsonObject | None) -> tuple[BranchMessage, .. raise ValueError("messages must be a list") return tuple( BranchMessage( - entry_id=str(_clone_json_object(item, field="messages[]").get("entryId", "")), + entry_id=str( + _clone_json_object(item, field="messages[]").get("entryId", "") + ), text=str(_clone_json_object(item, field="messages[]").get("text", "")), ) for item in messages @@ -1228,7 +1379,9 @@ def parse_branch_messages(payload: JsonObject | None) -> tuple[BranchMessage, .. def parse_session_stats(payload: JsonObject) -> SessionStats: - tokens_payload = _optional_json_object(payload.get("tokens"), field="sessionStats.tokens") or {} + tokens_payload = ( + _optional_json_object(payload.get("tokens"), field="sessionStats.tokens") or {} + ) return SessionStats( session_file=_optional_str(payload, "sessionFile"), session_id=str(payload.get("sessionId", "")), @@ -1254,10 +1407,16 @@ def parse_extension_ui_request(payload: JsonObject) -> ExtensionUiRequest: id=_require_str(payload, "id"), method=cast( ExtensionUiMethod, - _require_literal(payload.get("method"), _EXTENSION_UI_METHOD_VALUES, field="extension_ui_request.method"), + _require_literal( + payload.get("method"), + _EXTENSION_UI_METHOD_VALUES, + field="extension_ui_request.method", + ), ), title=_optional_str(payload, "title"), - options=_tuple_of_strings(payload.get("options"), field="extension_ui_request.options"), + options=_tuple_of_strings( + payload.get("options"), field="extension_ui_request.options" + ), message=_optional_str(payload, "message"), placeholder=_optional_str(payload, "placeholder"), prefill=_optional_str(payload, "prefill"), @@ -1266,12 +1425,18 @@ def parse_extension_ui_request(payload: JsonObject) -> ExtensionUiRequest: target_id=_optional_str(payload, "targetId"), notify_type=cast( NotifyType | None, - _optional_literal(payload.get("notifyType"), _NOTIFY_TYPE_VALUES, field="extension_ui_request.notifyType"), + _optional_literal( + payload.get("notifyType"), + _NOTIFY_TYPE_VALUES, + field="extension_ui_request.notifyType", + ), ), status_key=_optional_str(payload, "statusKey"), status_text=_optional_str(payload, "statusText"), widget_key=_optional_str(payload, "widgetKey"), - widget_lines=_tuple_of_strings(payload.get("widgetLines"), field="extension_ui_request.widgetLines"), + widget_lines=_tuple_of_strings( + payload.get("widgetLines"), field="extension_ui_request.widgetLines" + ), widget_placement=cast( WidgetPlacement | None, _optional_literal( @@ -1303,7 +1468,11 @@ def parse_notification(payload: JsonObject) -> RpcNotification: if event_type == "agent_start": return AgentStartEvent() if event_type == "agent_end": - return AgentEndEvent(messages=parse_agent_messages(cast(JsonValue | None, payload.get("messages")))) + return AgentEndEvent( + messages=parse_agent_messages( + cast(JsonValue | None, payload.get("messages")) + ) + ) if event_type == "turn_start": return TurnStartEvent() if event_type == "turn_end": @@ -1313,25 +1482,35 @@ def parse_notification(payload: JsonObject) -> RpcNotification: field="turn_end.message", ), tool_results=tuple( - _parse_tool_result_message(_clone_json_object(item, field="turn_end.toolResults[]"), field="turn_end.toolResults[]") + _parse_tool_result_message( + _clone_json_object(item, field="turn_end.toolResults[]"), + field="turn_end.toolResults[]", + ) for item in cast(list[Any], payload.get("toolResults") or []) ), ) if event_type == "message_start": return MessageStartEvent( message=_parse_agent_message( - _clone_json_object(payload.get("message"), field="message_start.message"), + _clone_json_object( + payload.get("message"), field="message_start.message" + ), field="message_start.message", ) ) if event_type == "message_update": return MessageUpdateEvent( message=_parse_agent_message( - _clone_json_object(payload.get("message"), field="message_update.message"), + _clone_json_object( + payload.get("message"), field="message_update.message" + ), field="message_update.message", ), assistant_message_event=parse_assistant_message_event( - _clone_json_object(payload.get("assistantMessageEvent"), field="message_update.assistantMessageEvent") + _clone_json_object( + payload.get("assistantMessageEvent"), + field="message_update.assistantMessageEvent", + ) ), ) if event_type == "message_end": @@ -1345,16 +1524,27 @@ def parse_notification(payload: JsonObject) -> RpcNotification: return ToolExecutionStartEvent( tool_call_id=str(payload.get("toolCallId", "")), tool_name=str(payload.get("toolName", "")), - args=_clone_json_value(payload.get("args"), field="tool_execution_start.args") if "args" in payload else None, + args=_clone_json_value( + payload.get("args"), field="tool_execution_start.args" + ) + if "args" in payload + else None, intent=_optional_str(payload, "intent"), ) if event_type == "tool_execution_update": return ToolExecutionUpdateEvent( tool_call_id=str(payload.get("toolCallId", "")), tool_name=str(payload.get("toolName", "")), - args=_clone_json_value(payload.get("args"), field="tool_execution_update.args") if "args" in payload else None, + args=_clone_json_value( + payload.get("args"), field="tool_execution_update.args" + ) + if "args" in payload + else None, partial_result=( - _clone_json_value(payload.get("partialResult"), field="tool_execution_update.partialResult") + _clone_json_value( + payload.get("partialResult"), + field="tool_execution_update.partialResult", + ) if "partialResult" in payload else None ), @@ -1363,18 +1553,30 @@ def parse_notification(payload: JsonObject) -> RpcNotification: return ToolExecutionEndEvent( tool_call_id=str(payload.get("toolCallId", "")), tool_name=str(payload.get("toolName", "")), - result=_clone_json_value(payload.get("result"), field="tool_execution_end.result") if "result" in payload else None, + result=_clone_json_value( + payload.get("result"), field="tool_execution_end.result" + ) + if "result" in payload + else None, is_error=_optional_bool(payload, "isError"), ) if event_type == "auto_compaction_start": return AutoCompactionStartEvent( reason=cast( Literal["threshold", "overflow", "idle"], - _require_literal(payload.get("reason", "threshold"), _AUTO_COMPACTION_REASON_VALUES, field="auto_compaction_start.reason"), + _require_literal( + payload.get("reason", "threshold"), + _AUTO_COMPACTION_REASON_VALUES, + field="auto_compaction_start.reason", + ), ), action=cast( Literal["context-full", "handoff"], - _require_literal(payload.get("action", "context-full"), _AUTO_COMPACTION_ACTION_VALUES, field="auto_compaction_start.action"), + _require_literal( + payload.get("action", "context-full"), + _AUTO_COMPACTION_ACTION_VALUES, + field="auto_compaction_start.action", + ), ), ) if event_type == "auto_compaction_end": @@ -1382,10 +1584,18 @@ def parse_notification(payload: JsonObject) -> RpcNotification: return AutoCompactionEndEvent( action=cast( Literal["context-full", "handoff"], - _require_literal(payload.get("action", "context-full"), _AUTO_COMPACTION_ACTION_VALUES, field="auto_compaction_end.action"), + _require_literal( + payload.get("action", "context-full"), + _AUTO_COMPACTION_ACTION_VALUES, + field="auto_compaction_end.action", + ), ), result=( - parse_compaction_result(_clone_json_object(result_payload, field="auto_compaction_end.result")) + parse_compaction_result( + _clone_json_object( + result_payload, field="auto_compaction_end.result" + ) + ) if result_payload is not None else None ), @@ -1414,9 +1624,15 @@ def parse_notification(payload: JsonObject) -> RpcNotification: role=str(payload.get("role", "")), ) if event_type == "retry_fallback_succeeded": - return RetryFallbackSucceededEvent(model=str(payload.get("model", "")), role=str(payload.get("role", ""))) + return RetryFallbackSucceededEvent( + model=str(payload.get("model", "")), role=str(payload.get("role", "")) + ) if event_type == "ttsr_triggered": - return TtsrTriggeredEvent(rules=_clone_json_objects(payload.get("rules"), field="ttsr_triggered.rules")) + return TtsrTriggeredEvent( + rules=_clone_json_objects( + payload.get("rules"), field="ttsr_triggered.rules" + ) + ) if event_type == "todo_reminder": return TodoReminderEvent( todos=tuple( @@ -1428,4 +1644,6 @@ def parse_notification(payload: JsonObject) -> RpcNotification: ) if event_type == "todo_auto_clear": return TodoAutoClearEvent() - return UnknownNotification(payload=_clone_json_object(payload, field="notification")) + return UnknownNotification( + payload=_clone_json_object(payload, field="notification") + ) diff --git a/python/omp-rpc/tests/test_client.py b/python/omp-rpc/tests/test_client.py index d46f005fc..b62696d9c 100644 --- a/python/omp-rpc/tests/test_client.py +++ b/python/omp-rpc/tests/test_client.py @@ -576,7 +576,12 @@ BROKEN_STARTUP_SERVER = textwrap.dedent( class RpcClientTests(unittest.TestCase): def make_client(self, server: str = FAKE_SERVER, **kwargs: object) -> RpcClient: - return RpcClient(command=[sys.executable, "-u", "-c", server], startup_timeout=2.0, request_timeout=2.0, **kwargs) + return RpcClient( + command=[sys.executable, "-u", "-c", server], + startup_timeout=2.0, + request_timeout=2.0, + **kwargs, + ) def test_command_builder_supports_common_rpc_options(self) -> None: client = RpcClient( @@ -622,7 +627,9 @@ class RpcClientTests(unittest.TestCase): with self.make_client() as client: state = client.get_state() self.assertEqual(state.session_id, "fake-session") - self.assertEqual(state.model.id if state.model else None, "claude-sonnet-4-5") + self.assertEqual( + state.model.id if state.model else None, "claude-sonnet-4-5" + ) result = client.bash("echo hello") self.assertEqual(result.output, "hello\n") @@ -658,11 +665,21 @@ class RpcClientTests(unittest.TestCase): self.assertEqual(state.dump_tools[-1].name, "echo_host") turn = client.prompt_and_wait("needs host tool", timeout=2.0) - update_events = [event for event in turn.events if getattr(event, "type", None) == "tool_execution_update"] - end_events = [event for event in turn.events if getattr(event, "type", None) == "tool_execution_end"] + update_events = [ + event + for event in turn.events + if getattr(event, "type", None) == "tool_execution_update" + ] + end_events = [ + event + for event in turn.events + if getattr(event, "type", None) == "tool_execution_end" + ] self.assertEqual(len(update_events), 1) - self.assertEqual(update_events[0].partial_result["content"][0]["text"], "working:hello") + self.assertEqual( + update_events[0].partial_result["content"][0]["text"], "working:hello" + ) self.assertEqual(len(end_events), 1) self.assertEqual(end_events[0].result["content"][0]["text"], "host:hello") @@ -679,7 +696,9 @@ class RpcClientTests(unittest.TestCase): seen_methods: list[str] = [] with self.make_client() as client: - client.install_headless_ui(on_request=lambda request: seen_methods.append(request.method)) + client.install_headless_ui( + on_request=lambda request: seen_methods.append(request.method) + ) client.prompt_and_wait("needs ui", timeout=2.0) self.assertEqual(seen_methods, ["input"]) @@ -690,7 +709,9 @@ class RpcClientTests(unittest.TestCase): notification_types: list[str] = [] client = self.make_client() client.on_ready(lambda event: ready_types.append(event.type)) - client.on_notification(lambda notification: notification_types.append(notification.type)) + client.on_notification( + lambda notification: notification_types.append(notification.type) + ) client.on_turn_start(lambda event: event_types.append(event.type)) client.on_message_update(lambda event: event_types.append(event.type)) client.on_agent_end(lambda event: event_types.append(event.type)) @@ -729,7 +750,10 @@ class RpcClientTests(unittest.TestCase): self.assertEqual(cycled.model.id, "claude-sonnet-4-5") available = client.get_available_models() - self.assertEqual([item.id for item in available], ["claude-sonnet-4-5", "claude-sonnet-4-6"]) + self.assertEqual( + [item.id for item in available], + ["claude-sonnet-4-5", "claude-sonnet-4-6"], + ) client.set_thinking_level("high") self.assertEqual(client.get_state().thinking_level, "high") @@ -853,8 +877,12 @@ class RpcClientTests(unittest.TestCase): seen_unknown: list[str] = [] with self.make_client() as client: - client.on_extension_error(lambda event: seen_extension_errors.append(event.error)) - client.on_unknown_notification(lambda event: seen_unknown.append(str(event.payload.get("type")))) + client.on_extension_error( + lambda event: seen_extension_errors.append(event.error) + ) + client.on_unknown_notification( + lambda event: seen_unknown.append(str(event.payload.get("type"))) + ) client.prompt_and_wait("notifications", timeout=2.0) self.assertEqual(seen_extension_errors, ["boom"]) @@ -879,20 +907,32 @@ class RpcClientTests(unittest.TestCase): errors: list[BaseException] = [] with self.make_client() as client: + def run_prompt() -> None: try: - results.append(client.prompt_and_wait("slow", timeout=2.0).require_assistant_text()) - except BaseException as exc: # pragma: no cover - defensive thread capture + results.append( + client.prompt_and_wait( + "slow", timeout=2.0 + ).require_assistant_text() + ) + except ( + BaseException + ) as exc: # pragma: no cover - defensive thread capture errors.append(exc) thread = threading.Thread(target=run_prompt) thread.start() deadline = time.time() + 1.0 - while client._prompt_lifecycle.active_operation != "prompt_and_wait" and time.time() < deadline: + while ( + client._prompt_lifecycle.active_operation != "prompt_and_wait" + and time.time() < deadline + ): time.sleep(0.01) - self.assertEqual(client._prompt_lifecycle.active_operation, "prompt_and_wait") + self.assertEqual( + client._prompt_lifecycle.active_operation, "prompt_and_wait" + ) with self.assertRaises(RpcConcurrencyError): client.collect_events(timeout=1.0) @@ -904,7 +944,11 @@ class RpcClientTests(unittest.TestCase): def test_listener_mutation_does_not_change_retained_turn(self) -> None: with self.make_client() as client: - client.on_message_end(lambda event: event.message["content"].__setitem__(0, {"type": "text", "text": "mutated"})) + client.on_message_end( + lambda event: event.message["content"].__setitem__( + 0, {"type": "text", "text": "mutated"} + ) + ) turn = client.prompt_and_wait("say hello", timeout=2.0) messages = client.get_messages() @@ -941,12 +985,16 @@ class RpcClientTests(unittest.TestCase): listener_errors: list[tuple[str, str | None, str]] = [] client = self.make_client() client.on_notification( - lambda notification: (_ for _ in ()).throw(RuntimeError("boom")) - if notification.type == "turn_start" - else None + lambda notification: ( + (_ for _ in ()).throw(RuntimeError("boom")) + if notification.type == "turn_start" + else None + ) ) client.on_listener_error( - lambda event: listener_errors.append((event.listener_kind, event.source_type, str(event.error))) + lambda event: listener_errors.append( + (event.listener_kind, event.source_type, str(event.error)) + ) ) try: @@ -986,7 +1034,6 @@ class RpcClientTests(unittest.TestCase): self.assertIn("max_event_history", str(ctx.exception)) - HANGING_SERVER = textwrap.dedent( """ import json @@ -1051,22 +1098,32 @@ class StopUnblocksPromptAndWaitTests(unittest.TestCase): # Wait until the prompt is in flight. deadline = time.time() + 2.0 - while client._prompt_lifecycle.active_operation != "prompt_and_wait" and time.time() < deadline: + while ( + client._prompt_lifecycle.active_operation != "prompt_and_wait" + and time.time() < deadline + ): time.sleep(0.01) - self.assertEqual(client._prompt_lifecycle.active_operation, "prompt_and_wait") + self.assertEqual( + client._prompt_lifecycle.active_operation, "prompt_and_wait" + ) t0 = time.time() client.stop() thread.join(timeout=2.0) elapsed = time.time() - t0 - self.assertFalse(thread.is_alive(), "prompt_and_wait did not return after stop()") - self.assertLess(elapsed, 2.0, f"stop() took {elapsed:.2f}s to unblock prompt_and_wait") + self.assertFalse( + thread.is_alive(), "prompt_and_wait did not return after stop()" + ) + self.assertLess( + elapsed, 2.0, f"stop() took {elapsed:.2f}s to unblock prompt_and_wait" + ) self.assertEqual(len(errors), 1) self.assertIsInstance(errors[0], RpcProcessExitError) finally: # stop() is idempotent; safe to call again on cleanup paths. client.stop() + if __name__ == "__main__": unittest.main() diff --git a/python/omp-rpc/tests/test_host_uris.py b/python/omp-rpc/tests/test_host_uris.py index 486d96a11..dc01799ec 100644 --- a/python/omp-rpc/tests/test_host_uris.py +++ b/python/omp-rpc/tests/test_host_uris.py @@ -2,12 +2,11 @@ from __future__ import annotations import sys import textwrap -import threading import time import unittest from omp_rpc import RpcClient, host_uri -from omp_rpc.host_uris import HostUri, normalize_read_result +from omp_rpc.host_uris import normalize_read_result URI_SERVER = textwrap.dedent( @@ -114,7 +113,9 @@ class HostUriHelperTests(unittest.TestCase): def test_normalize_read_result_rejects_invalid_content_type(self) -> None: with self.assertRaises(ValueError): - normalize_read_result({"content": "x", "content_type": "application/octet-stream"}) # type: ignore[arg-type] + normalize_read_result( + {"content": "x", "content_type": "application/octet-stream"} + ) # type: ignore[arg-type] def test_host_uri_helper_normalizes_scheme(self) -> None: uri = host_uri(scheme=" DB ", read=lambda url, ctx: "x") @@ -125,7 +126,9 @@ class HostUriHelperTests(unittest.TestCase): host_uri(scheme="", read=lambda url, ctx: "x") def test_host_uri_writable_when_write_supplied(self) -> None: - uri = host_uri(scheme="db", read=lambda url, ctx: "x", write=lambda url, content, ctx: None) + uri = host_uri( + scheme="db", read=lambda url, ctx: "x", write=lambda url, content, ctx: None + ) self.assertTrue(uri.writable) @@ -168,7 +171,9 @@ class RpcHostUriBridgeTests(unittest.TestCase): "immutable": True, } - with self._make_client(host_uris=(host_uri(scheme="db", read=read_db),)) as client: + with self._make_client( + host_uris=(host_uri(scheme="db", read=read_db),) + ) as client: client._request("trigger_read", url="db://users/42") # type: ignore[attr-defined] frame = self._await_echo(client) self.assertEqual(frame["content"], '{"name":"Alice"}') @@ -212,7 +217,9 @@ class RpcHostUriBridgeTests(unittest.TestCase): def read_db(_url: str, _ctx) -> str: raise RuntimeError("boom") - with self._make_client(host_uris=(host_uri(scheme="db", read=read_db),)) as client: + with self._make_client( + host_uris=(host_uri(scheme="db", read=read_db),) + ) as client: client._request("trigger_read", url="db://users/42") # type: ignore[attr-defined] frame = self._await_echo(client) self.assertTrue(frame.get("isError")) diff --git a/python/omp-rpc/tests/test_protocol.py b/python/omp-rpc/tests/test_protocol.py index c26519fcd..ff641f19b 100644 --- a/python/omp-rpc/tests/test_protocol.py +++ b/python/omp-rpc/tests/test_protocol.py @@ -207,7 +207,9 @@ class ProtocolParsingTests(unittest.TestCase): ) self.assertEqual(state.system_prompt, ()) - def test_parse_session_state_rejects_non_string_in_system_prompt_array(self) -> None: + def test_parse_session_state_rejects_non_string_in_system_prompt_array( + self, + ) -> None: with self.assertRaises(ValueError): parse_session_state( { @@ -233,7 +235,9 @@ class ProtocolParsingTests(unittest.TestCase): def test_parse_extension_ui_request_rejects_invalid_method(self) -> None: with self.assertRaises(ValueError): - parse_notification({"type": "extension_ui_request", "id": "ui-1", "method": "launch"}) + parse_notification( + {"type": "extension_ui_request", "id": "ui-1", "method": "launch"} + ) def test_parse_message_update_rejects_invalid_assistant_done_reason(self) -> None: with self.assertRaises(ValueError): diff --git a/python/omp-rpc/tests/test_user_group.py b/python/omp-rpc/tests/test_user_group.py index 8f6b41fdc..13bb1e9c5 100644 --- a/python/omp-rpc/tests/test_user_group.py +++ b/python/omp-rpc/tests/test_user_group.py @@ -13,7 +13,9 @@ class _Sentinel(Exception): def _start_and_capture(**kwargs): client = RpcClient(**kwargs) - with patch("omp_rpc.client.subprocess.Popen", side_effect=_Sentinel("aborted")) as mock_popen: + with patch( + "omp_rpc.client.subprocess.Popen", side_effect=_Sentinel("aborted") + ) as mock_popen: with pytest.raises(_Sentinel): client.start() assert mock_popen.call_count == 1