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.
This commit is contained in:
can1357
2026-04-08 04:20:06 +02:00
parent 5eeb9d54e8
commit 67e12b39fe
9 changed files with 2148 additions and 0 deletions
+70
View File
@@ -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).
+36
View File
@@ -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"]
+129
View File
@@ -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",
]
+600
View File
@@ -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
+963
View File
@@ -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))
+1
View File
@@ -0,0 +1 @@
+6
View File
@@ -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"))
+220
View File
@@ -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()
+123
View File
@@ -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()