From 67e12b39fe19c5d60179e2052c63c6624edb2b7f Mon Sep 17 00:00:00 2001 From: can1357 Date: Wed, 8 Apr 2026 04:20:06 +0200 Subject: [PATCH] feat(omp-rpc): introduced RpcClient with subprocess communication and event listeners - Added RpcClient class with subprocess-based RPC communication, thread-safe request/response handling, and 30+ public methods for agent operations. - Defined 65+ public API symbols including protocol types, dataclasses, parser functions, and custom exception types for RPC communication. - Implemented event listener system with on_event, on_ui_request, and on_extension_error callbacks for asynchronous RPC event handling. - Added comprehensive test suite with integration tests for RpcClient and unit tests for protocol parsing functions. - Configured Python package with setuptools, PEP 561 type hints support, and documentation for usage and API reference. --- python/omp-rpc/README.md | 70 ++ python/omp-rpc/pyproject.toml | 36 + python/omp-rpc/src/omp_rpc/__init__.py | 129 ++++ python/omp-rpc/src/omp_rpc/client.py | 600 +++++++++++++++ python/omp-rpc/src/omp_rpc/protocol.py | 963 +++++++++++++++++++++++++ python/omp-rpc/src/omp_rpc/py.typed | 1 + python/omp-rpc/tests/__init__.py | 6 + python/omp-rpc/tests/test_client.py | 220 ++++++ python/omp-rpc/tests/test_protocol.py | 123 ++++ 9 files changed, 2148 insertions(+) create mode 100644 python/omp-rpc/README.md create mode 100644 python/omp-rpc/pyproject.toml create mode 100644 python/omp-rpc/src/omp_rpc/__init__.py create mode 100644 python/omp-rpc/src/omp_rpc/client.py create mode 100644 python/omp-rpc/src/omp_rpc/protocol.py create mode 100644 python/omp-rpc/src/omp_rpc/py.typed create mode 100644 python/omp-rpc/tests/__init__.py create mode 100644 python/omp-rpc/tests/test_client.py create mode 100644 python/omp-rpc/tests/test_protocol.py diff --git a/python/omp-rpc/README.md b/python/omp-rpc/README.md new file mode 100644 index 000000000..6fe04c8e1 --- /dev/null +++ b/python/omp-rpc/README.md @@ -0,0 +1,70 @@ +# omp-rpc + +Typed Python bindings for the `omp --mode rpc` protocol used by the coding agent. + +This package wraps the newline-delimited JSON RPC transport exposed by the CLI and +provides: + +- typed command methods for the stable RPC surface +- typed protocol models for state, bash results, compaction, and session stats +- a process-backed client that manages request correlation over stdio +- helpers for collecting prompt runs and handling extension UI requests + +## Basic Usage + +```python +from omp_rpc import RpcClient + +with RpcClient(provider="anthropic", model="claude-sonnet-4-5") as client: + state = client.get_state() + print(state.model.id if state.model else "no model") + + turn = client.prompt_and_wait("Reply with just the word hello") + print(turn.require_assistant_text()) +``` + +By default the client runs: + +```bash +omp --mode rpc +``` + +You can also point it at a custom command, which is useful inside this repo while +developing against the Bun entrypoint: + +```python +from omp_rpc import RpcClient + +with RpcClient( + command=[ + "bun", + "packages/coding-agent/src/cli.ts", + "--mode", + "rpc", + "--provider", + "anthropic", + "--model", + "claude-sonnet-4-5", + ], +) as client: + print(client.get_state().session_id) +``` + +## Extension UI Requests + +Extensions in RPC mode can ask the host for input. Those requests are available as +typed `ExtensionUiRequest` instances: + +```python +request = client.next_ui_request(timeout=5.0) + +if request.method == "confirm": + client.send_ui_confirmation(request.id, True) +elif request.method in {"input", "editor"}: + client.send_ui_value(request.id, "approved") +``` + +## Protocol Reference + +The canonical wire protocol still lives in the repo at +[`docs/rpc.md`](../../docs/rpc.md). diff --git a/python/omp-rpc/pyproject.toml b/python/omp-rpc/pyproject.toml new file mode 100644 index 000000000..0a0802546 --- /dev/null +++ b/python/omp-rpc/pyproject.toml @@ -0,0 +1,36 @@ +[build-system] +requires = ["setuptools>=69"] +build-backend = "setuptools.build_meta" + +[project] +name = "omp-rpc" +version = "0.1.0" +description = "Typed Python client for the omp coding-agent RPC protocol" +readme = "README.md" +requires-python = ">=3.11" +license = { text = "MIT" } +authors = [{ name = "OpenAI Codex" }] +keywords = ["omp", "rpc", "agent", "coding-agent", "jsonl", "stdio"] +classifiers = [ + "Development Status :: 3 - Alpha", + "Intended Audience :: Developers", + "License :: OSI Approved :: MIT License", + "Programming Language :: Python :: 3", + "Programming Language :: Python :: 3.11", + "Programming Language :: Python :: 3.12", + "Topic :: Software Development :: Libraries :: Python Modules", + "Topic :: Software Development :: Libraries", +] + +[project.urls] +Repository = "https://github.com/can1357/oh-my-pi" +Documentation = "https://github.com/can1357/oh-my-pi/blob/main/docs/rpc.md" + +[tool.setuptools] +package-dir = { "" = "src" } + +[tool.setuptools.packages.find] +where = ["src"] + +[tool.setuptools.package-data] +omp_rpc = ["py.typed"] diff --git a/python/omp-rpc/src/omp_rpc/__init__.py b/python/omp-rpc/src/omp_rpc/__init__.py new file mode 100644 index 000000000..a68f6ea3e --- /dev/null +++ b/python/omp-rpc/src/omp_rpc/__init__.py @@ -0,0 +1,129 @@ +from .client import ( + AgentEventListener, + ExtensionErrorListener, + PromptTurn, + RpcClient, + RpcCommandError, + RpcError, + RpcProcessExitError, + RpcTimeoutError, + UiRequestListener, +) +from .protocol import ( + AgentEndEvent, + AgentMessage, + AssistantMessage, + AssistantMessageEvent, + AutoCompactionEndEvent, + AutoCompactionStartEvent, + AutoRetryEndEvent, + AutoRetryStartEvent, + BashResult, + BranchMessage, + BranchResult, + CancellationResult, + CompactionResult, + CompactionSummaryMessage, + CustomMessage, + DeveloperMessage, + ExtensionError, + ExtensionUiRequest, + FileMentionMessage, + HookMessage, + ImageContent, + MessageEndEvent, + MessageStartEvent, + MessageUpdateEvent, + ModelCycleResult, + ModelInfo, + ModelCost, + PythonExecutionMessage, + ReadyEvent, + RetryFallbackAppliedEvent, + RetryFallbackSucceededEvent, + RpcAgentEvent, + RpcNotification, + SessionState, + SessionStats, + ThinkingConfig, + ThinkingLevel, + ThinkingLevelCycleResult, + ToolDescriptor, + ToolExecutionEndEvent, + ToolExecutionStartEvent, + ToolExecutionUpdateEvent, + ToolResultMessage, + TurnEndEvent, + TurnStartEvent, + UnknownNotification, + UserMessage, + assistant_text, + image_from_path, + message_text, + parse_notification, + parse_session_state, +) + +__all__ = [ + "AgentEventListener", + "AgentEndEvent", + "AgentMessage", + "AssistantMessage", + "AssistantMessageEvent", + "AutoCompactionEndEvent", + "AutoCompactionStartEvent", + "AutoRetryEndEvent", + "AutoRetryStartEvent", + "BashResult", + "BranchMessage", + "BranchResult", + "CancellationResult", + "CompactionResult", + "CompactionSummaryMessage", + "CustomMessage", + "DeveloperMessage", + "ExtensionError", + "ExtensionErrorListener", + "ExtensionUiRequest", + "FileMentionMessage", + "HookMessage", + "ImageContent", + "MessageEndEvent", + "MessageStartEvent", + "MessageUpdateEvent", + "ModelCost", + "ModelCycleResult", + "ModelInfo", + "PromptTurn", + "PythonExecutionMessage", + "ReadyEvent", + "RetryFallbackAppliedEvent", + "RetryFallbackSucceededEvent", + "RpcAgentEvent", + "RpcClient", + "RpcCommandError", + "RpcError", + "RpcNotification", + "RpcProcessExitError", + "RpcTimeoutError", + "SessionState", + "SessionStats", + "ThinkingConfig", + "ThinkingLevel", + "ThinkingLevelCycleResult", + "ToolDescriptor", + "ToolExecutionEndEvent", + "ToolExecutionStartEvent", + "ToolExecutionUpdateEvent", + "ToolResultMessage", + "TurnEndEvent", + "TurnStartEvent", + "UiRequestListener", + "UnknownNotification", + "UserMessage", + "assistant_text", + "image_from_path", + "message_text", + "parse_notification", + "parse_session_state", +] diff --git a/python/omp-rpc/src/omp_rpc/client.py b/python/omp-rpc/src/omp_rpc/client.py new file mode 100644 index 000000000..d17537d4c --- /dev/null +++ b/python/omp-rpc/src/omp_rpc/client.py @@ -0,0 +1,600 @@ +from __future__ import annotations + +import json +import os +import queue +import subprocess +import threading +import time +from dataclasses import dataclass +from pathlib import Path +from typing import Any, Callable, Mapping, Sequence, TypeVar, cast + +from .protocol import ( + AgentEndEvent, + AgentMessage, + AssistantMessage, + BashResult, + BranchMessage, + BranchResult, + CancellationResult, + CompactionResult, + ExtensionError, + ExtensionUiRequest, + ImageContent, + InterruptMode, + JsonObject, + JsonValue, + ModelCycleResult, + ModelInfo, + ReadyEvent, + RpcAgentEvent, + SessionState, + SessionStats, + SteeringMode, + StreamingBehavior, + ThinkingLevel, + ThinkingLevelCycleResult, + assistant_text, + parse_bash_result, + parse_branch_messages, + parse_branch_result, + parse_cancellation_result, + parse_compaction_result, + parse_model_cycle_result, + parse_model_info, + parse_notification, + parse_session_state, + parse_session_stats, + parse_thinking_level_cycle_result, +) + +AgentEventListener = Callable[[RpcAgentEvent], None] +UiRequestListener = Callable[[ExtensionUiRequest], None] +ExtensionErrorListener = Callable[[ExtensionError], None] +TListener = TypeVar("TListener") + + +class RpcError(RuntimeError): + """Base exception for the Python RPC client.""" + + +class RpcTimeoutError(RpcError): + """Raised when the server does not respond before a timeout.""" + + +class RpcProcessExitError(RpcError): + """Raised when the RPC process exits while a request is pending.""" + + +class RpcCommandError(RpcError): + """Raised when the RPC server returns `success: false`.""" + + def __init__(self, command: str, error: str): + super().__init__(f"{command}: {error}") + self.command = command + self.error = error + + +@dataclass(slots=True, frozen=True) +class PromptTurn: + events: tuple[RpcAgentEvent, ...] + messages: tuple[AgentMessage, ...] + assistant_message: AssistantMessage | None + assistant_text: str | None + + def require_assistant_text(self) -> str: + if self.assistant_text is None: + raise RpcError("Prompt completed without a text assistant message") + return self.assistant_text + + +class RpcClient: + def __init__( + self, + *, + command: Sequence[str] | None = None, + executable: str = "omp", + provider: str | None = None, + model: str | None = None, + session_dir: str | Path | None = None, + cwd: str | Path | None = None, + env: Mapping[str, str] | None = None, + extra_args: Sequence[str] = (), + startup_timeout: float = 30.0, + request_timeout: float = 30.0, + ) -> None: + self._command = tuple(command) if command is not None else None + self._executable = executable + self._provider = provider + self._model = model + self._session_dir = Path(session_dir) if session_dir is not None else None + self._cwd = Path(cwd) if cwd is not None else None + self._env = dict(env or {}) + self._extra_args = tuple(extra_args) + self._startup_timeout = startup_timeout + self._request_timeout = request_timeout + + self._process: subprocess.Popen[str] | None = None + self._stdout_thread: threading.Thread | None = None + self._stderr_thread: threading.Thread | None = None + self._ready = threading.Event() + self._write_lock = threading.Lock() + self._state_lock = threading.Lock() + self._event_condition = threading.Condition() + self._pending: dict[str, queue.Queue[JsonObject | BaseException]] = {} + self._request_id = 0 + self._events: list[RpcAgentEvent] = [] + self._ui_requests: queue.Queue[ExtensionUiRequest] = queue.Queue() + self._stderr_chunks: list[str] = [] + self._closed_error: BaseException | None = None + self._stopping = False + + self._event_listeners: list[AgentEventListener] = [] + self._ui_request_listeners: list[UiRequestListener] = [] + self._extension_error_listeners: list[ExtensionErrorListener] = [] + + def __enter__(self) -> RpcClient: + return self.start() + + def __exit__(self, _exc_type: object, _exc: object, _tb: object) -> None: + self.stop() + + @property + def stderr(self) -> str: + return "".join(self._stderr_chunks) + + @property + def command(self) -> tuple[str, ...]: + return self._build_command() + + def start(self) -> RpcClient: + if self._process is not None: + raise RpcError("RPC client is already started") + + self._ready.clear() + self._stopping = False + self._closed_error = None + self._events.clear() + self._ui_requests = queue.Queue() + self._stderr_chunks.clear() + + process = subprocess.Popen( + list(self._build_command()), + cwd=str(self._cwd) if self._cwd is not None else None, + env={**os.environ, **self._env}, + stdin=subprocess.PIPE, + stdout=subprocess.PIPE, + stderr=subprocess.PIPE, + text=True, + encoding="utf-8", + bufsize=1, + ) + 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.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}") + + return self + + def stop(self) -> None: + process = self._process + if process is None: + return + + self._stopping = True + + try: + if process.stdin is not None: + try: + process.stdin.close() + except OSError: + pass + + if process.poll() is None: + process.terminate() + try: + process.wait(timeout=1.0) + except subprocess.TimeoutExpired: + process.kill() + process.wait(timeout=1.0) + finally: + if process.stdout is not None: + try: + process.stdout.close() + except OSError: + pass + if process.stderr is not None: + try: + process.stderr.close() + except OSError: + pass + self._fail_pending(RpcProcessExitError("RPC process stopped")) + self._process = None + self._ready.set() + with self._event_condition: + self._event_condition.notify_all() + if self._stdout_thread is not None: + self._stdout_thread.join(timeout=1.0) + if self._stderr_thread is not None: + self._stderr_thread.join(timeout=1.0) + self._stdout_thread = None + self._stderr_thread = None + + def on_event(self, listener: AgentEventListener) -> Callable[[], None]: + self._event_listeners.append(listener) + return lambda: self._remove_listener(self._event_listeners, listener) + + def on_ui_request(self, listener: UiRequestListener) -> Callable[[], None]: + 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]: + self._extension_error_listeners.append(listener) + return lambda: self._remove_listener(self._extension_error_listeners, listener) + + def next_ui_request(self, timeout: float | None = None) -> ExtensionUiRequest: + 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 + + def send_ui_value(self, request_id: str, value: str) -> None: + 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}) + + def cancel_ui_request(self, request_id: str, *, timed_out: bool = False) -> None: + payload: JsonObject = {"type": "extension_ui_response", "id": request_id, "cancelled": True} + if timed_out: + payload["timedOut"] = True + self._send_notification(payload) + + def get_state(self) -> SessionState: + payload = self._request("get_state") + return parse_session_state(payload) + + def set_model(self, provider: str, model_id: str) -> ModelInfo: + payload = self._request("set_model", provider=provider, modelId=model_id) + model = parse_model_info(payload) + if model is None: + raise RpcError("set_model returned an empty payload") + return model + + def cycle_model(self) -> ModelCycleResult | None: + return parse_model_cycle_result(self._request("cycle_model")) + + def get_available_models(self) -> tuple[ModelInfo, ...]: + payload = self._request("get_available_models") + models = cast(list[JsonObject], payload.get("models") or []) + return tuple(filter(None, (parse_model_info(model) for model in models))) + + def set_thinking_level(self, level: ThinkingLevel) -> None: + self._request("set_thinking_level", level=level) + + def cycle_thinking_level(self) -> ThinkingLevelCycleResult | None: + return parse_thinking_level_cycle_result(self._request("cycle_thinking_level")) + + def set_steering_mode(self, mode: SteeringMode) -> None: + self._request("set_steering_mode", mode=mode) + + def set_follow_up_mode(self, mode: SteeringMode) -> None: + self._request("set_follow_up_mode", mode=mode) + + def set_interrupt_mode(self, mode: InterruptMode) -> None: + self._request("set_interrupt_mode", mode=mode) + + def compact(self, custom_instructions: str | None = None) -> CompactionResult: + payload = self._request("compact", customInstructions=custom_instructions) + return parse_compaction_result(payload) + + def set_auto_compaction(self, enabled: bool) -> None: + self._request("set_auto_compaction", enabled=enabled) + + def set_auto_retry(self, enabled: bool) -> None: + self._request("set_auto_retry", enabled=enabled) + + def abort_retry(self) -> None: + self._request("abort_retry") + + def bash(self, command: str) -> BashResult: + payload = self._request("bash", command=command) + return parse_bash_result(payload) + + def abort_bash(self) -> None: + self._request("abort_bash") + + def get_session_stats(self) -> SessionStats: + payload = self._request("get_session_stats") + 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) + 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)) + + def switch_session(self, session_path: str | Path) -> CancellationResult: + 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)) + + def get_branch_messages(self) -> tuple[BranchMessage, ...]: + return parse_branch_messages(self._request("get_branch_messages")) + + def get_last_assistant_text(self) -> str | None: + payload = self._request("get_last_assistant_text") + value = payload.get("text") + return str(value) if isinstance(value, str) else None + + def set_session_name(self, name: str) -> None: + self._request("set_session_name", name=name) + + def get_messages(self) -> tuple[AgentMessage, ...]: + payload = self._request("get_messages") + return tuple(cast(list[AgentMessage], payload.get("messages") or [])) + + def prompt( + self, + message: str, + *, + images: Sequence[ImageContent] | None = None, + streaming_behavior: StreamingBehavior | None = None, + ) -> None: + self._request( + "prompt", + message=message, + images=list(images) if images is not None else None, + streamingBehavior=streaming_behavior, + ) + + 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 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 prompt_and_wait( + self, + message: str, + *, + images: Sequence[ImageContent] | None = None, + streaming_behavior: StreamingBehavior | None = None, + timeout: float | None = None, + ) -> PromptTurn: + start_index = self._current_event_index() + self.prompt(message, images=images, streaming_behavior=streaming_behavior) + events = self._wait_for_agent_end(start_index, timeout=timeout) + return self._build_prompt_turn(events) + + def wait_for_idle(self, timeout: float | None = None) -> None: + start_index = self._current_event_index() + self._wait_for_agent_end(start_index, timeout=timeout) + + def collect_events(self, timeout: float | None = None) -> tuple[RpcAgentEvent, ...]: + start_index = self._current_event_index() + return self._wait_for_agent_end(start_index, timeout=timeout) + + def request_raw(self, command_type: str, **payload: JsonValue) -> JsonObject: + return self._request(command_type, **payload) + + def _current_event_index(self) -> int: + with self._event_condition: + return len(self._events) + + def _build_prompt_turn(self, events: tuple[RpcAgentEvent, ...]) -> PromptTurn: + final_messages: tuple[AgentMessage, ...] = () + for event in reversed(events): + if isinstance(event, AgentEndEvent): + final_messages = event.messages + break + + assistant_message: AssistantMessage | None = None + for message in reversed(final_messages): + if message.get("role") == "assistant": + assistant_message = cast(AssistantMessage, message) + break + + if assistant_message is None: + for event in reversed(events): + if hasattr(event, "message"): + message = cast(AgentMessage | None, getattr(event, "message", None)) + if isinstance(message, dict) and message.get("role") == "assistant": + assistant_message = cast(AssistantMessage, message) + break + + return PromptTurn( + events=events, + messages=final_messages, + assistant_message=assistant_message, + assistant_text=assistant_text(assistant_message) if assistant_message is not None else None, + ) + + def _wait_for_agent_end(self, start_index: int, timeout: float | None = None) -> tuple[RpcAgentEvent, ...]: + deadline = time.monotonic() + (timeout if timeout is not None else 60.0) + with self._event_condition: + while True: + if self._closed_error is not None: + raise RpcProcessExitError(str(self._closed_error)) + + events = tuple(self._events[start_index:]) + if any(isinstance(event, AgentEndEvent) for event in events): + return events + + remaining = deadline - time.monotonic() + if remaining <= 0: + 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: + process = self._require_process() + request_id = self._next_request_id() + envelope: JsonObject = {"id": request_id, "type": command_type} + for key, value in payload.items(): + if value is not None: + envelope[key] = value + + response_queue: queue.Queue[JsonObject | BaseException] = queue.Queue(maxsize=1) + with self._state_lock: + self._pending[request_id] = response_queue + + self._write_json(process, envelope) + + try: + response = response_queue.get(timeout=self._request_timeout) + 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 + + 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", ""))) + + data = response.get("data") + return dict(cast(JsonObject, data or {})) + + def _send_notification(self, payload: JsonObject) -> None: + process = self._require_process() + self._write_json(process, payload) + + def _build_command(self) -> tuple[str, ...]: + if self._command is not None: + return self._command + + command: list[str] = [self._executable, "--mode", "rpc"] + if self._provider: + command.extend(["--provider", self._provider]) + if self._model: + command.extend(["--model", self._model]) + if self._session_dir is not None: + command.extend(["--session-dir", str(self._session_dir)]) + command.extend(self._extra_args) + return tuple(command) + + def _next_request_id(self) -> str: + with self._state_lock: + self._request_id += 1 + return f"req_{self._request_id}" + + def _require_process(self) -> subprocess.Popen[str]: + if self._process is None: + raise RpcError("RPC client is not started") + return self._process + + def _write_json(self, process: subprocess.Popen[str], payload: JsonObject) -> None: + if process.stdin is None: + raise RpcProcessExitError("RPC process stdin is unavailable") + with self._write_lock: + try: + process.stdin.write(json.dumps(payload)) + process.stdin.write("\n") + process.stdin.flush() + except (BrokenPipeError, OSError) as exc: + raise RpcProcessExitError(f"Failed to write RPC command: {exc}") from exc + + def _read_stdout_loop(self) -> None: + process = self._process + if process is None or process.stdout is None: + return + + try: + for line in process.stdout: + stripped = line.strip() + if not stripped: + continue + + payload = cast(JsonObject, json.loads(stripped)) + if payload.get("type") == "response": + request_id = payload.get("id") + if isinstance(request_id, str): + with self._state_lock: + pending = self._pending.pop(request_id, None) + if pending is not None: + pending.put(payload) + continue + + notification = parse_notification(payload) + if isinstance(notification, ReadyEvent): + self._ready.set() + continue + + if isinstance(notification, ExtensionUiRequest): + self._ui_requests.put(notification) + for listener in list(self._ui_request_listeners): + listener(notification) + continue + + if isinstance(notification, ExtensionError): + for listener in list(self._extension_error_listeners): + listener(notification) + continue + + if getattr(notification, "type", None) != "unknown": + with self._event_condition: + self._events.append(cast(RpcAgentEvent, notification)) + self._event_condition.notify_all() + for listener in list(self._event_listeners): + listener(cast(RpcAgentEvent, notification)) + except json.JSONDecodeError as exc: + self._mark_closed(RpcError(f"Failed to decode RPC output: {exc}")) + except Exception as exc: + self._mark_closed(exc) + else: + if not self._stopping: + exit_code = process.poll() + if exit_code is None: + try: + exit_code = process.wait(timeout=1.0) + except subprocess.TimeoutExpired: + 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}")) + + def _read_stderr_loop(self) -> None: + process = self._process + if process is None or process.stderr is None: + return + for chunk in process.stderr: + self._stderr_chunks.append(chunk) + + def _mark_closed(self, error: BaseException) -> None: + if self._closed_error is not None: + return + self._closed_error = error + self._ready.set() + self._fail_pending(error) + with self._event_condition: + self._event_condition.notify_all() + + def _fail_pending(self, error: BaseException) -> None: + with self._state_lock: + pending = list(self._pending.values()) + self._pending.clear() + for response_queue in pending: + response_queue.put(error) + + @staticmethod + def _remove_listener(listeners: list[TListener], listener: TListener) -> None: + try: + listeners.remove(listener) + except ValueError: + pass diff --git a/python/omp-rpc/src/omp_rpc/protocol.py b/python/omp-rpc/src/omp_rpc/protocol.py new file mode 100644 index 000000000..4fce000a2 --- /dev/null +++ b/python/omp-rpc/src/omp_rpc/protocol.py @@ -0,0 +1,963 @@ +from __future__ import annotations + +import base64 +import mimetypes +from dataclasses import dataclass +from pathlib import Path +from typing import Any, Literal, NotRequired, TypedDict, TypeAlias, cast + +JsonPrimitive: TypeAlias = str | int | float | bool | None +JsonValue: TypeAlias = JsonPrimitive | list["JsonValue"] | dict[str, "JsonValue"] +JsonObject: TypeAlias = dict[str, JsonValue] + +Attribution: TypeAlias = Literal["user", "agent"] +ThinkingLevel: TypeAlias = Literal["off", "minimal", "low", "medium", "high", "xhigh"] +StreamingBehavior: TypeAlias = Literal["steer", "followUp"] +SteeringMode: TypeAlias = Literal["all", "one-at-a-time"] +InterruptMode: TypeAlias = Literal["immediate", "wait"] +StopReason: TypeAlias = Literal["stop", "length", "toolUse", "error", "aborted"] +NotifyType: TypeAlias = Literal["info", "warning", "error"] +WidgetPlacement: TypeAlias = Literal["aboveEditor", "belowEditor"] + + +class TextContent(TypedDict, total=False): + type: Literal["text"] + text: str + textSignature: NotRequired[str] + + +class ThinkingContent(TypedDict, total=False): + type: Literal["thinking"] + thinking: str + thinkingSignature: NotRequired[str] + + +class RedactedThinkingContent(TypedDict, total=False): + type: Literal["redactedThinking"] + data: str + + +class ImageContent(TypedDict, total=False): + type: Literal["image"] + data: str + mimeType: str + + +class ToolCall(TypedDict, total=False): + type: Literal["toolCall"] + id: str + name: str + arguments: dict[str, Any] + thoughtSignature: NotRequired[str] + intent: NotRequired[str] + + +class UsageCost(TypedDict): + input: float + output: float + cacheRead: float + cacheWrite: float + total: float + + +class Usage(TypedDict, total=False): + input: int + output: int + cacheRead: int + cacheWrite: int + totalTokens: int + premiumRequests: NotRequired[int] + cost: UsageCost + + +class UserMessage(TypedDict, total=False): + role: Literal["user"] + content: str | list[TextContent | ImageContent] + synthetic: NotRequired[bool] + attribution: NotRequired[Attribution] + providerPayload: NotRequired[JsonObject] + timestamp: int + + +class DeveloperMessage(TypedDict, total=False): + role: Literal["developer"] + content: str | list[TextContent | ImageContent] + attribution: NotRequired[Attribution] + providerPayload: NotRequired[JsonObject] + timestamp: int + + +class AssistantMessage(TypedDict, total=False): + role: Literal["assistant"] + content: list[TextContent | ThinkingContent | RedactedThinkingContent | ToolCall] + api: str + provider: str + model: str + responseId: NotRequired[str] + usage: Usage + stopReason: StopReason + errorMessage: NotRequired[str] + providerPayload: NotRequired[JsonObject] + timestamp: int + duration: NotRequired[int] + ttft: NotRequired[int] + + +class ToolResultMessage(TypedDict, total=False): + role: Literal["toolResult"] + toolCallId: str + toolName: str + content: list[TextContent | ImageContent] + details: NotRequired[JsonValue] + isError: bool + attribution: NotRequired[Attribution] + prunedAt: NotRequired[int] + timestamp: int + + +class BashExecutionMessage(TypedDict, total=False): + role: Literal["bashExecution"] + command: str + output: str + exitCode: int | None + cancelled: bool + truncated: bool + meta: NotRequired[JsonObject] + timestamp: int + excludeFromContext: NotRequired[bool] + + +class PythonExecutionMessage(TypedDict, total=False): + role: Literal["pythonExecution"] + code: str + output: str + exitCode: int | None + cancelled: bool + truncated: bool + meta: NotRequired[JsonObject] + timestamp: int + excludeFromContext: NotRequired[bool] + + +class CustomMessage(TypedDict, total=False): + role: Literal["custom"] + customType: str + content: str | list[TextContent | ImageContent] + display: bool + details: NotRequired[JsonValue] + attribution: NotRequired[Attribution] + timestamp: int + + +class HookMessage(TypedDict, total=False): + role: Literal["hookMessage"] + customType: str + content: str | list[TextContent | ImageContent] + display: bool + details: NotRequired[JsonValue] + attribution: NotRequired[Attribution] + timestamp: int + + +class BranchSummaryMessage(TypedDict, total=False): + role: Literal["branchSummary"] + summary: str + fromId: str + timestamp: int + + +class CompactionSummaryMessage(TypedDict, total=False): + role: Literal["compactionSummary"] + summary: str + shortSummary: NotRequired[str] + tokensBefore: int + providerPayload: NotRequired[JsonObject] + timestamp: int + + +class FileMentionItem(TypedDict, total=False): + path: str + content: str + lineCount: NotRequired[int] + byteSize: NotRequired[int] + skippedReason: NotRequired[Literal["tooLarge"]] + image: NotRequired[ImageContent] + + +class FileMentionMessage(TypedDict, total=False): + role: Literal["fileMention"] + files: list[FileMentionItem] + timestamp: int + + +AgentMessage: TypeAlias = ( + UserMessage + | DeveloperMessage + | AssistantMessage + | ToolResultMessage + | BashExecutionMessage + | PythonExecutionMessage + | CustomMessage + | HookMessage + | BranchSummaryMessage + | CompactionSummaryMessage + | FileMentionMessage +) + + +class AssistantMessageStartEvent(TypedDict): + type: Literal["start"] + partial: AssistantMessage + + +class AssistantTextStartEvent(TypedDict): + type: Literal["text_start"] + contentIndex: int + partial: AssistantMessage + + +class AssistantTextDeltaEvent(TypedDict): + type: Literal["text_delta"] + contentIndex: int + delta: str + partial: AssistantMessage + + +class AssistantTextEndEvent(TypedDict): + type: Literal["text_end"] + contentIndex: int + content: str + partial: AssistantMessage + + +class AssistantThinkingStartEvent(TypedDict): + type: Literal["thinking_start"] + contentIndex: int + partial: AssistantMessage + + +class AssistantThinkingDeltaEvent(TypedDict): + type: Literal["thinking_delta"] + contentIndex: int + delta: str + partial: AssistantMessage + + +class AssistantThinkingEndEvent(TypedDict): + type: Literal["thinking_end"] + contentIndex: int + content: str + partial: AssistantMessage + + +class AssistantToolCallStartEvent(TypedDict): + type: Literal["toolcall_start"] + contentIndex: int + partial: AssistantMessage + + +class AssistantToolCallDeltaEvent(TypedDict): + type: Literal["toolcall_delta"] + contentIndex: int + delta: str + partial: AssistantMessage + + +class AssistantToolCallEndEvent(TypedDict): + type: Literal["toolcall_end"] + contentIndex: int + toolCall: ToolCall + partial: AssistantMessage + + +class AssistantDoneEvent(TypedDict): + type: Literal["done"] + reason: Literal["stop", "length", "toolUse"] + message: AssistantMessage + + +class AssistantErrorEvent(TypedDict): + type: Literal["error"] + reason: Literal["aborted", "error"] + error: AssistantMessage + + +AssistantMessageEvent: TypeAlias = ( + AssistantMessageStartEvent + | AssistantTextStartEvent + | AssistantTextDeltaEvent + | AssistantTextEndEvent + | AssistantThinkingStartEvent + | AssistantThinkingDeltaEvent + | AssistantThinkingEndEvent + | AssistantToolCallStartEvent + | AssistantToolCallDeltaEvent + | AssistantToolCallEndEvent + | AssistantDoneEvent + | AssistantErrorEvent +) + + +@dataclass(slots=True, frozen=True) +class ModelCost: + input: float + output: float + cache_read: float + cache_write: float + + +@dataclass(slots=True, frozen=True) +class ThinkingConfig: + min_level: ThinkingLevel + max_level: ThinkingLevel + mode: str + + +@dataclass(slots=True, frozen=True) +class ModelInfo: + id: str + name: str + api: str + provider: str + base_url: str + reasoning: bool + input_modalities: tuple[str, ...] + cost: ModelCost + context_window: int + max_tokens: int + headers: dict[str, str] | None = None + premium_multiplier: float | None = None + prefer_websockets: bool | None = None + context_promotion_target: str | None = None + priority: int | None = None + thinking: ThinkingConfig | None = None + compat: JsonObject | None = None + + +@dataclass(slots=True, frozen=True) +class ToolDescriptor: + name: str + description: str + parameters: JsonValue + + +@dataclass(slots=True, frozen=True) +class SessionState: + model: ModelInfo | None + thinking_level: ThinkingLevel | None + is_streaming: bool + is_compacting: bool + steering_mode: SteeringMode + follow_up_mode: SteeringMode + interrupt_mode: InterruptMode + session_file: str | None + session_id: str + session_name: str | None + auto_compaction_enabled: bool + message_count: int + queued_message_count: int + system_prompt: str | None = None + dump_tools: tuple[ToolDescriptor, ...] = () + + +@dataclass(slots=True, frozen=True) +class BashResult: + output: str + exit_code: int | None + cancelled: bool + truncated: bool + total_lines: int + total_bytes: int + output_lines: int + output_bytes: int + artifact_id: str | None = None + + +@dataclass(slots=True, frozen=True) +class CompactionResult: + summary: str + first_kept_entry_id: str + tokens_before: int + short_summary: str | None = None + details: JsonValue | None = None + preserve_data: JsonObject | None = None + + +@dataclass(slots=True, frozen=True) +class ModelCycleResult: + model: ModelInfo + thinking_level: ThinkingLevel | None + is_scoped: bool + + +@dataclass(slots=True, frozen=True) +class ThinkingLevelCycleResult: + level: ThinkingLevel + + +@dataclass(slots=True, frozen=True) +class CancellationResult: + cancelled: bool + + +@dataclass(slots=True, frozen=True) +class BranchMessage: + entry_id: str + text: str + + +@dataclass(slots=True, frozen=True) +class BranchResult: + text: str + cancelled: bool + + +@dataclass(slots=True, frozen=True) +class TokenUsage: + input: int + output: int + cache_read: int + cache_write: int + total: int + + +@dataclass(slots=True, frozen=True) +class SessionStats: + session_file: str | None + session_id: str + user_messages: int + assistant_messages: int + tool_calls: int + tool_results: int + total_messages: int + tokens: TokenUsage + premium_requests: int + cost: float + + +@dataclass(slots=True, frozen=True) +class ReadyEvent: + type: Literal["ready"] = "ready" + + +@dataclass(slots=True, frozen=True) +class ExtensionUiRequest: + id: str + method: str + title: str | None = None + options: tuple[str, ...] | None = None + message: str | None = None + placeholder: str | None = None + prefill: str | None = None + timeout: int | None = None + prompt_style: bool | None = None + target_id: str | None = None + notify_type: NotifyType | None = None + status_key: str | None = None + status_text: str | None = None + widget_key: str | None = None + widget_lines: tuple[str, ...] | None = None + widget_placement: WidgetPlacement | None = None + text: str | None = None + + +@dataclass(slots=True, frozen=True) +class ExtensionError: + extension_path: str + event: str + error: str + type: Literal["extension_error"] = "extension_error" + + +@dataclass(slots=True, frozen=True) +class AgentStartEvent: + type: Literal["agent_start"] = "agent_start" + + +@dataclass(slots=True, frozen=True) +class AgentEndEvent: + messages: tuple[AgentMessage, ...] + type: Literal["agent_end"] = "agent_end" + + +@dataclass(slots=True, frozen=True) +class TurnStartEvent: + type: Literal["turn_start"] = "turn_start" + + +@dataclass(slots=True, frozen=True) +class TurnEndEvent: + message: AgentMessage + tool_results: tuple[ToolResultMessage, ...] + type: Literal["turn_end"] = "turn_end" + + +@dataclass(slots=True, frozen=True) +class MessageStartEvent: + message: AgentMessage + type: Literal["message_start"] = "message_start" + + +@dataclass(slots=True, frozen=True) +class MessageUpdateEvent: + message: AgentMessage + assistant_message_event: AssistantMessageEvent + type: Literal["message_update"] = "message_update" + + +@dataclass(slots=True, frozen=True) +class MessageEndEvent: + message: AgentMessage + type: Literal["message_end"] = "message_end" + + +@dataclass(slots=True, frozen=True) +class ToolExecutionStartEvent: + tool_call_id: str + tool_name: str + args: JsonValue + intent: str | None = None + type: Literal["tool_execution_start"] = "tool_execution_start" + + +@dataclass(slots=True, frozen=True) +class ToolExecutionUpdateEvent: + tool_call_id: str + tool_name: str + args: JsonValue + partial_result: JsonValue + type: Literal["tool_execution_update"] = "tool_execution_update" + + +@dataclass(slots=True, frozen=True) +class ToolExecutionEndEvent: + tool_call_id: str + tool_name: str + result: JsonValue + is_error: bool | None = None + type: Literal["tool_execution_end"] = "tool_execution_end" + + +@dataclass(slots=True, frozen=True) +class AutoCompactionStartEvent: + reason: Literal["threshold", "overflow", "idle"] + action: Literal["context-full", "handoff"] + type: Literal["auto_compaction_start"] = "auto_compaction_start" + + +@dataclass(slots=True, frozen=True) +class AutoCompactionEndEvent: + action: Literal["context-full", "handoff"] + result: CompactionResult | None + aborted: bool + will_retry: bool + error_message: str | None = None + skipped: bool | None = None + type: Literal["auto_compaction_end"] = "auto_compaction_end" + + +@dataclass(slots=True, frozen=True) +class AutoRetryStartEvent: + attempt: int + max_attempts: int + delay_ms: int + error_message: str + type: Literal["auto_retry_start"] = "auto_retry_start" + + +@dataclass(slots=True, frozen=True) +class AutoRetryEndEvent: + success: bool + attempt: int + final_error: str | None = None + type: Literal["auto_retry_end"] = "auto_retry_end" + + +@dataclass(slots=True, frozen=True) +class RetryFallbackAppliedEvent: + from_model: str + to_model: str + role: str + type: Literal["retry_fallback_applied"] = "retry_fallback_applied" + + +@dataclass(slots=True, frozen=True) +class RetryFallbackSucceededEvent: + model: str + role: str + type: Literal["retry_fallback_succeeded"] = "retry_fallback_succeeded" + + +@dataclass(slots=True, frozen=True) +class TtsrTriggeredEvent: + rules: tuple[JsonObject, ...] + type: Literal["ttsr_triggered"] = "ttsr_triggered" + + +@dataclass(slots=True, frozen=True) +class TodoReminderEvent: + todos: tuple[JsonObject, ...] + attempt: int + max_attempts: int + type: Literal["todo_reminder"] = "todo_reminder" + + +@dataclass(slots=True, frozen=True) +class TodoAutoClearEvent: + type: Literal["todo_auto_clear"] = "todo_auto_clear" + + +@dataclass(slots=True, frozen=True) +class UnknownNotification: + payload: JsonObject + type: Literal["unknown"] = "unknown" + + +RpcAgentEvent: TypeAlias = ( + AgentStartEvent + | AgentEndEvent + | TurnStartEvent + | TurnEndEvent + | MessageStartEvent + | MessageUpdateEvent + | MessageEndEvent + | ToolExecutionStartEvent + | ToolExecutionUpdateEvent + | ToolExecutionEndEvent + | AutoCompactionStartEvent + | AutoCompactionEndEvent + | AutoRetryStartEvent + | AutoRetryEndEvent + | RetryFallbackAppliedEvent + | RetryFallbackSucceededEvent + | TtsrTriggeredEvent + | TodoReminderEvent + | TodoAutoClearEvent +) + +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" + return { + "type": "image", + "mimeType": resolved_mime_type, + "data": base64.b64encode(file_path.read_bytes()).decode("ascii"), + } + + +def message_text(message: AgentMessage) -> str | None: + role = message.get("role") + if role not in {"user", "developer", "assistant", "toolResult", "custom", "hookMessage"}: + return None + + content = message.get("content") + if isinstance(content, str): + return content + if not isinstance(content, list): + return None + + fragments: list[str] = [] + for block in content: + if not isinstance(block, dict): + continue + block_type = block.get("type") + if block_type == "text" and isinstance(block.get("text"), str): + fragments.append(cast(str, block["text"])) + elif block_type == "thinking" and isinstance(block.get("thinking"), str): + fragments.append(cast(str, block["thinking"])) + return "".join(fragments) or None + + +def assistant_text(message: AgentMessage) -> str | None: + if message.get("role") != "assistant": + return None + return message_text(message) + + +def parse_model_info(payload: JsonObject | None) -> ModelInfo | None: + if payload is None: + return None + cost_payload = cast(dict[str, Any], payload.get("cost") or {}) + thinking_payload = payload.get("thinking") + headers_payload = payload.get("headers") + compat_payload = payload.get("compat") + return ModelInfo( + id=str(payload["id"]), + name=str(payload["name"]), + api=str(payload["api"]), + provider=str(payload["provider"]), + base_url=str(payload["baseUrl"]), + reasoning=bool(payload.get("reasoning", False)), + input_modalities=tuple(str(item) for item in cast(list[Any], payload.get("input") or [])), + cost=ModelCost( + input=float(cost_payload.get("input", 0.0)), + output=float(cost_payload.get("output", 0.0)), + cache_read=float(cost_payload.get("cacheRead", 0.0)), + cache_write=float(cost_payload.get("cacheWrite", 0.0)), + ), + context_window=int(payload.get("contextWindow", 0)), + max_tokens=int(payload.get("maxTokens", 0)), + headers=dict(cast(dict[str, str], headers_payload)) if isinstance(headers_payload, dict) else None, + 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 + ), + priority=int(payload["priority"]) if "priority" in payload else None, + thinking=( + ThinkingConfig( + min_level=cast(ThinkingLevel, thinking_payload["minLevel"]), + max_level=cast(ThinkingLevel, thinking_payload["maxLevel"]), + mode=str(thinking_payload["mode"]), + ) + if isinstance(thinking_payload, dict) + else None + ), + compat=dict(cast(dict[str, JsonValue], compat_payload)) if isinstance(compat_payload, dict) else None, + ) + + +def parse_tool_descriptor(payload: JsonObject) -> ToolDescriptor: + return ToolDescriptor( + name=str(payload["name"]), + description=str(payload["description"]), + parameters=cast(JsonValue, payload.get("parameters")), + ) + + +def parse_session_state(payload: JsonObject) -> SessionState: + dump_tools = tuple( + parse_tool_descriptor(cast(JsonObject, item)) 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, payload.get("thinkingLevel")), + is_streaming=bool(payload.get("isStreaming", False)), + is_compacting=bool(payload.get("isCompacting", False)), + steering_mode=cast(SteeringMode, payload.get("steeringMode", "one-at-a-time")), + follow_up_mode=cast(SteeringMode, payload.get("followUpMode", "one-at-a-time")), + interrupt_mode=cast(InterruptMode, payload.get("interruptMode", "immediate")), + session_file=str(payload["sessionFile"]) if payload.get("sessionFile") is not None else None, + session_id=str(payload["sessionId"]), + session_name=str(payload["sessionName"]) if payload.get("sessionName") is not None else None, + auto_compaction_enabled=bool(payload.get("autoCompactionEnabled", False)), + message_count=int(payload.get("messageCount", 0)), + queued_message_count=int(payload.get("queuedMessageCount", 0)), + system_prompt=str(payload["systemPrompt"]) if payload.get("systemPrompt") is not None else None, + dump_tools=dump_tools, + ) + + +def parse_bash_result(payload: JsonObject) -> BashResult: + return BashResult( + output=str(payload.get("output", "")), + exit_code=int(payload["exitCode"]) if payload.get("exitCode") is not None else None, + cancelled=bool(payload.get("cancelled", False)), + truncated=bool(payload.get("truncated", False)), + total_lines=int(payload.get("totalLines", 0)), + total_bytes=int(payload.get("totalBytes", 0)), + output_lines=int(payload.get("outputLines", 0)), + output_bytes=int(payload.get("outputBytes", 0)), + artifact_id=str(payload["artifactId"]) if payload.get("artifactId") is not None else None, + ) + + +def parse_compaction_result(payload: JsonObject) -> CompactionResult: + return CompactionResult( + summary=str(payload.get("summary", "")), + short_summary=str(payload["shortSummary"]) if payload.get("shortSummary") is not None else None, + first_kept_entry_id=str(payload.get("firstKeptEntryId", "")), + tokens_before=int(payload.get("tokensBefore", 0)), + details=cast(JsonValue | None, payload.get("details")), + preserve_data=cast(JsonObject | None, payload.get("preserveData")), + ) + + +def parse_model_cycle_result(payload: JsonObject | None) -> ModelCycleResult | None: + if payload is None: + return None + model = parse_model_info(cast(JsonObject, payload.get("model"))) + if model is None: + raise ValueError("cycle_model response did not include a model") + return ModelCycleResult( + model=model, + thinking_level=cast(ThinkingLevel | None, payload.get("thinkingLevel")), + is_scoped=bool(payload.get("isScoped", False)), + ) + + +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"])) + + +def parse_cancellation_result(payload: JsonObject | None) -> CancellationResult: + return CancellationResult(cancelled=bool((payload or {}).get("cancelled", False))) + + +def parse_branch_result(payload: JsonObject | None) -> BranchResult: + payload = payload or {} + return BranchResult(text=str(payload.get("text", "")), cancelled=bool(payload.get("cancelled", False))) + + +def parse_branch_messages(payload: JsonObject | None) -> tuple[BranchMessage, ...]: + messages = cast(list[Any], (payload or {}).get("messages") or []) + return tuple( + BranchMessage(entry_id=str(item.get("entryId", "")), text=str(item.get("text", ""))) for item in messages + ) + + +def parse_session_stats(payload: JsonObject) -> SessionStats: + tokens_payload = cast(dict[str, Any], payload.get("tokens") or {}) + return SessionStats( + session_file=str(payload["sessionFile"]) if payload.get("sessionFile") is not None else None, + session_id=str(payload.get("sessionId", "")), + user_messages=int(payload.get("userMessages", 0)), + assistant_messages=int(payload.get("assistantMessages", 0)), + tool_calls=int(payload.get("toolCalls", 0)), + tool_results=int(payload.get("toolResults", 0)), + total_messages=int(payload.get("totalMessages", 0)), + tokens=TokenUsage( + input=int(tokens_payload.get("input", 0)), + output=int(tokens_payload.get("output", 0)), + cache_read=int(tokens_payload.get("cacheRead", 0)), + cache_write=int(tokens_payload.get("cacheWrite", 0)), + total=int(tokens_payload.get("total", 0)), + ), + premium_requests=int(payload.get("premiumRequests", 0)), + cost=float(payload.get("cost", 0.0)), + ) + + +def parse_extension_ui_request(payload: JsonObject) -> ExtensionUiRequest: + return ExtensionUiRequest( + id=str(payload["id"]), + method=str(payload["method"]), + title=str(payload["title"]) if payload.get("title") is not None else None, + options=tuple(str(item) for item in cast(list[Any], payload.get("options") or [])) or None, + message=str(payload["message"]) if payload.get("message") is not None else None, + placeholder=str(payload["placeholder"]) if payload.get("placeholder") is not None else None, + prefill=str(payload["prefill"]) if payload.get("prefill") is not None else None, + timeout=int(payload["timeout"]) if payload.get("timeout") is not None else None, + prompt_style=bool(payload["promptStyle"]) if "promptStyle" in payload else None, + target_id=str(payload["targetId"]) if payload.get("targetId") is not None else None, + notify_type=cast(NotifyType | None, payload.get("notifyType")), + status_key=str(payload["statusKey"]) if payload.get("statusKey") is not None else None, + status_text=str(payload["statusText"]) if payload.get("statusText") is not None else None, + widget_key=str(payload["widgetKey"]) if payload.get("widgetKey") is not None else None, + widget_lines=tuple(str(item) for item in cast(list[Any], payload.get("widgetLines") or [])) or None, + widget_placement=cast(WidgetPlacement | None, payload.get("widgetPlacement")), + text=str(payload["text"]) if payload.get("text") is not None else None, + ) + + +def parse_extension_error(payload: JsonObject) -> ExtensionError: + return ExtensionError( + extension_path=str(payload.get("extensionPath", "")), + event=str(payload.get("event", "")), + error=str(payload.get("error", "")), + ) + + +def parse_notification(payload: JsonObject) -> RpcNotification: + event_type = payload.get("type") + if event_type == "ready": + return ReadyEvent() + if event_type == "extension_ui_request": + return parse_extension_ui_request(payload) + if event_type == "extension_error": + return parse_extension_error(payload) + if event_type == "agent_start": + return AgentStartEvent() + if event_type == "agent_end": + return AgentEndEvent(messages=tuple(cast(list[AgentMessage], payload.get("messages") or []))) + if event_type == "turn_start": + return TurnStartEvent() + if event_type == "turn_end": + return TurnEndEvent( + message=cast(AgentMessage, payload["message"]), + tool_results=tuple(cast(list[ToolResultMessage], payload.get("toolResults") or [])), + ) + if event_type == "message_start": + return MessageStartEvent(message=cast(AgentMessage, payload["message"])) + if event_type == "message_update": + return MessageUpdateEvent( + message=cast(AgentMessage, payload["message"]), + assistant_message_event=cast(AssistantMessageEvent, payload["assistantMessageEvent"]), + ) + if event_type == "message_end": + return MessageEndEvent(message=cast(AgentMessage, payload["message"])) + if event_type == "tool_execution_start": + return ToolExecutionStartEvent( + tool_call_id=str(payload.get("toolCallId", "")), + tool_name=str(payload.get("toolName", "")), + args=cast(JsonValue, payload.get("args")), + intent=str(payload["intent"]) if payload.get("intent") is not None else None, + ) + if event_type == "tool_execution_update": + return ToolExecutionUpdateEvent( + tool_call_id=str(payload.get("toolCallId", "")), + tool_name=str(payload.get("toolName", "")), + args=cast(JsonValue, payload.get("args")), + partial_result=cast(JsonValue, payload.get("partialResult")), + ) + if event_type == "tool_execution_end": + return ToolExecutionEndEvent( + tool_call_id=str(payload.get("toolCallId", "")), + tool_name=str(payload.get("toolName", "")), + result=cast(JsonValue, payload.get("result")), + is_error=bool(payload["isError"]) if "isError" in payload else None, + ) + if event_type == "auto_compaction_start": + return AutoCompactionStartEvent( + reason=cast(Literal["threshold", "overflow", "idle"], payload.get("reason", "threshold")), + action=cast(Literal["context-full", "handoff"], payload.get("action", "context-full")), + ) + if event_type == "auto_compaction_end": + result_payload = payload.get("result") + return AutoCompactionEndEvent( + action=cast(Literal["context-full", "handoff"], payload.get("action", "context-full")), + result=( + parse_compaction_result(cast(JsonObject, result_payload)) if isinstance(result_payload, dict) else None + ), + aborted=bool(payload.get("aborted", False)), + will_retry=bool(payload.get("willRetry", False)), + error_message=str(payload["errorMessage"]) if payload.get("errorMessage") is not None else None, + skipped=bool(payload["skipped"]) if "skipped" in payload else None, + ) + if event_type == "auto_retry_start": + return AutoRetryStartEvent( + attempt=int(payload.get("attempt", 0)), + max_attempts=int(payload.get("maxAttempts", 0)), + delay_ms=int(payload.get("delayMs", 0)), + error_message=str(payload.get("errorMessage", "")), + ) + if event_type == "auto_retry_end": + return AutoRetryEndEvent( + success=bool(payload.get("success", False)), + attempt=int(payload.get("attempt", 0)), + final_error=str(payload["finalError"]) if payload.get("finalError") is not None else None, + ) + if event_type == "retry_fallback_applied": + return RetryFallbackAppliedEvent( + from_model=str(payload.get("from", "")), + to_model=str(payload.get("to", "")), + role=str(payload.get("role", "")), + ) + if event_type == "retry_fallback_succeeded": + return RetryFallbackSucceededEvent(model=str(payload.get("model", "")), role=str(payload.get("role", ""))) + if event_type == "ttsr_triggered": + return TtsrTriggeredEvent(rules=tuple(cast(list[JsonObject], payload.get("rules") or []))) + if event_type == "todo_reminder": + return TodoReminderEvent( + todos=tuple(cast(list[JsonObject], payload.get("todos") or [])), + attempt=int(payload.get("attempt", 0)), + max_attempts=int(payload.get("maxAttempts", 0)), + ) + if event_type == "todo_auto_clear": + return TodoAutoClearEvent() + return UnknownNotification(payload=dict(payload)) diff --git a/python/omp-rpc/src/omp_rpc/py.typed b/python/omp-rpc/src/omp_rpc/py.typed new file mode 100644 index 000000000..8b1378917 --- /dev/null +++ b/python/omp-rpc/src/omp_rpc/py.typed @@ -0,0 +1 @@ + diff --git a/python/omp-rpc/tests/__init__.py b/python/omp-rpc/tests/__init__.py new file mode 100644 index 000000000..5c8135e92 --- /dev/null +++ b/python/omp-rpc/tests/__init__.py @@ -0,0 +1,6 @@ +from __future__ import annotations + +import sys +from pathlib import Path + +sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src")) diff --git a/python/omp-rpc/tests/test_client.py b/python/omp-rpc/tests/test_client.py new file mode 100644 index 000000000..0e3928128 --- /dev/null +++ b/python/omp-rpc/tests/test_client.py @@ -0,0 +1,220 @@ +from __future__ import annotations + +import sys +import textwrap +import unittest + +from omp_rpc import RpcClient + + +FAKE_SERVER = textwrap.dedent( + """ + import json + import sys + + def usage(): + return { + "input": 1, + "output": 1, + "cacheRead": 0, + "cacheWrite": 0, + "totalTokens": 2, + "cost": { + "input": 0.0, + "output": 0.0, + "cacheRead": 0.0, + "cacheWrite": 0.0, + "total": 0.0, + }, + } + + def assistant_message(text: str): + return { + "role": "assistant", + "content": [{"type": "text", "text": text}], + "api": "anthropic-messages", + "provider": "anthropic", + "model": "claude-sonnet-4-5", + "usage": usage(), + "stopReason": "stop", + "timestamp": 1, + } + + print(json.dumps({"type": "ready"}), flush=True) + + for raw_line in sys.stdin: + raw_line = raw_line.strip() + if not raw_line: + continue + + command = json.loads(raw_line) + command_type = command["type"] + request_id = command.get("id") + + if command_type == "extension_ui_response": + print(json.dumps({"type": "agent_end", "messages": [assistant_message("ui acknowledged")]}), flush=True) + continue + + if command_type == "get_state": + print( + json.dumps( + { + "id": request_id, + "type": "response", + "command": "get_state", + "success": True, + "data": { + "model": { + "id": "claude-sonnet-4-5", + "name": "Claude Sonnet 4.5", + "api": "anthropic-messages", + "provider": "anthropic", + "baseUrl": "https://api.anthropic.com", + "reasoning": True, + "input": ["text"], + "cost": { + "input": 1.0, + "output": 2.0, + "cacheRead": 0.0, + "cacheWrite": 0.0, + }, + "contextWindow": 200000, + "maxTokens": 8192, + }, + "thinkingLevel": "medium", + "isStreaming": False, + "isCompacting": False, + "steeringMode": "one-at-a-time", + "followUpMode": "one-at-a-time", + "interruptMode": "immediate", + "sessionId": "fake-session", + "autoCompactionEnabled": True, + "messageCount": 0, + "queuedMessageCount": 0, + }, + } + ), + flush=True, + ) + elif command_type == "bash": + print( + json.dumps( + { + "id": request_id, + "type": "response", + "command": "bash", + "success": True, + "data": { + "output": "hello\\n", + "exitCode": 0, + "cancelled": False, + "truncated": False, + "totalLines": 1, + "totalBytes": 6, + "outputLines": 1, + "outputBytes": 6, + }, + } + ), + flush=True, + ) + elif command_type == "prompt": + print( + json.dumps( + { + "id": request_id, + "type": "response", + "command": "prompt", + "success": True, + } + ), + flush=True, + ) + if command["message"] == "needs ui": + print( + json.dumps( + { + "type": "extension_ui_request", + "id": "ui-1", + "method": "input", + "title": "Need input", + "placeholder": "value", + } + ), + flush=True, + ) + continue + + print(json.dumps({"type": "agent_start"}), flush=True) + partial = assistant_message("") + print( + json.dumps( + { + "type": "message_update", + "message": partial, + "assistantMessageEvent": { + "type": "text_delta", + "contentIndex": 0, + "delta": "pong", + "partial": partial, + }, + } + ), + flush=True, + ) + assistant = assistant_message("pong") + print(json.dumps({"type": "message_end", "message": assistant}), flush=True) + print(json.dumps({"type": "agent_end", "messages": [assistant]}), flush=True) + else: + print( + json.dumps( + { + "id": request_id, + "type": "response", + "command": command_type, + "success": False, + "error": f"unsupported: {command_type}", + } + ), + flush=True, + ) + """ +) + + +class RpcClientTests(unittest.TestCase): + def make_client(self) -> RpcClient: + return RpcClient( + command=[sys.executable, "-u", "-c", FAKE_SERVER], + startup_timeout=2.0, + request_timeout=2.0, + ) + + def test_get_state_and_bash(self) -> None: + 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") + + result = client.bash("echo hello") + self.assertEqual(result.output, "hello\n") + self.assertEqual(result.exit_code, 0) + + def test_prompt_and_wait_returns_assistant_text(self) -> None: + with self.make_client() as client: + turn = client.prompt_and_wait("say hello", timeout=2.0) + self.assertEqual(turn.require_assistant_text(), "pong") + self.assertGreaterEqual(len(turn.events), 3) + + def test_extension_ui_round_trip(self) -> None: + with self.make_client() as client: + client.prompt("needs ui") + request = client.next_ui_request(timeout=2.0) + self.assertEqual(request.method, "input") + + client.send_ui_value(request.id, "approved") + client.wait_for_idle(timeout=2.0) + + +if __name__ == "__main__": + unittest.main() diff --git a/python/omp-rpc/tests/test_protocol.py b/python/omp-rpc/tests/test_protocol.py new file mode 100644 index 000000000..bb7fb0795 --- /dev/null +++ b/python/omp-rpc/tests/test_protocol.py @@ -0,0 +1,123 @@ +from __future__ import annotations + +import unittest + +from omp_rpc import ( + AgentEndEvent, + ExtensionUiRequest, + SessionState, + assistant_text, + parse_notification, + parse_session_state, +) + + +class ProtocolParsingTests(unittest.TestCase): + def test_parse_session_state(self) -> None: + state = parse_session_state( + { + "model": { + "id": "claude-sonnet-4-5", + "name": "Claude Sonnet 4.5", + "api": "anthropic-messages", + "provider": "anthropic", + "baseUrl": "https://api.anthropic.com", + "reasoning": True, + "input": ["text", "image"], + "cost": { + "input": 1.0, + "output": 2.0, + "cacheRead": 0.1, + "cacheWrite": 0.2, + }, + "contextWindow": 200000, + "maxTokens": 8192, + "thinking": { + "minLevel": "minimal", + "maxLevel": "high", + "mode": "effort", + }, + }, + "thinkingLevel": "medium", + "isStreaming": False, + "isCompacting": False, + "steeringMode": "one-at-a-time", + "followUpMode": "all", + "interruptMode": "immediate", + "sessionFile": "/tmp/test.jsonl", + "sessionId": "session-123", + "sessionName": "Scratchpad", + "autoCompactionEnabled": True, + "messageCount": 4, + "queuedMessageCount": 1, + "systemPrompt": "You are useful.", + "dumpTools": [ + { + "name": "read", + "description": "Read files", + "parameters": {"type": "object"}, + } + ], + } + ) + + self.assertIsInstance(state, SessionState) + self.assertEqual(state.session_id, "session-123") + self.assertEqual(state.follow_up_mode, "all") + self.assertEqual(state.model.id if state.model else None, "claude-sonnet-4-5") + self.assertEqual(state.dump_tools[0].name, "read") + + def test_parse_agent_end_notification(self) -> None: + notification = parse_notification( + { + "type": "agent_end", + "messages": [ + { + "role": "assistant", + "content": [{"type": "text", "text": "hello"}], + "api": "anthropic-messages", + "provider": "anthropic", + "model": "claude-sonnet-4-5", + "usage": { + "input": 1, + "output": 1, + "cacheRead": 0, + "cacheWrite": 0, + "totalTokens": 2, + "cost": { + "input": 0.0, + "output": 0.0, + "cacheRead": 0.0, + "cacheWrite": 0.0, + "total": 0.0, + }, + }, + "stopReason": "stop", + "timestamp": 1, + } + ], + } + ) + + self.assertIsInstance(notification, AgentEndEvent) + self.assertEqual(assistant_text(notification.messages[0]), "hello") + + def test_parse_extension_ui_request(self) -> None: + notification = parse_notification( + { + "type": "extension_ui_request", + "id": "ui-1", + "method": "confirm", + "title": "Confirm", + "message": "Continue?", + "timeout": 1000, + } + ) + + self.assertIsInstance(notification, ExtensionUiRequest) + self.assertEqual(notification.method, "confirm") + self.assertEqual(notification.message, "Continue?") + + +if __name__ == "__main__": + unittest.main()