feat(omp-rpc): introduced host tool execution framework with custom tool registration

- Added host tool execution framework with HostTool, HostToolContext, and host_tool() factory for custom tool integration.
- Added RpcConcurrencyError exception and _PromptLifecycleCoordinator to enforce single-flight constraint on prompt lifecycle methods.
- Enhanced JSON parsing with 10 validation helpers and enum frozensets for safe field extraction with detailed error messages.
- Replaced manual event/error list management with _BoundedHistory for bounded-size history with offset tracking.
- Added deep JSON cloning to prevent external mutations of stored payloads and improved UTF-8 error handling in subprocess stderr.
- Added custom_tools parameter to RpcClient and set_custom_tools() method for runtime tool registration.
This commit is contained in:
can1357
2026-04-08 06:26:50 +02:00
parent 821ec9570d
commit d1d0187859
9 changed files with 1781 additions and 281 deletions
+51
View File
@@ -12,6 +12,7 @@ provides:
- a process-backed client that manages request correlation over stdio
- typed per-event listeners plus a typed catch-all notification hook
- helpers for collecting prompt runs and handling extension UI requests in manual or headless mode
- typed host-tool helpers so Python RPC owners can expose custom tools with JSON Schema metadata
## Basic Usage
@@ -97,6 +98,50 @@ with RpcClient(
print(client.get_state().session_id)
```
## Host-Owned Custom Tools
RPC hosts can expose custom tools to the agent with JSON Schema metadata. The
Python helper keeps the wire format simple while still giving the handler a
typed signature:
```python
from typing import TypedDict
from omp_rpc import RpcClient, host_tool
class EchoArgs(TypedDict):
message: str
def echo_host(args: EchoArgs, context) -> str:
context.send_update(f"working:{args['message']}")
return f"host:{args['message']}"
with RpcClient(
no_session=True,
custom_tools=(
host_tool(
name="echo_host",
description="Echo a value from the Python host",
parameters={
"type": "object",
"properties": {"message": {"type": "string"}},
"required": ["message"],
"additionalProperties": False,
},
execute=echo_host,
),
),
) as client:
client.prompt_and_wait("Use the echo_host tool with the value hello")
```
If you want runtime conversion into a richer Python type, pass `decode=` to
`host_tool(...)`. That lets you keep the JSON Schema contract on the wire while
parsing the incoming argument object into a dataclass or model in the handler.
## Extension UI Requests
Extensions in RPC mode can ask the host for input. Those requests are available as
@@ -152,6 +197,12 @@ If a single prompt streams more events than `max_event_history` allows,
`prompt_and_wait()` raises a clear error so hosts can increase the limit instead
of silently losing earlier events.
Prompt lifecycle collection is intentionally single-flight. Only one of
`prompt_and_wait()`, `wait_for_idle()`, or `collect_events()` may be active at a
time on a client instance. If a host needs concurrent orchestration, use
separate `RpcClient` instances instead of overlapping lifecycle waiters on one
session.
## Text Helpers
`assistant_text()` and `message_text()` now return visible text blocks only.
+8
View File
@@ -8,6 +8,7 @@ from .client import (
ReadyListener,
RpcClient,
RpcCommandError,
RpcConcurrencyError,
RpcError,
RpcProcessExitError,
RpcProtocolError,
@@ -15,6 +16,7 @@ from .client import (
ProtocolErrorListener,
UiRequestListener,
)
from .host_tools import HostTool, HostToolContext, HostToolResultPayload, HostToolResultValue, host_tool
from .protocol import (
AgentEndEvent,
AgentMessage,
@@ -100,6 +102,10 @@ __all__ = [
"ExtensionErrorListener",
"ExtensionUiRequest",
"FileMentionMessage",
"HostTool",
"HostToolContext",
"HostToolResultPayload",
"HostToolResultValue",
"HookMessage",
"ImageContent",
"ListenerErrorEvent",
@@ -121,6 +127,7 @@ __all__ = [
"RpcAgentEvent",
"RpcClient",
"RpcCommandError",
"RpcConcurrencyError",
"RpcError",
"RpcNotification",
"RpcProcessExitError",
@@ -154,4 +161,5 @@ __all__ = [
"parse_notification",
"parse_session_state",
"parse_todo_phases",
"host_tool",
]
+360 -74
View File
@@ -6,10 +6,11 @@ import queue
import subprocess
import threading
import time
from dataclasses import dataclass
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any, Callable, Mapping, Sequence, TypeVar, cast
from typing import Any, Callable, Generic, Mapping, Sequence, TypeVar, cast
from .host_tools import HostTool, HostToolContext
from .protocol import (
AgentStartEvent,
AgentEndEvent,
@@ -59,6 +60,7 @@ from .protocol import (
TurnStartEvent,
UnknownNotification,
assistant_text,
parse_agent_messages,
parse_bash_result,
parse_branch_messages,
parse_branch_result,
@@ -102,9 +104,32 @@ ProtocolErrorListener = Callable[["RpcProtocolError"], None]
ListenerErrorListener = Callable[["ListenerErrorEvent"], None]
TListener = TypeVar("TListener")
TEventListener = TypeVar("TEventListener", bound=Callable[..., None])
THistoryItem = TypeVar("THistoryItem")
_ASYNC_COMMANDS = frozenset({"prompt", "abort_and_prompt"})
_DEFAULT_ERROR_HISTORY_LIMIT = 128
_TODO_STATUS_VALUES = frozenset({"pending", "in_progress", "completed", "abandoned"})
def _clone_json_value(value: object) -> JsonValue:
if value is None or isinstance(value, (str, int, float, bool)):
return cast(JsonValue, value)
if isinstance(value, list):
return [_clone_json_value(item) for item in value]
if isinstance(value, dict):
cloned: JsonObject = {}
for key, item in value.items():
if not isinstance(key, str):
raise RpcError("RPC payload objects must use string keys")
cloned[key] = _clone_json_value(item)
return cloned
raise RpcError("RPC payload must be JSON-serializable")
def _clone_json_object(value: object) -> JsonObject:
if not isinstance(value, dict):
raise RpcError("RPC response payload must be an object")
return cast(JsonObject, _clone_json_value(value))
class RpcError(RuntimeError):
@@ -119,6 +144,10 @@ class RpcProcessExitError(RpcError):
"""Raised when the RPC process exits while a request is pending."""
class RpcConcurrencyError(RpcError):
"""Raised when overlapping prompt lifecycle collectors would be ambiguous."""
class RpcCommandError(RpcError):
"""Raised when the RPC server returns `success: false`."""
@@ -181,6 +210,57 @@ class _PendingRequest:
response_queue: queue.Queue[JsonObject | BaseException]
@dataclass(slots=True)
class _PendingHostToolCall:
cancel_event: threading.Event
@dataclass(slots=True)
class _BoundedHistory(Generic[THistoryItem]):
limit: int | None
items: list[THistoryItem] = field(default_factory=list)
offset: int = 0
def clear(self) -> None:
self.items.clear()
self.offset = 0
def append(self, item: THistoryItem) -> None:
self.items.append(item)
if self.limit is not None and len(self.items) > self.limit:
trim = len(self.items) - self.limit
del self.items[:trim]
self.offset += trim
def current_index(self) -> int:
return self.offset + len(self.items)
def snapshot(self) -> tuple[THistoryItem, ...]:
return tuple(self.items)
def snapshot_from(self, start_index: int) -> tuple[THistoryItem, ...]:
return tuple(self.items[start_index - self.offset :])
@dataclass(slots=True)
class _PromptLifecycleCoordinator:
lock: threading.Lock = field(default_factory=threading.Lock)
active_operation: str | None = None
def acquire(self, operation: str) -> None:
with self.lock:
if self.active_operation is not None:
raise RpcConcurrencyError(
f"Cannot start {operation} while {self.active_operation} is already collecting prompt lifecycle events"
)
self.active_operation = operation
def release(self, operation: str) -> None:
with self.lock:
if self.active_operation == operation:
self.active_operation = None
class RpcClient:
def __init__(
self,
@@ -196,6 +276,7 @@ class RpcClient:
append_system_prompt: str | None = None,
provider_session_id: str | None = None,
tools: Sequence[str] | None = None,
custom_tools: Sequence[HostTool[Any, Any]] | None = None,
no_session: bool = False,
no_skills: bool = False,
no_rules: bool = False,
@@ -218,6 +299,7 @@ class RpcClient:
self._append_system_prompt = append_system_prompt
self._provider_session_id = provider_session_id
self._tools = tuple(tools) if tools is not None else None
self._custom_tools = tuple(custom_tools) if custom_tools is not None else ()
self._no_session = no_session
self._no_skills = no_skills
self._no_rules = no_rules
@@ -237,17 +319,20 @@ class RpcClient:
self._state_lock = threading.Lock()
self._event_condition = threading.Condition()
self._pending: dict[str, _PendingRequest] = {}
self._pending_host_tool_calls: dict[str, _PendingHostToolCall] = {}
self._request_id = 0
self._events: list[RpcAgentEvent] = []
self._event_offset = 0
self._async_errors: list[BaseException] = []
self._async_error_offset = 0
self._events = _BoundedHistory[JsonObject](self._max_event_history)
self._async_errors = _BoundedHistory[BaseException](_DEFAULT_ERROR_HISTORY_LIMIT)
self._scheduled_agent_runs = 0
self._completed_agent_runs = 0
self._ui_requests: queue.Queue[ExtensionUiRequest] = queue.Queue()
self._stderr_chunks: list[str] = []
self._stderr_chunks = _BoundedHistory[str](self._max_stderr_chunks)
self._closed_error: BaseException | None = None
self._stopping = False
self._protocol_errors: list[RpcProtocolError] = []
self._listener_errors: list[ListenerErrorEvent] = []
self._ready_received = False
self._protocol_errors = _BoundedHistory[RpcProtocolError](_DEFAULT_ERROR_HISTORY_LIMIT)
self._listener_errors = _BoundedHistory[ListenerErrorEvent](_DEFAULT_ERROR_HISTORY_LIMIT)
self._prompt_lifecycle = _PromptLifecycleCoordinator()
self._notification_listeners: list[NotificationListener] = []
self._event_listeners: list[AgentEventListener] = []
@@ -267,7 +352,8 @@ class RpcClient:
@property
def stderr(self) -> str:
return "".join(self._stderr_chunks)
with self._state_lock:
return "".join(self._stderr_chunks.snapshot())
@property
def command(self) -> tuple[str, ...]:
@@ -276,12 +362,12 @@ class RpcClient:
@property
def protocol_errors(self) -> tuple[RpcProtocolError, ...]:
with self._state_lock:
return tuple(self._protocol_errors)
return self._protocol_errors.snapshot()
@property
def listener_errors(self) -> tuple[ListenerErrorEvent, ...]:
with self._state_lock:
return tuple(self._listener_errors)
return self._listener_errors.snapshot()
def start(self) -> RpcClient:
if self._process is not None:
@@ -290,12 +376,14 @@ class RpcClient:
self._ready.clear()
self._stopping = False
self._closed_error = None
self._ready_received = False
self._events.clear()
self._event_offset = 0
self._async_errors.clear()
self._async_error_offset = 0
self._scheduled_agent_runs = 0
self._completed_agent_runs = 0
self._ui_requests = queue.Queue()
self._stderr_chunks.clear()
with self._state_lock:
self._stderr_chunks.clear()
with self._state_lock:
self._protocol_errors.clear()
self._listener_errors.clear()
@@ -309,6 +397,7 @@ class RpcClient:
stderr=subprocess.PIPE,
text=True,
encoding="utf-8",
errors="replace",
bufsize=1,
)
self._process = process
@@ -323,6 +412,18 @@ class RpcClient:
self.stop()
raise RpcTimeoutError(f"Timed out waiting for RPC ready signal. Stderr: {stderr}")
if not self._ready_received:
error = self._closed_error
stderr = self.stderr
self.stop()
if isinstance(error, RpcError):
raise error
if error is not None:
raise RpcProcessExitError(f"RPC process stopped before ready: {error}. Stderr: {stderr}") from error
raise RpcTimeoutError(f"Timed out waiting for RPC ready signal. Stderr: {stderr}")
if self._custom_tools:
self.set_custom_tools(self._custom_tools)
return self
def stop(self) -> None:
@@ -331,6 +432,8 @@ class RpcClient:
return
self._stopping = True
for pending_call in self._pending_host_tool_calls.values():
pending_call.cancel_event.set()
try:
if process.stdin is not None:
@@ -358,6 +461,7 @@ class RpcClient:
except OSError:
pass
self._fail_pending(RpcProcessExitError("RPC process stopped"))
self._pending_host_tool_calls.clear()
self._process = None
self._ready.set()
with self._event_condition:
@@ -626,7 +730,33 @@ class RpcClient:
def get_messages(self) -> tuple[AgentMessage, ...]:
payload = self._request("get_messages")
return tuple(cast(list[AgentMessage], payload.get("messages") or []))
return parse_agent_messages(cast(JsonValue | None, payload.get("messages")))
def set_custom_tools(self, tools: Sequence[HostTool[Any, Any]]) -> tuple[str, ...]:
self._custom_tools = tuple(tools)
if self._process is None:
return tuple(tool.name for tool in self._custom_tools)
payload = self._request(
"set_host_tools",
tools=cast(
JsonValue,
[
{
"name": tool.name,
"label": tool.label,
"description": tool.description,
"parameters": tool.parameters,
"hidden": tool.hidden,
}
for tool in self._custom_tools
],
),
)
tool_names = payload.get("toolNames") or []
if not isinstance(tool_names, list):
raise RpcError("set_host_tools response did not include toolNames")
return tuple(str(name) for name in tool_names)
def prompt(
self,
@@ -641,6 +771,7 @@ class RpcClient:
images=list(images) if images is not None else None,
streamingBehavior=streaming_behavior,
)
self._mark_agent_run_scheduled()
def steer(self, message: str, *, images: Sequence[ImageContent] | None = None) -> None:
self._request("steer", message=message, images=list(images) if images is not None else None)
@@ -653,6 +784,7 @@ class RpcClient:
def abort_and_prompt(self, message: str, *, images: Sequence[ImageContent] | None = None) -> None:
self._request("abort_and_prompt", message=message, images=list(images) if images is not None else None)
self._mark_agent_run_scheduled()
def prompt_and_wait(
self,
@@ -662,32 +794,62 @@ class RpcClient:
streaming_behavior: StreamingBehavior | None = None,
timeout: float | None = None,
) -> PromptTurn:
start_index = self._current_event_index()
start_async_error_index = self._current_async_error_index()
self.prompt(message, images=images, streaming_behavior=streaming_behavior)
events = self._wait_for_agent_end(start_index, start_async_error_index, timeout=timeout)
return self._build_prompt_turn(events)
operation = "prompt_and_wait"
self._prompt_lifecycle.acquire(operation)
try:
start_index = self._current_event_index()
start_async_error_index = self._current_async_error_index()
self.prompt(message, images=images, streaming_behavior=streaming_behavior)
events = self._wait_for_agent_end(start_index, start_async_error_index, timeout=timeout)
return self._build_prompt_turn(events)
finally:
self._prompt_lifecycle.release(operation)
def wait_for_idle(self, timeout: float | None = None) -> None:
start_index = self._current_event_index()
start_async_error_index = self._current_async_error_index()
self._wait_for_agent_end(start_index, start_async_error_index, timeout=timeout)
operation = "wait_for_idle"
self._prompt_lifecycle.acquire(operation)
try:
if self._is_agent_idle():
return
start_index = self._current_event_index()
start_async_error_index = self._current_async_error_index()
self._wait_for_agent_end(start_index, start_async_error_index, timeout=timeout)
finally:
self._prompt_lifecycle.release(operation)
def collect_events(self, timeout: float | None = None) -> tuple[RpcAgentEvent, ...]:
start_index = self._current_event_index()
start_async_error_index = self._current_async_error_index()
return self._wait_for_agent_end(start_index, start_async_error_index, timeout=timeout)
operation = "collect_events"
self._prompt_lifecycle.acquire(operation)
try:
start_index = self._current_event_index()
start_async_error_index = self._current_async_error_index()
return self._wait_for_agent_end(start_index, start_async_error_index, timeout=timeout)
finally:
self._prompt_lifecycle.release(operation)
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 self._event_offset + len(self._events)
return self._events.current_index()
def _current_async_error_index(self) -> int:
with self._event_condition:
return self._async_error_offset + len(self._async_errors)
return self._async_errors.current_index()
def _mark_agent_run_scheduled(self) -> None:
with self._event_condition:
self._scheduled_agent_runs += 1
def _mark_agent_run_completed(self) -> None:
with self._event_condition:
self._completed_agent_runs += 1
self._event_condition.notify_all()
def _is_agent_idle(self) -> bool:
with self._event_condition:
return self._scheduled_agent_runs == self._completed_agent_runs
def _build_prompt_turn(self, events: tuple[RpcAgentEvent, ...]) -> PromptTurn:
final_messages: tuple[AgentMessage, ...] = ()
@@ -729,25 +891,25 @@ class RpcClient:
if self._closed_error is not None:
raise RpcProcessExitError(str(self._closed_error))
if start_index < self._event_offset:
if start_index < self._events.offset:
raise RpcError(
"Event history limit was exceeded while waiting for agent_end. "
"Increase max_event_history to retain more streamed events."
)
if start_async_error_index < self._async_error_offset:
if start_async_error_index < self._async_errors.offset:
raise RpcError(
"Async error history limit was exceeded while waiting for agent_end. "
"Increase max_event_history if your host needs to retain more background failures."
)
async_error_index = start_async_error_index - self._async_error_offset
if async_error_index < len(self._async_errors):
raise self._async_errors[async_error_index]
async_errors = self._async_errors.snapshot_from(start_async_error_index)
if len(async_errors) > 0:
raise async_errors[0]
event_index = start_index - self._event_offset
events = tuple(self._events[event_index:])
if any(isinstance(event, AgentEndEvent) for event in events):
event_payloads = self._events.snapshot_from(start_index)
if any(payload.get("type") == "agent_end" for payload in event_payloads):
events = tuple(cast(RpcAgentEvent, parse_notification(payload)) for payload in event_payloads)
return events
remaining = deadline - time.monotonic()
@@ -767,7 +929,12 @@ class RpcClient:
with self._state_lock:
self._pending[request_id] = _PendingRequest(command=command_type, response_queue=response_queue)
self._write_json(process, envelope)
try:
self._write_json(process, envelope)
except BaseException:
with self._state_lock:
self._pending.pop(request_id, None)
raise
try:
response = response_queue.get(timeout=self._request_timeout)
@@ -783,12 +950,101 @@ class RpcClient:
raise RpcCommandError(command=str(response.get("command", command_type)), error=str(response.get("error", "")))
data = response.get("data")
return dict(cast(JsonObject, data or {}))
if data is None:
return {}
return _clone_json_object(data)
def _send_notification(self, payload: JsonObject) -> None:
process = self._require_process()
self._write_json(process, payload)
def _normalize_host_tool_result(self, result: object) -> JsonObject:
if isinstance(result, str):
return {"content": [{"type": "text", "text": result}]}
if isinstance(result, Mapping):
return cast(JsonObject, dict(result))
raise RpcError("Host tool handlers must return a string or a result mapping")
def _handle_host_tool_call(self, payload: JsonObject) -> None:
request_id = payload.get("id")
tool_name = payload.get("toolName")
tool_call_id = payload.get("toolCallId")
raw_arguments = payload.get("arguments")
if not isinstance(request_id, str) or not isinstance(tool_name, str) or not isinstance(tool_call_id, str):
return
if not isinstance(raw_arguments, Mapping):
self._send_notification(
{
"type": "host_tool_result",
"id": request_id,
"result": {"content": [{"type": "text", "text": "Host tool arguments must be an object"}], "details": {}},
"isError": True,
}
)
return
tool = next((candidate for candidate in self._custom_tools if candidate.name == tool_name), None)
if tool is None:
self._send_notification(
{
"type": "host_tool_result",
"id": request_id,
"result": {
"content": [{"type": "text", "text": f'Host tool "{tool_name}" is not registered'}],
"details": {},
},
"isError": True,
}
)
return
pending_call = _PendingHostToolCall(cancel_event=threading.Event())
self._pending_host_tool_calls[request_id] = pending_call
def run_tool() -> None:
try:
params = tool.parse_params(cast(JsonObject, dict(raw_arguments)))
context = HostToolContext(
tool_call_id=tool_call_id,
_cancel_event=pending_call.cancel_event,
_send_update=lambda result: self._send_notification(
{"type": "host_tool_update", "id": request_id, "partialResult": result}
),
)
result = tool.execute(params, context)
if pending_call.cancel_event.is_set():
return
self._send_notification(
{
"type": "host_tool_result",
"id": request_id,
"result": self._normalize_host_tool_result(result),
}
)
except Exception as exc:
if pending_call.cancel_event.is_set():
return
self._send_notification(
{
"type": "host_tool_result",
"id": request_id,
"result": {"content": [{"type": "text", "text": str(exc)}], "details": {}},
"isError": True,
}
)
finally:
self._pending_host_tool_calls.pop(request_id, None)
threading.Thread(target=run_tool, name=f"omp-rpc-host-tool:{tool_name}", daemon=True).start()
def _handle_host_tool_cancel(self, payload: JsonObject) -> None:
target_id = payload.get("targetId")
if not isinstance(target_id, str):
return
pending_call = self._pending_host_tool_calls.get(target_id)
if pending_call is not None:
pending_call.cancel_event.set()
def _add_typed_event_listener(self, event_type: str, listener: TEventListener) -> Callable[[], None]:
listeners = self._typed_event_listeners.setdefault(event_type, [])
typed_listener = cast(AgentEventListener, listener)
@@ -813,6 +1069,8 @@ class RpcClient:
return {"id": next_task(), "content": seed, "status": cast(JsonValue, "pending")}
if isinstance(seed, TodoItem):
if seed.status not in _TODO_STATUS_VALUES:
raise RpcError(f"Unsupported todo status: {seed.status}")
return {
"id": seed.id or next_task(),
"content": seed.content,
@@ -829,7 +1087,12 @@ class RpcClient:
raw_status = seed.get("status")
raw_notes = seed.get("notes")
raw_details = seed.get("details")
status: TodoStatus = cast(TodoStatus, raw_status) if isinstance(raw_status, str) else "pending"
if isinstance(raw_status, str):
if raw_status not in _TODO_STATUS_VALUES:
raise RpcError(f"Unsupported todo status: {raw_status}")
status: TodoStatus = cast(TodoStatus, raw_status)
else:
status = "pending"
return {
"id": str(raw_id) if isinstance(raw_id, str) and raw_id else next_task(),
"content": content,
@@ -934,50 +1197,84 @@ class RpcClient:
if process is None or process.stdout is None:
return
line_number = 0
try:
for line in process.stdout:
line_number += 1
stripped = line.strip()
if not stripped:
continue
payload = cast(JsonObject, json.loads(stripped))
try:
payload = cast(JsonObject, json.loads(stripped))
except json.JSONDecodeError as exc:
snippet = stripped
if len(snippet) > 240:
snippet = f"{snippet[:237]}..."
raise RpcError(
f"Failed to decode RPC output on line {line_number}: {exc}. Frame: {snippet!r}"
) from exc
if payload.get("type") == "response":
self._handle_response(payload)
continue
if payload.get("type") == "host_tool_call":
self._handle_host_tool_call(payload)
continue
if payload.get("type") == "host_tool_cancel":
self._handle_host_tool_cancel(payload)
continue
notification = parse_notification(payload)
self._dispatch_listeners("notification", notification.type, self._notification_listeners, notification)
listener_notification = parse_notification(payload)
self._dispatch_listeners(
"notification",
listener_notification.type,
self._notification_listeners,
listener_notification,
)
if isinstance(notification, ReadyEvent):
self._ready_received = True
self._ready.set()
self._dispatch_listeners("ready", notification.type, self._ready_listeners, notification)
self._dispatch_listeners("ready", listener_notification.type, self._ready_listeners, listener_notification)
continue
if isinstance(notification, ExtensionUiRequest):
self._ui_requests.put(notification)
self._dispatch_listeners("ui_request", notification.type, self._ui_request_listeners, notification)
self._dispatch_listeners(
"ui_request",
listener_notification.type,
self._ui_request_listeners,
cast(ExtensionUiRequest, listener_notification),
)
continue
if isinstance(notification, ExtensionError):
self._dispatch_listeners(
"extension_error", notification.type, self._extension_error_listeners, notification
"extension_error",
listener_notification.type,
self._extension_error_listeners,
cast(ExtensionError, listener_notification),
)
continue
if isinstance(notification, UnknownNotification):
self._dispatch_listeners(
"unknown_notification", notification.type, self._unknown_notification_listeners, notification
"unknown_notification",
listener_notification.type,
self._unknown_notification_listeners,
cast(UnknownNotification, listener_notification),
)
continue
event = cast(RpcAgentEvent, notification)
self._append_event(event)
self._dispatch_listeners("event", event.type, self._event_listeners, event)
listener_event = cast(RpcAgentEvent, listener_notification)
self._append_event(payload)
if listener_event.type == "agent_end":
self._mark_agent_run_completed()
self._dispatch_listeners("event", listener_event.type, self._event_listeners, listener_event)
self._dispatch_listeners(
"typed_event", event.type, self._typed_event_listeners.get(event.type, []), event
"typed_event", listener_event.type, self._typed_event_listeners.get(listener_event.type, []), listener_event
)
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:
@@ -995,11 +1292,13 @@ class RpcClient:
process = self._process
if process is None or process.stderr is None:
return
for chunk in process.stderr:
self._stderr_chunks.append(chunk)
if self._max_stderr_chunks is not None and len(self._stderr_chunks) > self._max_stderr_chunks:
trim = len(self._stderr_chunks) - self._max_stderr_chunks
del self._stderr_chunks[:trim]
try:
for chunk in process.stderr:
with self._state_lock:
self._stderr_chunks.append(chunk)
except Exception as exc:
if not self._stopping:
self._mark_closed(RpcError(f"Failed to read RPC stderr: {exc}"))
def _mark_closed(self, error: BaseException) -> None:
if self._closed_error is not None:
@@ -1035,6 +1334,7 @@ class RpcClient:
if protocol_error.command in _ASYNC_COMMANDS and protocol_error.remote_error is not None:
self._append_async_error(RpcCommandError(protocol_error.command, protocol_error.remote_error))
self._mark_agent_run_completed()
self._record_protocol_error(protocol_error)
@@ -1067,40 +1367,26 @@ class RpcClient:
return None
if bool(payload.get("success", False)):
return None
return RpcProtocolError(payload)
return RpcProtocolError(_clone_json_object(payload))
def _append_event(self, event: RpcAgentEvent) -> None:
def _append_event(self, payload: JsonObject) -> None:
with self._event_condition:
self._events.append(event)
if self._max_event_history is not None and len(self._events) > self._max_event_history:
trim = len(self._events) - self._max_event_history
del self._events[:trim]
self._event_offset += trim
self._events.append(_clone_json_object(payload))
self._event_condition.notify_all()
def _append_async_error(self, error: BaseException) -> None:
with self._event_condition:
self._async_errors.append(error)
if len(self._async_errors) > _DEFAULT_ERROR_HISTORY_LIMIT:
trim = len(self._async_errors) - _DEFAULT_ERROR_HISTORY_LIMIT
del self._async_errors[:trim]
self._async_error_offset += trim
self._event_condition.notify_all()
def _record_protocol_error(self, error: RpcProtocolError) -> None:
with self._state_lock:
self._protocol_errors.append(error)
if len(self._protocol_errors) > _DEFAULT_ERROR_HISTORY_LIMIT:
trim = len(self._protocol_errors) - _DEFAULT_ERROR_HISTORY_LIMIT
del self._protocol_errors[:trim]
self._dispatch_listeners("protocol_error", error.command, self._protocol_error_listeners, error)
def _record_listener_error(self, event: ListenerErrorEvent) -> None:
with self._state_lock:
self._listener_errors.append(event)
if len(self._listener_errors) > _DEFAULT_ERROR_HISTORY_LIMIT:
trim = len(self._listener_errors) - _DEFAULT_ERROR_HISTORY_LIMIT
del self._listener_errors[:trim]
for listener in list(self._listener_error_listeners):
try:
+80
View File
@@ -0,0 +1,80 @@
from __future__ import annotations
import threading
from dataclasses import dataclass
from typing import Callable, Generic, TypeAlias, TypeVar, TypedDict, cast
from .protocol import ImageContent, JsonObject, JsonValue, TextContent
TParams = TypeVar("TParams")
TDetails = TypeVar("TDetails")
class HostToolResultPayload(TypedDict, total=False):
content: list[TextContent | ImageContent]
details: JsonValue
HostToolResultValue: TypeAlias = HostToolResultPayload | str
def _normalize_result(result: HostToolResultValue) -> JsonObject:
if isinstance(result, str):
return {"content": [{"type": "text", "text": result}]}
return dict(result)
@dataclass(slots=True)
class HostToolContext(Generic[TDetails]):
tool_call_id: str
_cancel_event: threading.Event
_send_update: Callable[[JsonObject], None]
@property
def cancelled(self) -> bool:
return self._cancel_event.is_set()
def send_update(self, result: HostToolResultValue) -> None:
if self.cancelled:
return
self._send_update(_normalize_result(result))
@dataclass(slots=True, frozen=True)
class HostTool(Generic[TParams, TDetails]):
name: str
description: str
parameters: JsonObject
execute: Callable[[TParams, HostToolContext[TDetails]], HostToolResultValue]
label: str | None = None
hidden: bool = False
decode: Callable[[JsonObject], TParams] | None = None
def parse_params(self, payload: JsonObject) -> TParams:
if self.decode is not None:
return self.decode(payload)
return cast(TParams, payload)
def normalize_result(self, result: HostToolResultValue) -> JsonObject:
return _normalize_result(result)
def host_tool(
*,
name: str,
description: str,
parameters: JsonObject,
execute: Callable[[TParams, HostToolContext[TDetails]], HostToolResultValue],
label: str | None = None,
hidden: bool = False,
decode: Callable[[JsonObject], TParams] | None = None,
) -> HostTool[TParams, TDetails]:
return HostTool(
name=name,
description=description,
parameters=dict(parameters),
execute=execute,
label=label,
hidden=hidden,
decode=decode,
)
+439 -82
View File
@@ -42,6 +42,278 @@ INTERACTIVE_EXTENSION_UI_METHODS: Final[frozenset[InteractiveExtensionUiMethod]]
{"select", "confirm", "input", "editor"}
)
VALUE_EXTENSION_UI_METHODS: Final[frozenset[ValueExtensionUiMethod]] = frozenset({"select", "input", "editor"})
_THINKING_LEVEL_VALUES: Final[frozenset[str]] = frozenset({"off", "minimal", "low", "medium", "high", "xhigh"})
_STEERING_MODE_VALUES: Final[frozenset[str]] = frozenset({"all", "one-at-a-time"})
_INTERRUPT_MODE_VALUES: Final[frozenset[str]] = frozenset({"immediate", "wait"})
_STOP_REASON_VALUES: Final[frozenset[str]] = frozenset({"stop", "length", "toolUse", "error", "aborted"})
_NOTIFY_TYPE_VALUES: Final[frozenset[str]] = frozenset({"info", "warning", "error"})
_WIDGET_PLACEMENT_VALUES: Final[frozenset[str]] = frozenset({"aboveEditor", "belowEditor"})
_TODO_STATUS_VALUES: Final[frozenset[str]] = frozenset({"pending", "in_progress", "completed", "abandoned"})
_EXTENSION_UI_METHOD_VALUES: Final[frozenset[str]] = frozenset(
{
"select",
"confirm",
"input",
"editor",
"cancel",
"notify",
"setStatus",
"setWidget",
"setTitle",
"set_editor_text",
}
)
_AGENT_MESSAGE_ROLE_VALUES: Final[frozenset[str]] = frozenset(
{
"user",
"developer",
"assistant",
"toolResult",
"bashExecution",
"pythonExecution",
"custom",
"hookMessage",
"branchSummary",
"compactionSummary",
"fileMention",
}
)
_ASSISTANT_MESSAGE_EVENT_TYPE_VALUES: Final[frozenset[str]] = frozenset(
{
"start",
"text_start",
"text_delta",
"text_end",
"thinking_start",
"thinking_delta",
"thinking_end",
"toolcall_start",
"toolcall_delta",
"toolcall_end",
"done",
"error",
}
)
_ASSISTANT_DONE_REASON_VALUES: Final[frozenset[str]] = frozenset({"stop", "length", "toolUse"})
_ASSISTANT_ERROR_REASON_VALUES: Final[frozenset[str]] = frozenset({"aborted", "error"})
_AUTO_COMPACTION_REASON_VALUES: Final[frozenset[str]] = frozenset({"threshold", "overflow", "idle"})
_AUTO_COMPACTION_ACTION_VALUES: Final[frozenset[str]] = frozenset({"context-full", "handoff"})
def _clone_json_value(value: object, *, field: str) -> JsonValue:
if value is None or isinstance(value, (str, int, float, bool)):
return cast(JsonValue, value)
if isinstance(value, list):
return [_clone_json_value(item, field=field) for item in value]
if isinstance(value, dict):
cloned: JsonObject = {}
for key, item in value.items():
if not isinstance(key, str):
raise ValueError(f"{field} must contain string keys")
cloned[key] = _clone_json_value(item, field=field)
return cloned
raise ValueError(f"{field} must be JSON-serializable")
def _clone_json_object(value: object, *, field: str) -> JsonObject:
if not isinstance(value, dict):
raise ValueError(f"{field} must be an object")
return cast(JsonObject, _clone_json_value(value, field=field))
def _optional_json_object(value: object, *, field: str) -> JsonObject | None:
if value is None:
return None
return _clone_json_object(value, field=field)
def _clone_json_objects(values: object, *, field: str) -> tuple[JsonObject, ...]:
if values is None:
return ()
if not isinstance(values, list):
raise ValueError(f"{field} must be a list")
return tuple(_clone_json_object(item, field=f"{field}[]") for item in values)
def _require_literal(value: object, allowed: frozenset[str], *, field: str) -> str:
if not isinstance(value, str) or value not in allowed:
expected = ", ".join(sorted(allowed))
raise ValueError(f"{field} must be one of: {expected}")
return value
def _optional_literal(value: object, allowed: frozenset[str], *, field: str) -> str | None:
if value is None:
return None
return _require_literal(value, allowed, field=field)
def _require_str(payload: JsonObject, field: str) -> str:
value = payload.get(field)
if not isinstance(value, str):
raise ValueError(f"{field} must be a string")
return value
def _optional_str(payload: JsonObject, field: str) -> str | None:
value = payload.get(field)
if value is None:
return None
if not isinstance(value, str):
raise ValueError(f"{field} must be a string")
return value
def _optional_bool(payload: JsonObject, field: str) -> bool | None:
value = payload.get(field)
if value is None:
return None
if not isinstance(value, bool):
raise ValueError(f"{field} must be a boolean")
return value
def _optional_int(payload: JsonObject, field: str) -> int | None:
value = payload.get(field)
if value is None:
return None
if isinstance(value, bool) or not isinstance(value, int):
raise ValueError(f"{field} must be an integer")
return value
def _tuple_of_strings(values: object, *, field: str) -> tuple[str, ...] | None:
if values is None:
return None
if not isinstance(values, list):
raise ValueError(f"{field} must be a list")
result: list[str] = []
for item in values:
if not isinstance(item, str):
raise ValueError(f"{field} must contain only strings")
result.append(item)
return tuple(result) or None
def _parse_agent_message(payload: JsonObject, *, field: str) -> AgentMessage:
_require_literal(payload.get("role"), _AGENT_MESSAGE_ROLE_VALUES, field=f"{field}.role")
return cast(AgentMessage, _clone_json_object(payload, field=field))
def _parse_assistant_message(payload: JsonObject, *, field: str) -> AssistantMessage:
message = _parse_agent_message(payload, field=field)
if message.get("role") != "assistant":
raise ValueError(f"{field}.role must be 'assistant'")
return cast(AssistantMessage, message)
def _parse_tool_result_message(payload: JsonObject, *, field: str) -> ToolResultMessage:
message = _parse_agent_message(payload, field=field)
if message.get("role") != "toolResult":
raise ValueError(f"{field}.role must be 'toolResult'")
return cast(ToolResultMessage, message)
def parse_agent_messages(payload: JsonValue | None) -> tuple[AgentMessage, ...]:
if payload is None:
return ()
if not isinstance(payload, list):
raise ValueError("messages must be a list")
messages: list[AgentMessage] = []
for index, item in enumerate(payload):
messages.append(_parse_agent_message(_clone_json_object(item, field=f"messages[{index}]"), field=f"messages[{index}]"))
return tuple(messages)
def parse_assistant_message_event(payload: JsonObject) -> AssistantMessageEvent:
event_type = _require_literal(
payload.get("type"),
_ASSISTANT_MESSAGE_EVENT_TYPE_VALUES,
field="assistantMessageEvent.type",
)
if event_type == "start":
return AssistantMessageStartEvent(
partial=_parse_assistant_message(
_clone_json_object(payload.get("partial"), field="assistantMessageEvent.partial"),
field="assistantMessageEvent.partial",
)
)
if event_type in {"text_start", "thinking_start", "toolcall_start"}:
partial = _parse_assistant_message(
_clone_json_object(payload.get("partial"), field="assistantMessageEvent.partial"),
field="assistantMessageEvent.partial",
)
content_index = _optional_int(payload, "contentIndex")
if content_index is None:
raise ValueError("assistantMessageEvent.contentIndex must be an integer")
if event_type == "text_start":
return AssistantTextStartEvent(contentIndex=content_index, partial=partial)
if event_type == "thinking_start":
return AssistantThinkingStartEvent(contentIndex=content_index, partial=partial)
return AssistantToolCallStartEvent(contentIndex=content_index, partial=partial)
if event_type in {"text_delta", "thinking_delta", "toolcall_delta"}:
partial = _parse_assistant_message(
_clone_json_object(payload.get("partial"), field="assistantMessageEvent.partial"),
field="assistantMessageEvent.partial",
)
content_index = _optional_int(payload, "contentIndex")
delta = _optional_str(payload, "delta")
if content_index is None:
raise ValueError("assistantMessageEvent.contentIndex must be an integer")
if delta is None:
raise ValueError("assistantMessageEvent.delta must be a string")
if event_type == "text_delta":
return AssistantTextDeltaEvent(contentIndex=content_index, delta=delta, partial=partial)
if event_type == "thinking_delta":
return AssistantThinkingDeltaEvent(contentIndex=content_index, delta=delta, partial=partial)
return AssistantToolCallDeltaEvent(contentIndex=content_index, delta=delta, partial=partial)
if event_type in {"text_end", "thinking_end"}:
partial = _parse_assistant_message(
_clone_json_object(payload.get("partial"), field="assistantMessageEvent.partial"),
field="assistantMessageEvent.partial",
)
content_index = _optional_int(payload, "contentIndex")
content = _optional_str(payload, "content")
if content_index is None:
raise ValueError("assistantMessageEvent.contentIndex must be an integer")
if content is None:
raise ValueError("assistantMessageEvent.content must be a string")
if event_type == "text_end":
return AssistantTextEndEvent(contentIndex=content_index, content=content, partial=partial)
return AssistantThinkingEndEvent(contentIndex=content_index, content=content, partial=partial)
if event_type == "toolcall_end":
partial = _parse_assistant_message(
_clone_json_object(payload.get("partial"), field="assistantMessageEvent.partial"),
field="assistantMessageEvent.partial",
)
content_index = _optional_int(payload, "contentIndex")
if content_index is None:
raise ValueError("assistantMessageEvent.contentIndex must be an integer")
tool_call = _clone_json_object(payload.get("toolCall"), field="assistantMessageEvent.toolCall")
return AssistantToolCallEndEvent(contentIndex=content_index, toolCall=cast(ToolCall, tool_call), partial=partial)
if event_type == "done":
return AssistantDoneEvent(
reason=cast(
Literal["stop", "length", "toolUse"],
_require_literal(payload.get("reason"), _ASSISTANT_DONE_REASON_VALUES, field="assistantMessageEvent.reason"),
),
message=_parse_assistant_message(
_clone_json_object(payload.get("message"), field="assistantMessageEvent.message"),
field="assistantMessageEvent.message",
),
)
return AssistantErrorEvent(
reason=cast(
Literal["aborted", "error"],
_require_literal(payload.get("reason"), _ASSISTANT_ERROR_REASON_VALUES, field="assistantMessageEvent.reason"),
),
error=_parse_assistant_message(
_clone_json_object(payload.get("error"), field="assistantMessageEvent.error"),
field="assistantMessageEvent.error",
),
)
class TextContent(TypedDict, total=False):
@@ -742,18 +1014,18 @@ def assistant_text_with_thinking(message: AgentMessage) -> str | None:
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 {})
cost_payload = _optional_json_object(payload.get("cost"), field="model.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"]),
id=_require_str(payload, "id"),
name=_require_str(payload, "name"),
api=_require_str(payload, "api"),
provider=_require_str(payload, "provider"),
base_url=_require_str(payload, "baseUrl"),
reasoning=bool(payload.get("reasoning", False)),
input_modalities=tuple(str(item) for item in cast(list[Any], payload.get("input") or [])),
input_modalities=_tuple_of_strings(payload.get("input"), field="model.input") or (),
cost=ModelCost(
input=float(cost_payload.get("input", 0.0)),
output=float(cost_payload.get("output", 0.0)),
@@ -762,7 +1034,7 @@ def parse_model_info(payload: JsonObject | None) -> ModelInfo | None:
),
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,
headers=cast(dict[str, str] | None, _optional_json_object(headers_payload, field="model.headers")),
premium_multiplier=float(payload["premiumMultiplier"]) if "premiumMultiplier" in payload else None,
prefer_websockets=bool(payload["preferWebsockets"]) if "preferWebsockets" in payload else None,
context_promotion_target=(
@@ -771,40 +1043,55 @@ def parse_model_info(payload: JsonObject | None) -> ModelInfo | 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"]),
min_level=cast(
ThinkingLevel,
_require_literal(thinking_payload.get("minLevel"), _THINKING_LEVEL_VALUES, field="model.thinking.minLevel"),
),
max_level=cast(
ThinkingLevel,
_require_literal(thinking_payload.get("maxLevel"), _THINKING_LEVEL_VALUES, field="model.thinking.maxLevel"),
),
mode=_require_str(cast(JsonObject, 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,
compat=_optional_json_object(compat_payload, field="model.compat"),
)
def parse_tool_descriptor(payload: JsonObject) -> ToolDescriptor:
return ToolDescriptor(
name=str(payload["name"]),
description=str(payload["description"]),
parameters=cast(JsonValue, payload.get("parameters")),
name=_require_str(payload, "name"),
description=_require_str(payload, "description"),
parameters=_clone_json_value(payload.get("parameters"), field="tool.parameters"),
)
def parse_todo_item(payload: JsonObject) -> TodoItem:
return TodoItem(
id=str(payload.get("id", "")),
content=str(payload.get("content", "")),
status=cast(TodoStatus, payload.get("status", "pending")),
notes=str(payload["notes"]) if payload.get("notes") is not None else None,
details=str(payload["details"]) if payload.get("details") is not None else None,
content=_require_str(payload, "content"),
status=cast(
TodoStatus,
_require_literal(payload.get("status", "pending"), _TODO_STATUS_VALUES, field="todo.status"),
),
notes=_optional_str(payload, "notes"),
details=_optional_str(payload, "details"),
)
def parse_todo_phase(payload: JsonObject) -> TodoPhase:
tasks = tuple(parse_todo_item(cast(JsonObject, item)) for item in cast(list[Any], payload.get("tasks") or []))
raw_tasks = payload.get("tasks")
if raw_tasks is None:
tasks = ()
else:
if not isinstance(raw_tasks, list):
raise ValueError("tasks must be a list")
tasks = tuple(parse_todo_item(_clone_json_object(item, field="tasks[]")) for item in raw_tasks)
return TodoPhase(
id=str(payload.get("id", "")),
name=str(payload.get("name", "")),
name=_require_str(payload, "name"),
tasks=tasks,
)
@@ -817,24 +1104,36 @@ def parse_todo_phases(payload: JsonValue | None) -> tuple[TodoPhase, ...]:
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 [])
parse_tool_descriptor(_clone_json_object(item, field="dumpTools[]")) for item in cast(list[Any], payload.get("dumpTools") or [])
)
return SessionState(
model=parse_model_info(cast(JsonObject | None, payload.get("model"))),
thinking_level=cast(ThinkingLevel | None, payload.get("thinkingLevel")),
thinking_level=cast(
ThinkingLevel | None,
_optional_literal(payload.get("thinkingLevel"), _THINKING_LEVEL_VALUES, field="thinkingLevel"),
),
is_streaming=bool(payload.get("isStreaming", False)),
is_compacting=bool(payload.get("isCompacting", False)),
steering_mode=cast(SteeringMode, 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,
steering_mode=cast(
SteeringMode,
_require_literal(payload.get("steeringMode", "one-at-a-time"), _STEERING_MODE_VALUES, field="steeringMode"),
),
follow_up_mode=cast(
SteeringMode,
_require_literal(payload.get("followUpMode", "one-at-a-time"), _STEERING_MODE_VALUES, field="followUpMode"),
),
interrupt_mode=cast(
InterruptMode,
_require_literal(payload.get("interruptMode", "immediate"), _INTERRUPT_MODE_VALUES, field="interruptMode"),
),
session_file=_optional_str(payload, "sessionFile"),
session_id=_require_str(payload, "sessionId"),
session_name=_optional_str(payload, "sessionName"),
auto_compaction_enabled=bool(payload.get("autoCompactionEnabled", False)),
message_count=int(payload.get("messageCount", 0)),
queued_message_count=int(payload.get("queuedMessageCount", 0)),
todo_phases=parse_todo_phases(cast(JsonValue | None, payload.get("todoPhases"))),
system_prompt=str(payload["systemPrompt"]) if payload.get("systemPrompt") is not None else None,
system_prompt=_optional_str(payload, "systemPrompt"),
dump_tools=dump_tools,
)
@@ -842,25 +1141,25 @@ def parse_session_state(payload: JsonObject) -> SessionState:
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,
exit_code=_optional_int(payload, "exitCode"),
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,
artifact_id=_optional_str(payload, "artifactId"),
)
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,
short_summary=_optional_str(payload, "shortSummary"),
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")),
details=_clone_json_value(payload.get("details"), field="compaction.details") if "details" in payload else None,
preserve_data=_optional_json_object(payload.get("preserveData"), field="compaction.preserveData"),
)
@@ -893,16 +1192,22 @@ def parse_branch_result(payload: JsonObject | None) -> BranchResult:
def parse_branch_messages(payload: JsonObject | None) -> tuple[BranchMessage, ...]:
messages = cast(list[Any], (payload or {}).get("messages") or [])
messages = (payload or {}).get("messages") or []
if not isinstance(messages, list):
raise ValueError("messages must be a list")
return tuple(
BranchMessage(entry_id=str(item.get("entryId", "")), text=str(item.get("text", ""))) for item in messages
BranchMessage(
entry_id=str(_clone_json_object(item, field="messages[]").get("entryId", "")),
text=str(_clone_json_object(item, field="messages[]").get("text", "")),
)
for item in messages
)
def parse_session_stats(payload: JsonObject) -> SessionStats:
tokens_payload = cast(dict[str, Any], payload.get("tokens") or {})
tokens_payload = _optional_json_object(payload.get("tokens"), field="sessionStats.tokens") or {}
return SessionStats(
session_file=str(payload["sessionFile"]) if payload.get("sessionFile") is not None else None,
session_file=_optional_str(payload, "sessionFile"),
session_id=str(payload.get("sessionId", "")),
user_messages=int(payload.get("userMessages", 0)),
assistant_messages=int(payload.get("assistantMessages", 0)),
@@ -923,31 +1228,44 @@ def parse_session_stats(payload: JsonObject) -> SessionStats:
def parse_extension_ui_request(payload: JsonObject) -> ExtensionUiRequest:
return ExtensionUiRequest(
id=str(payload["id"]),
method=cast(ExtensionUiMethod, 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,
id=_require_str(payload, "id"),
method=cast(
ExtensionUiMethod,
_require_literal(payload.get("method"), _EXTENSION_UI_METHOD_VALUES, field="extension_ui_request.method"),
),
title=_optional_str(payload, "title"),
options=_tuple_of_strings(payload.get("options"), field="extension_ui_request.options"),
message=_optional_str(payload, "message"),
placeholder=_optional_str(payload, "placeholder"),
prefill=_optional_str(payload, "prefill"),
timeout=_optional_int(payload, "timeout"),
prompt_style=_optional_bool(payload, "promptStyle"),
target_id=_optional_str(payload, "targetId"),
notify_type=cast(
NotifyType | None,
_optional_literal(payload.get("notifyType"), _NOTIFY_TYPE_VALUES, field="extension_ui_request.notifyType"),
),
status_key=_optional_str(payload, "statusKey"),
status_text=_optional_str(payload, "statusText"),
widget_key=_optional_str(payload, "widgetKey"),
widget_lines=_tuple_of_strings(payload.get("widgetLines"), field="extension_ui_request.widgetLines"),
widget_placement=cast(
WidgetPlacement | None,
_optional_literal(
payload.get("widgetPlacement"),
_WIDGET_PLACEMENT_VALUES,
field="extension_ui_request.widgetPlacement",
),
),
text=_optional_str(payload, "text"),
)
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", "")),
extension_path=_require_str(payload, "extensionPath"),
event=_require_str(payload, "event"),
error=_require_str(payload, "error"),
)
@@ -962,60 +1280,96 @@ def parse_notification(payload: JsonObject) -> RpcNotification:
if event_type == "agent_start":
return AgentStartEvent()
if event_type == "agent_end":
return AgentEndEvent(messages=tuple(cast(list[AgentMessage], payload.get("messages") or [])))
return AgentEndEvent(messages=parse_agent_messages(cast(JsonValue | None, payload.get("messages"))))
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 [])),
message=_parse_agent_message(
_clone_json_object(payload.get("message"), field="turn_end.message"),
field="turn_end.message",
),
tool_results=tuple(
_parse_tool_result_message(_clone_json_object(item, field="turn_end.toolResults[]"), field="turn_end.toolResults[]")
for item in cast(list[Any], payload.get("toolResults") or [])
),
)
if event_type == "message_start":
return MessageStartEvent(message=cast(AgentMessage, payload["message"]))
return MessageStartEvent(
message=_parse_agent_message(
_clone_json_object(payload.get("message"), field="message_start.message"),
field="message_start.message",
)
)
if event_type == "message_update":
return MessageUpdateEvent(
message=cast(AgentMessage, payload["message"]),
assistant_message_event=cast(AssistantMessageEvent, payload["assistantMessageEvent"]),
message=_parse_agent_message(
_clone_json_object(payload.get("message"), field="message_update.message"),
field="message_update.message",
),
assistant_message_event=parse_assistant_message_event(
_clone_json_object(payload.get("assistantMessageEvent"), field="message_update.assistantMessageEvent")
),
)
if event_type == "message_end":
return MessageEndEvent(message=cast(AgentMessage, payload["message"]))
return MessageEndEvent(
message=_parse_agent_message(
_clone_json_object(payload.get("message"), field="message_end.message"),
field="message_end.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,
args=_clone_json_value(payload.get("args"), field="tool_execution_start.args") if "args" in payload else None,
intent=_optional_str(payload, "intent"),
)
if event_type == "tool_execution_update":
return ToolExecutionUpdateEvent(
tool_call_id=str(payload.get("toolCallId", "")),
tool_name=str(payload.get("toolName", "")),
args=cast(JsonValue, payload.get("args")),
partial_result=cast(JsonValue, payload.get("partialResult")),
args=_clone_json_value(payload.get("args"), field="tool_execution_update.args") if "args" in payload else None,
partial_result=(
_clone_json_value(payload.get("partialResult"), field="tool_execution_update.partialResult")
if "partialResult" in payload
else None
),
)
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,
result=_clone_json_value(payload.get("result"), field="tool_execution_end.result") if "result" in payload else None,
is_error=_optional_bool(payload, "isError"),
)
if event_type == "auto_compaction_start":
return AutoCompactionStartEvent(
reason=cast(Literal["threshold", "overflow", "idle"], payload.get("reason", "threshold")),
action=cast(Literal["context-full", "handoff"], payload.get("action", "context-full")),
reason=cast(
Literal["threshold", "overflow", "idle"],
_require_literal(payload.get("reason", "threshold"), _AUTO_COMPACTION_REASON_VALUES, field="auto_compaction_start.reason"),
),
action=cast(
Literal["context-full", "handoff"],
_require_literal(payload.get("action", "context-full"), _AUTO_COMPACTION_ACTION_VALUES, field="auto_compaction_start.action"),
),
)
if event_type == "auto_compaction_end":
result_payload = payload.get("result")
return AutoCompactionEndEvent(
action=cast(Literal["context-full", "handoff"], payload.get("action", "context-full")),
action=cast(
Literal["context-full", "handoff"],
_require_literal(payload.get("action", "context-full"), _AUTO_COMPACTION_ACTION_VALUES, field="auto_compaction_end.action"),
),
result=(
parse_compaction_result(cast(JsonObject, result_payload)) if isinstance(result_payload, dict) else None
parse_compaction_result(_clone_json_object(result_payload, field="auto_compaction_end.result"))
if result_payload is not None
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,
error_message=_optional_str(payload, "errorMessage"),
skipped=_optional_bool(payload, "skipped"),
)
if event_type == "auto_retry_start":
return AutoRetryStartEvent(
@@ -1028,7 +1382,7 @@ def parse_notification(payload: JsonObject) -> RpcNotification:
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,
final_error=_optional_str(payload, "finalError"),
)
if event_type == "retry_fallback_applied":
return RetryFallbackAppliedEvent(
@@ -1039,13 +1393,16 @@ def parse_notification(payload: JsonObject) -> RpcNotification:
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 [])))
return TtsrTriggeredEvent(rules=_clone_json_objects(payload.get("rules"), field="ttsr_triggered.rules"))
if event_type == "todo_reminder":
return TodoReminderEvent(
todos=tuple(parse_todo_item(cast(JsonObject, item)) for item in cast(list[Any], payload.get("todos") or [])),
todos=tuple(
parse_todo_item(_clone_json_object(item, field="todo_reminder.todos[]"))
for item in cast(list[Any], 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))
return UnknownNotification(payload=_clone_json_object(payload, field="notification"))
+632 -124
View File
@@ -2,15 +2,18 @@ from __future__ import annotations
import sys
import textwrap
import threading
import time
import unittest
from omp_rpc import RpcClient, RpcCommandError, RpcError
from omp_rpc import RpcClient, RpcCommandError, RpcConcurrencyError, RpcError, host_tool
FAKE_SERVER = textwrap.dedent(
"""
import json
import sys
import time
def usage():
return {
@@ -28,20 +31,195 @@ FAKE_SERVER = textwrap.dedent(
},
}
def model_info(model_id: str, provider: str = "anthropic"):
return {
"id": model_id,
"name": f"Model {model_id}",
"api": "anthropic-messages",
"provider": provider,
"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,
}
def assistant_message(text: str):
return {
"role": "assistant",
"content": [{"type": "text", "text": text}],
"api": "anthropic-messages",
"provider": "anthropic",
"model": "claude-sonnet-4-5",
"provider": model_provider,
"model": model_id,
"usage": usage(),
"stopReason": "stop",
"timestamp": 1,
}
registered_host_tools = []
def current_state():
return {
"model": model_info(model_id, model_provider),
"thinkingLevel": thinking_level,
"isStreaming": False,
"isCompacting": False,
"steeringMode": steering_mode,
"followUpMode": follow_up_mode,
"interruptMode": interrupt_mode,
"sessionId": "fake-session",
"sessionName": session_name,
"autoCompactionEnabled": auto_compaction_enabled,
"messageCount": len(messages),
"queuedMessageCount": 0,
"todoPhases": todo_phases,
"dumpTools": [{"name": "read", "description": "Read files", "parameters": {"type": "object"}}] + registered_host_tools,
}
def emit_prompt_turn(text: str, delay: float = 0.0, include_extra_events: bool = False):
global last_assistant_text, messages
print(json.dumps({"type": "agent_start"}), flush=True)
print(json.dumps({"type": "turn_start"}), flush=True)
partial = assistant_message("")
print(json.dumps({"type": "message_start", "message": partial}), flush=True)
print(
json.dumps(
{
"type": "message_update",
"message": partial,
"assistantMessageEvent": {
"type": "text_delta",
"contentIndex": 0,
"delta": text,
"partial": partial,
},
}
),
flush=True,
)
if delay:
time.sleep(delay)
if include_extra_events:
print(
json.dumps(
{
"type": "tool_execution_start",
"toolCallId": "tool-1",
"toolName": "read",
"args": {"path": "README.md"},
"intent": "Inspect docs",
}
),
flush=True,
)
print(
json.dumps(
{
"type": "tool_execution_update",
"toolCallId": "tool-1",
"toolName": "read",
"args": {"path": "README.md"},
"partialResult": {"bytes": 12},
}
),
flush=True,
)
print(
json.dumps(
{
"type": "tool_execution_end",
"toolCallId": "tool-1",
"toolName": "read",
"result": {"text": "docs"},
"isError": False,
}
),
flush=True,
)
print(json.dumps({"type": "auto_compaction_start", "reason": "threshold", "action": "context-full"}), flush=True)
print(
json.dumps(
{
"type": "auto_compaction_end",
"action": "context-full",
"result": {
"summary": "trimmed",
"shortSummary": "trimmed",
"firstKeptEntryId": "entry-1",
"tokensBefore": 123,
},
"aborted": False,
"willRetry": False,
}
),
flush=True,
)
print(
json.dumps(
{
"type": "auto_retry_start",
"attempt": 1,
"maxAttempts": 3,
"delayMs": 25,
"errorMessage": "retrying",
}
),
flush=True,
)
print(json.dumps({"type": "auto_retry_end", "success": True, "attempt": 1}), flush=True)
print(json.dumps({"type": "retry_fallback_applied", "from": "a", "to": "b", "role": "primary"}), flush=True)
print(json.dumps({"type": "retry_fallback_succeeded", "model": "b", "role": "primary"}), flush=True)
print(json.dumps({"type": "ttsr_triggered", "rules": [{"id": "rule-1"}]}), flush=True)
print(
json.dumps(
{
"type": "todo_reminder",
"attempt": 1,
"maxAttempts": 2,
"todos": [{"id": "task-1", "content": "Map tools", "status": "pending"}],
}
),
flush=True,
)
print(json.dumps({"type": "todo_auto_clear"}), flush=True)
assistant = assistant_message(text)
print(json.dumps({"type": "message_end", "message": assistant}), flush=True)
print(json.dumps({"type": "turn_end", "message": assistant, "toolResults": []}), flush=True)
print(json.dumps({"type": "agent_end", "messages": [assistant]}), flush=True)
last_assistant_text = text
messages = [assistant]
def respond(request_id, command, data=None, success=True, error=None):
payload = {"id": request_id, "type": "response", "command": command, "success": success}
if success and data is not None:
payload["data"] = data
if not success:
payload["error"] = error
print(json.dumps(payload), flush=True)
print(json.dumps({"type": "ready"}), flush=True)
todo_phases = []
messages = []
branch_messages = [{"entryId": "entry-1", "text": "branch message"}]
model_provider = "anthropic"
model_id = "claude-sonnet-4-5"
thinking_level = "medium"
steering_mode = "one-at-a-time"
follow_up_mode = "one-at-a-time"
interrupt_mode = "immediate"
auto_compaction_enabled = True
auto_retry_enabled = True
session_name = "Scratchpad"
last_assistant_text = None
for raw_line in sys.stdin:
raw_line = raw_line.strip()
@@ -53,151 +231,185 @@ FAKE_SERVER = textwrap.dedent(
request_id = command.get("id")
if command_type == "extension_ui_response":
print(json.dumps({"type": "agent_end", "messages": [assistant_message("ui acknowledged")]}), flush=True)
emit_prompt_turn("ui acknowledged")
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,
"todoPhases": todo_phases,
},
}
),
flush=True,
respond(request_id, "get_state", current_state())
elif command_type == "set_host_tools":
registered_host_tools = command.get("tools", [])
respond(
request_id,
"set_host_tools",
{"toolNames": [tool.get("name", "") for tool in registered_host_tools]},
)
elif command_type == "set_todos":
todo_phases = command.get("phases", [])
print(
json.dumps(
{
"id": request_id,
"type": "response",
"command": "set_todos",
"success": True,
"data": {
"todoPhases": todo_phases,
},
}
),
flush=True,
respond(request_id, "set_todos", {"todoPhases": todo_phases})
elif command_type == "get_messages":
respond(request_id, "get_messages", {"messages": messages})
elif command_type == "set_host_tools":
tool_names = [tool.get("name", "") for tool in command.get("tools", [])]
respond(request_id, "set_host_tools", {"toolNames": tool_names})
elif command_type == "set_model":
model_provider = command["provider"]
model_id = command["modelId"]
respond(request_id, "set_model", model_info(model_id, model_provider))
elif command_type == "cycle_model":
model_id = "claude-sonnet-4-6" if model_id == "claude-sonnet-4-5" else "claude-sonnet-4-5"
respond(request_id, "cycle_model", {"model": model_info(model_id, model_provider), "thinkingLevel": thinking_level, "isScoped": False})
elif command_type == "get_available_models":
respond(
request_id,
"get_available_models",
{
"models": [
model_info("claude-sonnet-4-5", "anthropic"),
model_info("claude-sonnet-4-6", "anthropic"),
]
},
)
elif command_type == "set_thinking_level":
thinking_level = command["level"]
respond(request_id, "set_thinking_level", {})
elif command_type == "cycle_thinking_level":
thinking_level = "high" if thinking_level != "high" else "low"
respond(request_id, "cycle_thinking_level", {"level": thinking_level})
elif command_type == "set_steering_mode":
steering_mode = command["mode"]
respond(request_id, "set_steering_mode", {})
elif command_type == "set_follow_up_mode":
follow_up_mode = command["mode"]
respond(request_id, "set_follow_up_mode", {})
elif command_type == "set_interrupt_mode":
interrupt_mode = command["mode"]
respond(request_id, "set_interrupt_mode", {})
elif command_type == "compact":
respond(
request_id,
"compact",
{"summary": "trimmed", "shortSummary": "trimmed", "firstKeptEntryId": "entry-1", "tokensBefore": 123},
)
elif command_type == "set_auto_compaction":
auto_compaction_enabled = command["enabled"]
respond(request_id, "set_auto_compaction", {})
elif command_type == "set_auto_retry":
auto_retry_enabled = command["enabled"]
respond(request_id, "set_auto_retry", {})
elif command_type == "abort_retry":
respond(request_id, "abort_retry", {})
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,
respond(
request_id,
"bash",
{
"output": "hello\\n",
"exitCode": 0,
"cancelled": False,
"truncated": False,
"totalLines": 1,
"totalBytes": 6,
"outputLines": 1,
"outputBytes": 6,
},
)
elif command_type == "prompt":
print(
json.dumps(
{
"id": request_id,
"type": "response",
"command": "prompt",
"success": True,
}
),
flush=True,
elif command_type == "abort_bash":
respond(request_id, "abort_bash", {})
elif command_type == "get_session_stats":
respond(
request_id,
"get_session_stats",
{
"sessionFile": "/tmp/fake-session.jsonl",
"sessionId": "fake-session",
"userMessages": 1,
"assistantMessages": len(messages),
"toolCalls": 1,
"toolResults": 1,
"totalMessages": len(messages) + 1,
"tokens": {"input": 10, "output": 5, "cacheRead": 0, "cacheWrite": 0, "total": 15},
"premiumRequests": 0,
"cost": 0.0,
},
)
if command["message"] == "needs ui":
elif command_type == "export_html":
respond(request_id, "export_html", {"path": command.get("outputPath") or "/tmp/session.html"})
elif command_type == "new_session":
respond(request_id, "new_session", {"cancelled": False})
elif command_type == "switch_session":
respond(request_id, "switch_session", {"cancelled": False})
elif command_type == "branch":
branch_messages = [{"entryId": command["entryId"], "text": "branch message"}]
respond(request_id, "branch", {"text": "branch created", "cancelled": False})
elif command_type == "get_branch_messages":
respond(request_id, "get_branch_messages", {"messages": branch_messages})
elif command_type == "get_last_assistant_text":
respond(request_id, "get_last_assistant_text", {"text": last_assistant_text})
elif command_type == "set_session_name":
session_name = command["name"]
respond(request_id, "set_session_name", {})
elif command_type in {"steer", "follow_up", "abort"}:
respond(request_id, command_type, {})
elif command_type in {"prompt", "abort_and_prompt"}:
respond(request_id, command_type, {})
message = command["message"]
if message == "needs ui":
print(json.dumps({"type": "extension_ui_request", "id": "ui-1", "method": "input", "title": "Need input", "placeholder": "value"}), flush=True)
continue
if message == "needs confirm":
print(json.dumps({"type": "extension_ui_request", "id": "ui-2", "method": "confirm", "title": "Confirm", "message": "Continue?"}), flush=True)
continue
if message == "needs cancel":
print(json.dumps({"type": "extension_ui_request", "id": "ui-3", "method": "editor", "title": "Edit", "placeholder": "value"}), flush=True)
continue
if message == "needs host tool":
print(json.dumps({"type": "agent_start"}), flush=True)
print(
json.dumps(
{
"type": "extension_ui_request",
"id": "ui-1",
"method": "input",
"title": "Need input",
"placeholder": "value",
"type": "host_tool_call",
"id": "host-call-1",
"toolCallId": "toolu_host_1",
"toolName": "echo_host",
"arguments": {"message": "hello"},
}
),
flush=True,
)
continue
print(json.dumps({"type": "agent_start"}), flush=True)
print(json.dumps({"type": "turn_start"}), flush=True)
partial = assistant_message("")
if message == "notifications":
print(json.dumps({"type": "extension_error", "extensionPath": "/tmp/ext.py", "event": "run", "error": "boom"}), flush=True)
print(json.dumps({"type": "unknown_future_event", "value": 1}), flush=True)
emit_prompt_turn("pong", delay=0.3 if message == "slow" else 0.0, include_extra_events=message == "all events")
elif command_type == "host_tool_update":
print(
json.dumps(
{
"type": "message_update",
"message": partial,
"assistantMessageEvent": {
"type": "text_delta",
"contentIndex": 0,
"delta": "pong",
"partial": partial,
},
"type": "tool_execution_update",
"toolCallId": "toolu_host_1",
"toolName": "echo_host",
"args": {"message": "hello"},
"partialResult": command["partialResult"],
}
),
flush=True,
)
assistant = assistant_message("pong")
print(json.dumps({"type": "message_end", "message": assistant}), flush=True)
print(json.dumps({"type": "turn_end", "message": assistant, "toolResults": []}), flush=True)
print(json.dumps({"type": "agent_end", "messages": [assistant]}), flush=True)
elif command_type == "host_tool_result":
print(
json.dumps(
{
"type": "tool_execution_end",
"toolCallId": "toolu_host_1",
"toolName": "echo_host",
"result": command["result"],
"isError": command.get("isError", False),
}
),
flush=True,
)
print(json.dumps({"type": "agent_end", "messages": []}), flush=True)
else:
print(
json.dumps(
{
"id": request_id,
"type": "response",
"command": command_type,
"success": False,
"error": f"unsupported: {command_type}",
}
),
flush=True,
)
respond(request_id, command_type, success=False, error=f"unsupported: {command_type}")
"""
)
@@ -214,6 +426,20 @@ IDLESS_ERROR_SERVER = textwrap.dedent(
continue
command = json.loads(raw_line)
if command["type"] == "set_host_tools":
print(
json.dumps(
{
"id": command.get("id"),
"type": "response",
"command": "set_host_tools",
"success": True,
"data": {"toolNames": []},
}
),
flush=True,
)
continue
print(
json.dumps(
{
@@ -242,6 +468,20 @@ LATE_PROMPT_FAILURE_SERVER = textwrap.dedent(
command = json.loads(raw_line)
request_id = command.get("id")
if command["type"] == "set_host_tools":
print(
json.dumps(
{
"id": request_id,
"type": "response",
"command": "set_host_tools",
"success": True,
"data": {"toolNames": []},
}
),
flush=True,
)
continue
if command["type"] == "prompt":
print(
json.dumps(
@@ -292,8 +532,44 @@ STDERR_SERVER = textwrap.dedent(
sys.stderr.flush()
print(json.dumps({"type": "ready"}), flush=True)
for _ in sys.stdin:
pass
for raw_line in sys.stdin:
raw_line = raw_line.strip()
if not raw_line:
continue
command = json.loads(raw_line)
if command["type"] == "set_host_tools":
print(
json.dumps(
{
"id": command.get("id"),
"type": "response",
"command": "set_host_tools",
"success": True,
"data": {"toolNames": []},
}
),
flush=True,
)
"""
)
INVALID_JSON_SERVER = textwrap.dedent(
"""
import sys
sys.stdout.write('{"type":"ready"}\\n')
sys.stdout.flush()
sys.stdout.write('{"type":"broken"\\n')
sys.stdout.flush()
"""
)
BROKEN_STARTUP_SERVER = textwrap.dedent(
"""
import sys
sys.stdout.write('not-json\\n')
sys.stdout.flush()
"""
)
@@ -358,6 +634,38 @@ class RpcClientTests(unittest.TestCase):
self.assertEqual(turn.require_assistant_text(), "pong")
self.assertGreaterEqual(len(turn.events), 3)
def test_custom_tools_are_registered_and_executed_via_rpc(self) -> None:
def echo_host(args: dict[str, str], context) -> str:
context.send_update(f"working:{args['message']}")
return f"host:{args['message']}"
with self.make_client(
custom_tools=(
host_tool(
name="echo_host",
description="Echo from the Python host process",
parameters={
"type": "object",
"properties": {"message": {"type": "string"}},
"required": ["message"],
"additionalProperties": False,
},
execute=echo_host,
),
)
) as client:
state = client.get_state()
self.assertEqual(state.dump_tools[-1].name, "echo_host")
turn = client.prompt_and_wait("needs host tool", timeout=2.0)
update_events = [event for event in turn.events if getattr(event, "type", None) == "tool_execution_update"]
end_events = [event for event in turn.events if getattr(event, "type", None) == "tool_execution_end"]
self.assertEqual(len(update_events), 1)
self.assertEqual(update_events[0].partial_result["content"][0]["text"], "working:hello")
self.assertEqual(len(end_events), 1)
self.assertEqual(end_events[0].result["content"][0]["text"], "host:hello")
def test_extension_ui_round_trip(self) -> None:
with self.make_client() as client:
client.prompt("needs ui")
@@ -411,6 +719,198 @@ class RpcClientTests(unittest.TestCase):
state = client.get_state()
self.assertEqual(state.todo_phases[0].tasks[1].content, "Exercise edits")
def test_model_mode_and_session_commands(self) -> None:
with self.make_client() as client:
model = client.set_model("anthropic", "claude-sonnet-4-6")
self.assertEqual(model.id, "claude-sonnet-4-6")
cycled = client.cycle_model()
self.assertIsNotNone(cycled)
self.assertEqual(cycled.model.id, "claude-sonnet-4-5")
available = client.get_available_models()
self.assertEqual([item.id for item in available], ["claude-sonnet-4-5", "claude-sonnet-4-6"])
client.set_thinking_level("high")
self.assertEqual(client.get_state().thinking_level, "high")
cycled_level = client.cycle_thinking_level()
self.assertIsNotNone(cycled_level)
self.assertEqual(cycled_level.level, "low")
client.set_steering_mode("all")
client.set_follow_up_mode("all")
client.set_interrupt_mode("wait")
client.set_auto_compaction(False)
client.set_auto_retry(False)
client.set_session_name("Renamed")
state = client.get_state()
self.assertEqual(state.steering_mode, "all")
self.assertEqual(state.follow_up_mode, "all")
self.assertEqual(state.interrupt_mode, "wait")
self.assertFalse(state.auto_compaction_enabled)
self.assertEqual(state.session_name, "Renamed")
compacted = client.compact()
self.assertEqual(compacted.summary, "trimmed")
stats = client.get_session_stats()
self.assertEqual(stats.session_id, "fake-session")
self.assertEqual(stats.tokens.total, 15)
exported = client.export_html("/tmp/custom.html")
self.assertEqual(str(exported), "/tmp/custom.html")
new_session = client.new_session()
switched = client.switch_session("/tmp/session.jsonl")
self.assertFalse(new_session.cancelled)
self.assertFalse(switched.cancelled)
branch = client.branch("entry-9")
self.assertEqual(branch.text, "branch created")
branch_messages = client.get_branch_messages()
self.assertEqual(branch_messages[0].entry_id, "entry-9")
def test_message_and_control_commands(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.assertEqual(client.get_last_assistant_text(), "pong")
messages = client.get_messages()
self.assertEqual(len(messages), 1)
self.assertEqual(messages[0]["role"], "assistant")
client.clear_todos()
self.assertEqual(client.get_todos(), ())
client.steer("nudge")
client.follow_up("later")
client.abort()
client.abort_retry()
client.abort_bash()
client.abort_and_prompt("say hello")
client.wait_for_idle(timeout=2.0)
self.assertEqual(client.get_last_assistant_text(), "pong")
def test_collect_events_returns_turn_events(self) -> None:
with self.make_client() as client:
client.prompt("slow")
events = client.collect_events(timeout=2.0)
self.assertGreaterEqual(len(events), 1)
self.assertEqual(events[-1].type, "agent_end")
def test_all_typed_event_listeners_receive_eventful_prompt(self) -> None:
seen: list[str] = []
with self.make_client() as client:
client.on_event(lambda event: seen.append(f"event:{event.type}"))
client.on_agent_start(lambda event: seen.append(event.type))
client.on_turn_end(lambda event: seen.append(event.type))
client.on_message_start(lambda event: seen.append(event.type))
client.on_message_end(lambda event: seen.append(event.type))
client.on_tool_execution_start(lambda event: seen.append(event.type))
client.on_tool_execution_update(lambda event: seen.append(event.type))
client.on_tool_execution_end(lambda event: seen.append(event.type))
client.on_auto_compaction_start(lambda event: seen.append(event.type))
client.on_auto_compaction_end(lambda event: seen.append(event.type))
client.on_auto_retry_start(lambda event: seen.append(event.type))
client.on_auto_retry_end(lambda event: seen.append(event.type))
client.on_retry_fallback_applied(lambda event: seen.append(event.type))
client.on_retry_fallback_succeeded(lambda event: seen.append(event.type))
client.on_ttsr_triggered(lambda event: seen.append(event.type))
client.on_todo_reminder(lambda event: seen.append(event.type))
client.on_todo_auto_clear(lambda event: seen.append(event.type))
turn = client.prompt_and_wait("all events", timeout=2.0)
self.assertEqual(turn.require_assistant_text(), "pong")
for expected in [
"agent_start",
"message_start",
"message_end",
"turn_end",
"tool_execution_start",
"tool_execution_update",
"tool_execution_end",
"auto_compaction_start",
"auto_compaction_end",
"auto_retry_start",
"auto_retry_end",
"retry_fallback_applied",
"retry_fallback_succeeded",
"ttsr_triggered",
"todo_reminder",
"todo_auto_clear",
]:
self.assertIn(expected, seen)
def test_extension_and_unknown_notification_listeners(self) -> None:
seen_extension_errors: list[str] = []
seen_unknown: list[str] = []
with self.make_client() as client:
client.on_extension_error(lambda event: seen_extension_errors.append(event.error))
client.on_unknown_notification(lambda event: seen_unknown.append(str(event.payload.get("type"))))
client.prompt_and_wait("notifications", timeout=2.0)
self.assertEqual(seen_extension_errors, ["boom"])
self.assertEqual(seen_unknown, ["unknown_future_event"])
def test_ui_confirmation_and_cancel_round_trip(self) -> None:
with self.make_client() as client:
client.prompt("needs confirm")
confirm_request = client.next_ui_request(timeout=2.0)
self.assertEqual(confirm_request.method, "confirm")
client.send_ui_confirmation(confirm_request.id, True)
client.wait_for_idle(timeout=2.0)
client.prompt("needs cancel")
editor_request = client.next_ui_request(timeout=2.0)
self.assertEqual(editor_request.method, "editor")
client.cancel_ui_request(editor_request.id)
client.wait_for_idle(timeout=2.0)
def test_prompt_lifecycle_collectors_are_single_flight(self) -> None:
results: list[str] = []
errors: list[BaseException] = []
with self.make_client() as client:
def run_prompt() -> None:
try:
results.append(client.prompt_and_wait("slow", timeout=2.0).require_assistant_text())
except BaseException as exc: # pragma: no cover - defensive thread capture
errors.append(exc)
thread = threading.Thread(target=run_prompt)
thread.start()
deadline = time.time() + 1.0
while client._prompt_lifecycle.active_operation != "prompt_and_wait" and time.time() < deadline:
time.sleep(0.01)
self.assertEqual(client._prompt_lifecycle.active_operation, "prompt_and_wait")
with self.assertRaises(RpcConcurrencyError):
client.collect_events(timeout=1.0)
thread.join(timeout=2.0)
self.assertFalse(thread.is_alive())
self.assertEqual(errors, [])
self.assertEqual(results, ["pong"])
def test_listener_mutation_does_not_change_retained_turn(self) -> None:
with self.make_client() as client:
client.on_message_end(lambda event: event.message["content"].__setitem__(0, {"type": "text", "text": "mutated"}))
turn = client.prompt_and_wait("say hello", timeout=2.0)
messages = client.get_messages()
self.assertEqual(turn.require_assistant_text(), "pong")
self.assertEqual(messages[0]["content"][0]["text"], "pong")
def test_id_less_error_responses_are_correlated(self) -> None:
with self.make_client(server=IDLESS_ERROR_SERVER) as client:
with self.assertRaises(RpcCommandError) as ctx:
@@ -470,6 +970,14 @@ class RpcClientTests(unittest.TestCase):
self.assertEqual(client.stderr, "second\n")
def test_broken_startup_frame_is_reported(self) -> None:
client = self.make_client(server=BROKEN_STARTUP_SERVER)
with self.assertRaises(RpcError) as ctx:
client.start()
self.assertIn("Frame: 'not-json'", str(ctx.exception))
def test_event_history_limit_reports_overflow(self) -> None:
with self.make_client(max_event_history=2) as client:
with self.assertRaises(RpcError) as ctx:
+110
View File
@@ -170,6 +170,116 @@ class ProtocolParsingTests(unittest.TestCase):
self.assertEqual(assistant_text(message), "visible")
self.assertEqual(assistant_text_with_thinking(message), "internalvisible")
def test_parse_session_state_rejects_invalid_thinking_level(self) -> None:
with self.assertRaises(ValueError):
parse_session_state(
{
"sessionId": "session-123",
"thinkingLevel": "extreme",
"steeringMode": "one-at-a-time",
"followUpMode": "one-at-a-time",
"interruptMode": "immediate",
}
)
def test_parse_extension_ui_request_rejects_invalid_method(self) -> None:
with self.assertRaises(ValueError):
parse_notification({"type": "extension_ui_request", "id": "ui-1", "method": "launch"})
def test_parse_message_update_rejects_invalid_assistant_done_reason(self) -> None:
with self.assertRaises(ValueError):
parse_notification(
{
"type": "message_update",
"message": {
"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,
},
"assistantMessageEvent": {
"type": "done",
"reason": "error",
"message": {
"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,
},
},
}
)
def test_parse_notification_deep_clones_nested_messages(self) -> None:
payload = {
"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,
}
],
}
notification = parse_notification(payload)
payload["messages"][0]["content"][0]["text"] = "mutated"
self.assertIsInstance(notification, AgentEndEvent)
self.assertEqual(notification.messages[0]["content"][0]["text"], "hello")
if __name__ == "__main__":
unittest.main()