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:
@@ -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,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",
|
||||
]
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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
@@ -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:
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user