chore: reformat

This commit is contained in:
can1357
2026-06-02 08:43:23 +02:00
parent 0bc9bc25b4
commit 2ecb5fd9fa
8 changed files with 716 additions and 198 deletions
+7 -1
View File
@@ -16,7 +16,13 @@ from .client import (
ProtocolErrorListener,
UiRequestListener,
)
from .host_tools import HostTool, HostToolContext, HostToolResultPayload, HostToolResultValue, host_tool
from .host_tools import (
HostTool,
HostToolContext,
HostToolResultPayload,
HostToolResultValue,
host_tool,
)
from .host_uris import (
HostUri,
HostUriContentType,
+302 -82
View File
@@ -322,8 +322,12 @@ class RpcClient:
self._extra_args = tuple(extra_args)
self._startup_timeout = startup_timeout
self._request_timeout = request_timeout
self._max_event_history = self._validate_history_limit("max_event_history", max_event_history)
self._max_stderr_chunks = self._validate_history_limit("max_stderr_chunks", max_stderr_chunks)
self._max_event_history = self._validate_history_limit(
"max_event_history", max_event_history
)
self._max_stderr_chunks = self._validate_history_limit(
"max_stderr_chunks", max_stderr_chunks
)
self._process: subprocess.Popen[str] | None = None
self._stdout_thread: threading.Thread | None = None
@@ -337,7 +341,9 @@ class RpcClient:
self._pending_host_uri_requests: dict[str, _PendingHostUriRequest] = {}
self._request_id = 0
self._events = _BoundedHistory[JsonObject](self._max_event_history)
self._async_errors = _BoundedHistory[BaseException](_DEFAULT_ERROR_HISTORY_LIMIT)
self._async_errors = _BoundedHistory[BaseException](
_DEFAULT_ERROR_HISTORY_LIMIT
)
self._scheduled_agent_runs = 0
self._completed_agent_runs = 0
self._last_schedule_async_error_index = 0
@@ -346,8 +352,12 @@ class RpcClient:
self._closed_error: BaseException | None = None
self._stopping = False
self._ready_received = False
self._protocol_errors = _BoundedHistory[RpcProtocolError](_DEFAULT_ERROR_HISTORY_LIMIT)
self._listener_errors = _BoundedHistory[ListenerErrorEvent](_DEFAULT_ERROR_HISTORY_LIMIT)
self._protocol_errors = _BoundedHistory[RpcProtocolError](
_DEFAULT_ERROR_HISTORY_LIMIT
)
self._listener_errors = _BoundedHistory[ListenerErrorEvent](
_DEFAULT_ERROR_HISTORY_LIMIT
)
self._prompt_lifecycle = _PromptLifecycleCoordinator()
self._notification_listeners: list[NotificationListener] = []
@@ -422,15 +432,21 @@ class RpcClient:
)
self._process = process
self._stdout_thread = threading.Thread(target=self._read_stdout_loop, name="omp-rpc-stdout", daemon=True)
self._stderr_thread = threading.Thread(target=self._read_stderr_loop, name="omp-rpc-stderr", daemon=True)
self._stdout_thread = threading.Thread(
target=self._read_stdout_loop, name="omp-rpc-stdout", daemon=True
)
self._stderr_thread = threading.Thread(
target=self._read_stderr_loop, name="omp-rpc-stderr", daemon=True
)
self._stdout_thread.start()
self._stderr_thread.start()
if not self._ready.wait(self._startup_timeout):
stderr = self.stderr
self.stop()
raise RpcTimeoutError(f"Timed out waiting for RPC ready signal. Stderr: {stderr}")
raise RpcTimeoutError(
f"Timed out waiting for RPC ready signal. Stderr: {stderr}"
)
if not self._ready_received:
error = self._closed_error
@@ -439,8 +455,12 @@ class RpcClient:
if isinstance(error, RpcError):
raise error
if error is not None:
raise RpcProcessExitError(f"RPC process stopped before ready: {error}. Stderr: {stderr}") from error
raise RpcTimeoutError(f"Timed out waiting for RPC ready signal. Stderr: {stderr}")
raise RpcProcessExitError(
f"RPC process stopped before ready: {error}. Stderr: {stderr}"
) from error
raise RpcTimeoutError(
f"Timed out waiting for RPC ready signal. Stderr: {stderr}"
)
if self._custom_tools:
self.set_custom_tools(self._custom_tools)
@@ -536,31 +556,47 @@ class RpcClient:
def on_message_end(self, listener: MessageEndListener) -> Callable[[], None]:
return self._add_typed_event_listener("message_end", listener)
def on_tool_execution_start(self, listener: ToolExecutionStartListener) -> Callable[[], None]:
def on_tool_execution_start(
self, listener: ToolExecutionStartListener
) -> Callable[[], None]:
return self._add_typed_event_listener("tool_execution_start", listener)
def on_tool_execution_update(self, listener: ToolExecutionUpdateListener) -> Callable[[], None]:
def on_tool_execution_update(
self, listener: ToolExecutionUpdateListener
) -> Callable[[], None]:
return self._add_typed_event_listener("tool_execution_update", listener)
def on_tool_execution_end(self, listener: ToolExecutionEndListener) -> Callable[[], None]:
def on_tool_execution_end(
self, listener: ToolExecutionEndListener
) -> Callable[[], None]:
return self._add_typed_event_listener("tool_execution_end", listener)
def on_auto_compaction_start(self, listener: AutoCompactionStartListener) -> Callable[[], None]:
def on_auto_compaction_start(
self, listener: AutoCompactionStartListener
) -> Callable[[], None]:
return self._add_typed_event_listener("auto_compaction_start", listener)
def on_auto_compaction_end(self, listener: AutoCompactionEndListener) -> Callable[[], None]:
def on_auto_compaction_end(
self, listener: AutoCompactionEndListener
) -> Callable[[], None]:
return self._add_typed_event_listener("auto_compaction_end", listener)
def on_auto_retry_start(self, listener: AutoRetryStartListener) -> Callable[[], None]:
def on_auto_retry_start(
self, listener: AutoRetryStartListener
) -> Callable[[], None]:
return self._add_typed_event_listener("auto_retry_start", listener)
def on_auto_retry_end(self, listener: AutoRetryEndListener) -> Callable[[], None]:
return self._add_typed_event_listener("auto_retry_end", listener)
def on_retry_fallback_applied(self, listener: RetryFallbackAppliedListener) -> Callable[[], None]:
def on_retry_fallback_applied(
self, listener: RetryFallbackAppliedListener
) -> Callable[[], None]:
return self._add_typed_event_listener("retry_fallback_applied", listener)
def on_retry_fallback_succeeded(self, listener: RetryFallbackSucceededListener) -> Callable[[], None]:
def on_retry_fallback_succeeded(
self, listener: RetryFallbackSucceededListener
) -> Callable[[], None]:
return self._add_typed_event_listener("retry_fallback_succeeded", listener)
def on_ttsr_triggered(self, listener: TtsrTriggeredListener) -> Callable[[], None]:
@@ -576,7 +612,9 @@ class RpcClient:
self._ui_request_listeners.append(listener)
return lambda: self._remove_listener(self._ui_request_listeners, listener)
def on_extension_error(self, listener: ExtensionErrorListener) -> Callable[[], None]:
def on_extension_error(
self, listener: ExtensionErrorListener
) -> Callable[[], None]:
self._extension_error_listeners.append(listener)
return lambda: self._remove_listener(self._extension_error_listeners, listener)
@@ -588,9 +626,13 @@ class RpcClient:
self._listener_error_listeners.append(listener)
return lambda: self._remove_listener(self._listener_error_listeners, listener)
def on_unknown_notification(self, listener: UnknownNotificationListener) -> Callable[[], None]:
def on_unknown_notification(
self, listener: UnknownNotificationListener
) -> Callable[[], None]:
self._unknown_notification_listeners.append(listener)
return lambda: self._remove_listener(self._unknown_notification_listeners, listener)
return lambda: self._remove_listener(
self._unknown_notification_listeners, listener
)
def install_headless_ui(
self,
@@ -651,16 +693,26 @@ class RpcClient:
try:
return self._ui_requests.get(timeout=timeout)
except queue.Empty as exc:
raise RpcTimeoutError("Timed out waiting for an extension UI request") from exc
raise RpcTimeoutError(
"Timed out waiting for an extension UI request"
) from exc
def send_ui_value(self, request_id: str, value: str) -> None:
self._send_notification({"type": "extension_ui_response", "id": request_id, "value": value})
self._send_notification(
{"type": "extension_ui_response", "id": request_id, "value": value}
)
def send_ui_confirmation(self, request_id: str, confirmed: bool) -> None:
self._send_notification({"type": "extension_ui_response", "id": request_id, "confirmed": confirmed})
self._send_notification(
{"type": "extension_ui_response", "id": request_id, "confirmed": confirmed}
)
def cancel_ui_request(self, request_id: str, *, timed_out: bool = False) -> None:
payload: JsonObject = {"type": "extension_ui_response", "id": request_id, "cancelled": True}
payload: JsonObject = {
"type": "extension_ui_response",
"id": request_id,
"cancelled": True,
}
if timed_out:
payload["timedOut"] = True
self._send_notification(payload)
@@ -724,14 +776,21 @@ class RpcClient:
return parse_session_stats(payload)
def export_html(self, output_path: str | Path | None = None) -> Path:
payload = self._request("export_html", outputPath=str(output_path) if output_path is not None else None)
payload = self._request(
"export_html",
outputPath=str(output_path) if output_path is not None else None,
)
return Path(str(payload["path"]))
def new_session(self, parent_session: str | None = None) -> CancellationResult:
return parse_cancellation_result(self._request("new_session", parentSession=parent_session))
return parse_cancellation_result(
self._request("new_session", parentSession=parent_session)
)
def switch_session(self, session_path: str | Path) -> CancellationResult:
return parse_cancellation_result(self._request("switch_session", sessionPath=str(session_path)))
return parse_cancellation_result(
self._request("switch_session", sessionPath=str(session_path))
)
def branch(self, entry_id: str) -> BranchResult:
return parse_branch_result(self._request("branch", entryId=entry_id))
@@ -750,7 +809,9 @@ class RpcClient:
def get_todos(self) -> tuple[TodoPhase, ...]:
return self.get_state().todo_phases
def set_todos(self, todos: Sequence[TodoSeed | TodoPhaseSeed]) -> tuple[TodoPhase, ...]:
def set_todos(
self, todos: Sequence[TodoSeed | TodoPhaseSeed]
) -> tuple[TodoPhase, ...]:
phases = self._normalize_todo_phases(todos)
payload = self._request("set_todos", phases=cast(JsonValue, phases))
return parse_todo_phases(payload.get("todoPhases"))
@@ -795,7 +856,11 @@ class RpcClient:
schemes_payload: list[JsonObject] = []
for uri in self._host_uris:
entry: JsonObject = {"scheme": uri.scheme, "writable": uri.writable, "immutable": uri.immutable}
entry: JsonObject = {
"scheme": uri.scheme,
"writable": uri.writable,
"immutable": uri.immutable,
}
if uri.description is not None:
entry["description"] = uri.description
schemes_payload.append(entry)
@@ -824,17 +889,35 @@ class RpcClient:
)
self._mark_agent_run_scheduled()
def steer(self, message: str, *, images: Sequence[ImageContent] | None = None) -> None:
self._request("steer", message=message, images=list(images) if images is not None else None)
def steer(
self, message: str, *, images: Sequence[ImageContent] | None = None
) -> None:
self._request(
"steer",
message=message,
images=list(images) if images is not None else None,
)
def follow_up(self, message: str, *, images: Sequence[ImageContent] | None = None) -> None:
self._request("follow_up", message=message, images=list(images) if images is not None else None)
def follow_up(
self, message: str, *, images: Sequence[ImageContent] | None = None
) -> None:
self._request(
"follow_up",
message=message,
images=list(images) if images is not None else None,
)
def abort(self) -> None:
self._request("abort")
def abort_and_prompt(self, message: str, *, images: Sequence[ImageContent] | None = None) -> None:
self._request("abort_and_prompt", message=message, images=list(images) if images is not None else None)
def abort_and_prompt(
self, message: str, *, images: Sequence[ImageContent] | None = None
) -> None:
self._request(
"abort_and_prompt",
message=message,
images=list(images) if images is not None else None,
)
self._mark_agent_run_scheduled()
def prompt_and_wait(
@@ -851,7 +934,9 @@ class RpcClient:
start_index = self._current_event_index()
start_async_error_index = self._current_async_error_index()
self.prompt(message, images=images, streaming_behavior=streaming_behavior)
events = self._wait_for_agent_end(start_index, start_async_error_index, timeout=timeout)
events = self._wait_for_agent_end(
start_index, start_async_error_index, timeout=timeout
)
return self._build_prompt_turn(events)
finally:
self._prompt_lifecycle.release(operation)
@@ -865,7 +950,9 @@ class RpcClient:
return
start_index = self._current_event_index()
start_async_error_index = self._current_async_error_index()
self._wait_for_agent_end(start_index, start_async_error_index, timeout=timeout)
self._wait_for_agent_end(
start_index, start_async_error_index, timeout=timeout
)
finally:
self._prompt_lifecycle.release(operation)
@@ -875,7 +962,9 @@ class RpcClient:
try:
start_index = self._current_event_index()
start_async_error_index = self._current_async_error_index()
return self._wait_for_agent_end(start_index, start_async_error_index, timeout=timeout)
return self._wait_for_agent_end(
start_index, start_async_error_index, timeout=timeout
)
finally:
self._prompt_lifecycle.release(operation)
@@ -894,6 +983,7 @@ class RpcClient:
with self._event_condition:
self._scheduled_agent_runs += 1
self._last_schedule_async_error_index = self._async_errors.current_index()
def _mark_agent_run_completed(self) -> None:
with self._event_condition:
self._completed_agent_runs += 1
@@ -905,7 +995,9 @@ class RpcClient:
def _check_async_errors(self) -> None:
with self._event_condition:
errors = self._async_errors.snapshot_from(self._last_schedule_async_error_index)
errors = self._async_errors.snapshot_from(
self._last_schedule_async_error_index
)
if errors:
raise errors[0]
@@ -934,7 +1026,9 @@ class RpcClient:
events=events,
messages=final_messages,
assistant_message=assistant_message,
assistant_text=assistant_text(assistant_message) if assistant_message is not None else None,
assistant_text=assistant_text(assistant_message)
if assistant_message is not None
else None,
)
def _wait_for_agent_end(
@@ -966,13 +1060,20 @@ class RpcClient:
raise async_errors[0]
event_payloads = self._events.snapshot_from(start_index)
if any(payload.get("type") == "agent_end" for payload in event_payloads):
events = tuple(cast(RpcAgentEvent, parse_notification(payload)) for payload in event_payloads)
if any(
payload.get("type") == "agent_end" for payload in event_payloads
):
events = tuple(
cast(RpcAgentEvent, parse_notification(payload))
for payload in event_payloads
)
return events
remaining = deadline - time.monotonic()
if remaining <= 0:
raise RpcTimeoutError(f"Timed out waiting for agent_end. Stderr: {self.stderr}")
raise RpcTimeoutError(
f"Timed out waiting for agent_end. Stderr: {self.stderr}"
)
self._event_condition.wait(remaining)
def _request(self, command_type: str, **payload: JsonValue) -> JsonObject:
@@ -985,7 +1086,9 @@ class RpcClient:
response_queue: queue.Queue[JsonObject | BaseException] = queue.Queue(maxsize=1)
with self._state_lock:
self._pending[request_id] = _PendingRequest(command=command_type, response_queue=response_queue)
self._pending[request_id] = _PendingRequest(
command=command_type, response_queue=response_queue
)
try:
self._write_json(process, envelope)
@@ -999,13 +1102,18 @@ class RpcClient:
except queue.Empty as exc:
with self._state_lock:
self._pending.pop(request_id, None)
raise RpcTimeoutError(f"Timed out waiting for response to {command_type}. Stderr: {self.stderr}") from exc
raise RpcTimeoutError(
f"Timed out waiting for response to {command_type}. Stderr: {self.stderr}"
) from exc
if isinstance(response, BaseException):
raise response
if not bool(response.get("success", False)):
raise RpcCommandError(command=str(response.get("command", command_type)), error=str(response.get("error", "")))
raise RpcCommandError(
command=str(response.get("command", command_type)),
error=str(response.get("error", "")),
)
data = response.get("data")
if data is None:
@@ -1028,27 +1136,51 @@ class RpcClient:
tool_name = payload.get("toolName")
tool_call_id = payload.get("toolCallId")
raw_arguments = payload.get("arguments")
if not isinstance(request_id, str) or not isinstance(tool_name, str) or not isinstance(tool_call_id, str):
if (
not isinstance(request_id, str)
or not isinstance(tool_name, str)
or not isinstance(tool_call_id, str)
):
return
if not isinstance(raw_arguments, Mapping):
self._send_notification(
{
"type": "host_tool_result",
"id": request_id,
"result": {"content": [{"type": "text", "text": "Host tool arguments must be an object"}], "details": {}},
"result": {
"content": [
{
"type": "text",
"text": "Host tool arguments must be an object",
}
],
"details": {},
},
"isError": True,
}
)
return
tool = next((candidate for candidate in self._custom_tools if candidate.name == tool_name), None)
tool = next(
(
candidate
for candidate in self._custom_tools
if candidate.name == tool_name
),
None,
)
if tool is None:
self._send_notification(
{
"type": "host_tool_result",
"id": request_id,
"result": {
"content": [{"type": "text", "text": f'Host tool "{tool_name}" is not registered'}],
"content": [
{
"type": "text",
"text": f'Host tool "{tool_name}" is not registered',
}
],
"details": {},
},
"isError": True,
@@ -1066,7 +1198,11 @@ class RpcClient:
tool_call_id=tool_call_id,
_cancel_event=pending_call.cancel_event,
_send_update=lambda result: self._send_notification(
{"type": "host_tool_update", "id": request_id, "partialResult": result}
{
"type": "host_tool_update",
"id": request_id,
"partialResult": result,
}
),
)
result = tool.execute(params, context)
@@ -1086,14 +1222,19 @@ class RpcClient:
{
"type": "host_tool_result",
"id": request_id,
"result": {"content": [{"type": "text", "text": str(exc)}], "details": {}},
"result": {
"content": [{"type": "text", "text": str(exc)}],
"details": {},
},
"isError": True,
}
)
finally:
self._pending_host_tool_calls.pop(request_id, None)
threading.Thread(target=run_tool, name=f"omp-rpc-host-tool:{tool_name}", daemon=True).start()
threading.Thread(
target=run_tool, name=f"omp-rpc-host-tool:{tool_name}", daemon=True
).start()
def _handle_host_tool_cancel(self, payload: JsonObject) -> None:
target_id = payload.get("targetId")
@@ -1117,10 +1258,16 @@ class RpcClient:
request_id = payload.get("id")
operation = payload.get("operation")
url = payload.get("url")
if not isinstance(request_id, str) or not isinstance(operation, str) or not isinstance(url, str):
if (
not isinstance(request_id, str)
or not isinstance(operation, str)
or not isinstance(url, str)
):
return
if operation not in ("read", "write"):
self._send_host_uri_error(request_id, f"Unsupported host URI operation: {operation}")
self._send_host_uri_error(
request_id, f"Unsupported host URI operation: {operation}"
)
return
try:
@@ -1131,14 +1278,20 @@ class RpcClient:
self._send_host_uri_error(request_id, f"Could not parse host URI: {url}")
return
scheme = (parsed.scheme or "").lower()
uri = next((candidate for candidate in self._host_uris if candidate.scheme == scheme), None)
uri = next(
(candidate for candidate in self._host_uris if candidate.scheme == scheme),
None,
)
if uri is None:
self._send_host_uri_error(request_id, f'Host URI scheme "{scheme}://" is not registered')
self._send_host_uri_error(
request_id, f'Host URI scheme "{scheme}://" is not registered'
)
return
if operation == "write" and uri.write is None:
self._send_host_uri_error(
request_id, f'Host URI scheme "{scheme}://" was not registered with a write handler'
request_id,
f'Host URI scheme "{scheme}://" was not registered with a write handler',
)
return
@@ -1147,7 +1300,11 @@ class RpcClient:
def run() -> None:
try:
context = HostUriContext(url=url, operation=cast(Any, operation), _cancel_event=pending.cancel_event)
context = HostUriContext(
url=url,
operation=cast(Any, operation),
_cancel_event=pending.cancel_event,
)
if operation == "read":
value = uri.read(url, context)
if pending.cancel_event.is_set():
@@ -1167,7 +1324,9 @@ class RpcClient:
uri.write(url, content, context)
if pending.cancel_event.is_set():
return
self._send_notification({"type": "host_uri_result", "id": request_id})
self._send_notification(
{"type": "host_uri_result", "id": request_id}
)
except Exception as exc:
if pending.cancel_event.is_set():
return
@@ -1175,7 +1334,9 @@ class RpcClient:
finally:
self._pending_host_uri_requests.pop(request_id, None)
threading.Thread(target=run, name=f"omp-rpc-host-uri:{scheme}:{operation}", daemon=True).start()
threading.Thread(
target=run, name=f"omp-rpc-host-uri:{scheme}:{operation}", daemon=True
).start()
def _handle_host_uri_cancel(self, payload: JsonObject) -> None:
target_id = payload.get("targetId")
@@ -1185,14 +1346,18 @@ class RpcClient:
if pending is not None:
pending.cancel_event.set()
def _add_typed_event_listener(self, event_type: str, listener: TEventListener) -> Callable[[], None]:
def _add_typed_event_listener(
self, event_type: str, listener: TEventListener
) -> Callable[[], None]:
listeners = self._typed_event_listeners.setdefault(event_type, [])
typed_listener = cast(AgentEventListener, listener)
listeners.append(typed_listener)
return lambda: self._remove_listener(listeners, typed_listener)
@staticmethod
def _normalize_todo_phases(todos: Sequence[TodoSeed | TodoPhaseSeed]) -> list[JsonObject]:
def _normalize_todo_phases(
todos: Sequence[TodoSeed | TodoPhaseSeed],
) -> list[JsonObject]:
if len(todos) == 0:
return []
@@ -1206,7 +1371,11 @@ class RpcClient:
def normalize_todo_item(seed: TodoSeed) -> JsonObject:
if isinstance(seed, str):
return {"id": next_task(), "content": seed, "status": cast(JsonValue, "pending")}
return {
"id": next_task(),
"content": seed,
"status": cast(JsonValue, "pending"),
}
if isinstance(seed, TodoItem):
if seed.status not in _TODO_STATUS_VALUES:
@@ -1234,7 +1403,9 @@ class RpcClient:
else:
status = "pending"
return {
"id": str(raw_id) if isinstance(raw_id, str) and raw_id else next_task(),
"id": str(raw_id)
if isinstance(raw_id, str) and raw_id
else next_task(),
"content": content,
"status": cast(JsonValue, status),
"notes": raw_notes if isinstance(raw_notes, str) else None,
@@ -1259,11 +1430,19 @@ class RpcClient:
raise RpcError("Todo phases must provide a non-empty 'name' value")
phase_id_value = seed.get("id")
raw_tasks = seed.get("tasks") or ()
if not isinstance(raw_tasks, Sequence) or isinstance(raw_tasks, (str, bytes)):
if not isinstance(raw_tasks, Sequence) or isinstance(
raw_tasks, (str, bytes)
):
raise RpcError("Todo phase 'tasks' must be a sequence")
phase_id = str(phase_id_value) if isinstance(phase_id_value, str) and phase_id_value else f"phase-{index}"
phase_id = (
str(phase_id_value)
if isinstance(phase_id_value, str) and phase_id_value
else f"phase-{index}"
)
name = raw_name
tasks = [normalize_todo_item(cast(TodoSeed, task)) for task in raw_tasks]
tasks = [
normalize_todo_item(cast(TodoSeed, task)) for task in raw_tasks
]
return {"id": phase_id, "name": name, "tasks": tasks}
@@ -1271,11 +1450,19 @@ class RpcClient:
phases: list[JsonObject] = []
for index, seed in enumerate(todos, start=1):
if not is_phase_seed(seed):
raise RpcError("Cannot mix flat todo items with todo phases in one set_todos() call")
raise RpcError(
"Cannot mix flat todo items with todo phases in one set_todos() call"
)
phases.append(normalize_phase(cast(TodoPhaseSeed, seed), index))
return phases
return [{"id": "phase-1", "name": "Todos", "tasks": [normalize_todo_item(cast(TodoSeed, todo)) for todo in todos]}]
return [
{
"id": "phase-1",
"name": "Todos",
"tasks": [normalize_todo_item(cast(TodoSeed, todo)) for todo in todos],
}
]
def _build_command(self) -> tuple[str, ...]:
if self._command is not None:
@@ -1305,7 +1492,9 @@ class RpcClient:
command.append("--no-skills")
if self._no_rules:
command.append("--no-rules")
emit_no_title = self._no_title if self._no_title is not None else self._rpc_defaults
emit_no_title = (
self._no_title if self._no_title is not None else self._rpc_defaults
)
if emit_no_title:
command.append("--no-title")
command.extend(self._extra_args)
@@ -1330,7 +1519,9 @@ class RpcClient:
process.stdin.write("\n")
process.stdin.flush()
except (BrokenPipeError, OSError) as exc:
raise RpcProcessExitError(f"Failed to write RPC command: {exc}") from exc
raise RpcProcessExitError(
f"Failed to write RPC command: {exc}"
) from exc
def _read_stdout_loop(self) -> None:
process = self._process
@@ -1382,7 +1573,12 @@ class RpcClient:
if isinstance(notification, ReadyEvent):
self._ready_received = True
self._ready.set()
self._dispatch_listeners("ready", listener_notification.type, self._ready_listeners, listener_notification)
self._dispatch_listeners(
"ready",
listener_notification.type,
self._ready_listeners,
listener_notification,
)
continue
if isinstance(notification, ExtensionUiRequest):
@@ -1417,9 +1613,14 @@ class RpcClient:
self._append_event(payload)
if listener_event.type == "agent_end":
self._mark_agent_run_completed()
self._dispatch_listeners("event", listener_event.type, self._event_listeners, listener_event)
self._dispatch_listeners(
"typed_event", listener_event.type, self._typed_event_listeners.get(listener_event.type, []), listener_event
"event", listener_event.type, self._event_listeners, listener_event
)
self._dispatch_listeners(
"typed_event",
listener_event.type,
self._typed_event_listeners.get(listener_event.type, []),
listener_event,
)
except Exception as exc:
self._mark_closed(exc)
@@ -1430,9 +1631,17 @@ class RpcClient:
try:
exit_code = process.wait(timeout=1.0)
except subprocess.TimeoutExpired:
self._mark_closed(RpcProcessExitError("RPC process stdout closed before the process exited"))
self._mark_closed(
RpcProcessExitError(
"RPC process stdout closed before the process exited"
)
)
return
self._mark_closed(RpcProcessExitError(f"RPC process exited with code {exit_code}. Stderr: {self.stderr}"))
self._mark_closed(
RpcProcessExitError(
f"RPC process exited with code {exit_code}. Stderr: {self.stderr}"
)
)
def _read_stderr_loop(self) -> None:
process = self._process
@@ -1478,8 +1687,13 @@ class RpcClient:
if protocol_error is None:
return
if protocol_error.command in _ASYNC_COMMANDS and protocol_error.remote_error is not None:
self._append_async_error(RpcCommandError(protocol_error.command, protocol_error.remote_error))
if (
protocol_error.command in _ASYNC_COMMANDS
and protocol_error.remote_error is not None
):
self._append_async_error(
RpcCommandError(protocol_error.command, protocol_error.remote_error)
)
self._mark_agent_run_completed()
self._record_protocol_error(protocol_error)
@@ -1493,7 +1707,11 @@ class RpcClient:
return False
with self._state_lock:
matching_ids = [request_id for request_id, pending in self._pending.items() if pending.command == command]
matching_ids = [
request_id
for request_id, pending in self._pending.items()
if pending.command == command
]
target_id: str | None = None
if len(matching_ids) == 1:
target_id = matching_ids[0]
@@ -1528,7 +1746,9 @@ class RpcClient:
def _record_protocol_error(self, error: RpcProtocolError) -> None:
with self._state_lock:
self._protocol_errors.append(error)
self._dispatch_listeners("protocol_error", error.command, self._protocol_error_listeners, error)
self._dispatch_listeners(
"protocol_error", error.command, self._protocol_error_listeners, error
)
def _record_listener_error(self, event: ListenerErrorEvent) -> None:
with self._state_lock:
+6 -2
View File
@@ -9,7 +9,9 @@ from .protocol import JsonObject
TPayload = TypeVar("TPayload")
HostUriContentType: TypeAlias = Literal["text/markdown", "application/json", "text/plain"]
HostUriContentType: TypeAlias = Literal[
"text/markdown", "application/json", "text/plain"
]
class HostUriReadResult(TypedDict, total=False):
@@ -99,7 +101,9 @@ def normalize_read_result(value: HostUriReadValue) -> JsonObject:
if isinstance(value, str):
return {"content": value}
if not isinstance(value, dict):
raise TypeError("Host URI read handlers must return a string or a HostUriReadResult mapping")
raise TypeError(
"Host URI read handlers must return a string or a HostUriReadResult mapping"
)
payload: JsonObject = {}
if "content" not in value:
+298 -80
View File
@@ -31,24 +31,38 @@ ExtensionUiMethod: TypeAlias = Literal[
"setTitle",
"set_editor_text",
]
InteractiveExtensionUiMethod: TypeAlias = Literal["select", "confirm", "input", "editor"]
PassiveExtensionUiMethod: TypeAlias = Literal["notify", "setStatus", "setWidget", "setTitle", "set_editor_text"]
InteractiveExtensionUiMethod: TypeAlias = Literal[
"select", "confirm", "input", "editor"
]
PassiveExtensionUiMethod: TypeAlias = Literal[
"notify", "setStatus", "setWidget", "setTitle", "set_editor_text"
]
ValueExtensionUiMethod: TypeAlias = Literal["select", "input", "editor"]
PASSIVE_EXTENSION_UI_METHODS: Final[frozenset[PassiveExtensionUiMethod]] = frozenset(
{"notify", "setStatus", "setWidget", "setTitle", "set_editor_text"}
)
INTERACTIVE_EXTENSION_UI_METHODS: Final[frozenset[InteractiveExtensionUiMethod]] = frozenset(
{"select", "confirm", "input", "editor"}
INTERACTIVE_EXTENSION_UI_METHODS: Final[frozenset[InteractiveExtensionUiMethod]] = (
frozenset({"select", "confirm", "input", "editor"})
)
VALUE_EXTENSION_UI_METHODS: Final[frozenset[ValueExtensionUiMethod]] = frozenset(
{"select", "input", "editor"}
)
_THINKING_LEVEL_VALUES: Final[frozenset[str]] = frozenset(
{"off", "minimal", "low", "medium", "high", "xhigh"}
)
VALUE_EXTENSION_UI_METHODS: Final[frozenset[ValueExtensionUiMethod]] = frozenset({"select", "input", "editor"})
_THINKING_LEVEL_VALUES: Final[frozenset[str]] = frozenset({"off", "minimal", "low", "medium", "high", "xhigh"})
_STEERING_MODE_VALUES: Final[frozenset[str]] = frozenset({"all", "one-at-a-time"})
_INTERRUPT_MODE_VALUES: Final[frozenset[str]] = frozenset({"immediate", "wait"})
_STOP_REASON_VALUES: Final[frozenset[str]] = frozenset({"stop", "length", "toolUse", "error", "aborted"})
_STOP_REASON_VALUES: Final[frozenset[str]] = frozenset(
{"stop", "length", "toolUse", "error", "aborted"}
)
_NOTIFY_TYPE_VALUES: Final[frozenset[str]] = frozenset({"info", "warning", "error"})
_WIDGET_PLACEMENT_VALUES: Final[frozenset[str]] = frozenset({"aboveEditor", "belowEditor"})
_TODO_STATUS_VALUES: Final[frozenset[str]] = frozenset({"pending", "in_progress", "completed", "abandoned"})
_WIDGET_PLACEMENT_VALUES: Final[frozenset[str]] = frozenset(
{"aboveEditor", "belowEditor"}
)
_TODO_STATUS_VALUES: Final[frozenset[str]] = frozenset(
{"pending", "in_progress", "completed", "abandoned"}
)
_EXTENSION_UI_METHOD_VALUES: Final[frozenset[str]] = frozenset(
{
"select",
@@ -94,10 +108,16 @@ _ASSISTANT_MESSAGE_EVENT_TYPE_VALUES: Final[frozenset[str]] = frozenset(
"error",
}
)
_ASSISTANT_DONE_REASON_VALUES: Final[frozenset[str]] = frozenset({"stop", "length", "toolUse"})
_ASSISTANT_DONE_REASON_VALUES: Final[frozenset[str]] = frozenset(
{"stop", "length", "toolUse"}
)
_ASSISTANT_ERROR_REASON_VALUES: Final[frozenset[str]] = frozenset({"aborted", "error"})
_AUTO_COMPACTION_REASON_VALUES: Final[frozenset[str]] = frozenset({"threshold", "overflow", "idle"})
_AUTO_COMPACTION_ACTION_VALUES: Final[frozenset[str]] = frozenset({"context-full", "handoff"})
_AUTO_COMPACTION_REASON_VALUES: Final[frozenset[str]] = frozenset(
{"threshold", "overflow", "idle"}
)
_AUTO_COMPACTION_ACTION_VALUES: Final[frozenset[str]] = frozenset(
{"context-full", "handoff"}
)
def _clone_json_value(value: object, *, field: str) -> JsonValue:
@@ -142,7 +162,9 @@ def _require_literal(value: object, allowed: frozenset[str], *, field: str) -> s
return value
def _optional_literal(value: object, allowed: frozenset[str], *, field: str) -> str | None:
def _optional_literal(
value: object, allowed: frozenset[str], *, field: str
) -> str | None:
if value is None:
return None
return _require_literal(value, allowed, field=field)
@@ -220,7 +242,9 @@ def _tuple_of_strings(values: object, *, field: str) -> tuple[str, ...] | None:
def _parse_agent_message(payload: JsonObject, *, field: str) -> AgentMessage:
_require_literal(payload.get("role"), _AGENT_MESSAGE_ROLE_VALUES, field=f"{field}.role")
_require_literal(
payload.get("role"), _AGENT_MESSAGE_ROLE_VALUES, field=f"{field}.role"
)
return cast(AgentMessage, _clone_json_object(payload, field=field))
@@ -246,7 +270,12 @@ def parse_agent_messages(payload: JsonValue | None) -> tuple[AgentMessage, ...]:
messages: list[AgentMessage] = []
for index, item in enumerate(payload):
messages.append(_parse_agent_message(_clone_json_object(item, field=f"messages[{index}]"), field=f"messages[{index}]"))
messages.append(
_parse_agent_message(
_clone_json_object(item, field=f"messages[{index}]"),
field=f"messages[{index}]",
)
)
return tuple(messages)
@@ -259,13 +288,17 @@ def parse_assistant_message_event(payload: JsonObject) -> AssistantMessageEvent:
if event_type == "start":
return AssistantMessageStartEvent(
partial=_parse_assistant_message(
_clone_json_object(payload.get("partial"), field="assistantMessageEvent.partial"),
_clone_json_object(
payload.get("partial"), field="assistantMessageEvent.partial"
),
field="assistantMessageEvent.partial",
)
)
if event_type in {"text_start", "thinking_start", "toolcall_start"}:
partial = _parse_assistant_message(
_clone_json_object(payload.get("partial"), field="assistantMessageEvent.partial"),
_clone_json_object(
payload.get("partial"), field="assistantMessageEvent.partial"
),
field="assistantMessageEvent.partial",
)
content_index = _optional_int(payload, "contentIndex")
@@ -274,11 +307,15 @@ def parse_assistant_message_event(payload: JsonObject) -> AssistantMessageEvent:
if event_type == "text_start":
return AssistantTextStartEvent(contentIndex=content_index, partial=partial)
if event_type == "thinking_start":
return AssistantThinkingStartEvent(contentIndex=content_index, partial=partial)
return AssistantThinkingStartEvent(
contentIndex=content_index, partial=partial
)
return AssistantToolCallStartEvent(contentIndex=content_index, partial=partial)
if event_type in {"text_delta", "thinking_delta", "toolcall_delta"}:
partial = _parse_assistant_message(
_clone_json_object(payload.get("partial"), field="assistantMessageEvent.partial"),
_clone_json_object(
payload.get("partial"), field="assistantMessageEvent.partial"
),
field="assistantMessageEvent.partial",
)
content_index = _optional_int(payload, "contentIndex")
@@ -288,13 +325,21 @@ def parse_assistant_message_event(payload: JsonObject) -> AssistantMessageEvent:
if delta is None:
raise ValueError("assistantMessageEvent.delta must be a string")
if event_type == "text_delta":
return AssistantTextDeltaEvent(contentIndex=content_index, delta=delta, partial=partial)
return AssistantTextDeltaEvent(
contentIndex=content_index, delta=delta, partial=partial
)
if event_type == "thinking_delta":
return AssistantThinkingDeltaEvent(contentIndex=content_index, delta=delta, partial=partial)
return AssistantToolCallDeltaEvent(contentIndex=content_index, delta=delta, partial=partial)
return AssistantThinkingDeltaEvent(
contentIndex=content_index, delta=delta, partial=partial
)
return AssistantToolCallDeltaEvent(
contentIndex=content_index, delta=delta, partial=partial
)
if event_type in {"text_end", "thinking_end"}:
partial = _parse_assistant_message(
_clone_json_object(payload.get("partial"), field="assistantMessageEvent.partial"),
_clone_json_object(
payload.get("partial"), field="assistantMessageEvent.partial"
),
field="assistantMessageEvent.partial",
)
content_index = _optional_int(payload, "contentIndex")
@@ -304,36 +349,60 @@ def parse_assistant_message_event(payload: JsonObject) -> AssistantMessageEvent:
if content is None:
raise ValueError("assistantMessageEvent.content must be a string")
if event_type == "text_end":
return AssistantTextEndEvent(contentIndex=content_index, content=content, partial=partial)
return AssistantThinkingEndEvent(contentIndex=content_index, content=content, partial=partial)
return AssistantTextEndEvent(
contentIndex=content_index, content=content, partial=partial
)
return AssistantThinkingEndEvent(
contentIndex=content_index, content=content, partial=partial
)
if event_type == "toolcall_end":
partial = _parse_assistant_message(
_clone_json_object(payload.get("partial"), field="assistantMessageEvent.partial"),
_clone_json_object(
payload.get("partial"), field="assistantMessageEvent.partial"
),
field="assistantMessageEvent.partial",
)
content_index = _optional_int(payload, "contentIndex")
if content_index is None:
raise ValueError("assistantMessageEvent.contentIndex must be an integer")
tool_call = _clone_json_object(payload.get("toolCall"), field="assistantMessageEvent.toolCall")
return AssistantToolCallEndEvent(contentIndex=content_index, toolCall=cast(ToolCall, tool_call), partial=partial)
tool_call = _clone_json_object(
payload.get("toolCall"), field="assistantMessageEvent.toolCall"
)
return AssistantToolCallEndEvent(
contentIndex=content_index,
toolCall=cast(ToolCall, tool_call),
partial=partial,
)
if event_type == "done":
return AssistantDoneEvent(
reason=cast(
Literal["stop", "length", "toolUse"],
_require_literal(payload.get("reason"), _ASSISTANT_DONE_REASON_VALUES, field="assistantMessageEvent.reason"),
_require_literal(
payload.get("reason"),
_ASSISTANT_DONE_REASON_VALUES,
field="assistantMessageEvent.reason",
),
),
message=_parse_assistant_message(
_clone_json_object(payload.get("message"), field="assistantMessageEvent.message"),
_clone_json_object(
payload.get("message"), field="assistantMessageEvent.message"
),
field="assistantMessageEvent.message",
),
)
return AssistantErrorEvent(
reason=cast(
Literal["aborted", "error"],
_require_literal(payload.get("reason"), _ASSISTANT_ERROR_REASON_VALUES, field="assistantMessageEvent.reason"),
_require_literal(
payload.get("reason"),
_ASSISTANT_ERROR_REASON_VALUES,
field="assistantMessageEvent.reason",
),
),
error=_parse_assistant_message(
_clone_json_object(payload.get("error"), field="assistantMessageEvent.error"),
_clone_json_object(
payload.get("error"), field="assistantMessageEvent.error"
),
field="assistantMessageEvent.error",
),
)
@@ -984,12 +1053,22 @@ RpcAgentEvent: TypeAlias = (
| TodoAutoClearEvent
)
RpcNotification: TypeAlias = ReadyEvent | ExtensionUiRequest | ExtensionError | RpcAgentEvent | UnknownNotification
RpcNotification: TypeAlias = (
ReadyEvent
| ExtensionUiRequest
| ExtensionError
| RpcAgentEvent
| UnknownNotification
)
def image_from_path(path: str | Path, mime_type: str | None = None) -> ImageContent:
file_path = Path(path)
resolved_mime_type = mime_type or mimetypes.guess_type(file_path.name)[0] or "application/octet-stream"
resolved_mime_type = (
mime_type
or mimetypes.guess_type(file_path.name)[0]
or "application/octet-stream"
)
return {
"type": "image",
"mimeType": resolved_mime_type,
@@ -997,9 +1076,18 @@ def image_from_path(path: str | Path, mime_type: str | None = None) -> ImageCont
}
def message_text(message: AgentMessage, *, include_thinking: bool = False) -> str | None:
def message_text(
message: AgentMessage, *, include_thinking: bool = False
) -> str | None:
role = message.get("role")
if role not in {"user", "developer", "assistant", "toolResult", "custom", "hookMessage"}:
if role not in {
"user",
"developer",
"assistant",
"toolResult",
"custom",
"hookMessage",
}:
return None
content = message.get("content")
@@ -1015,7 +1103,11 @@ def message_text(message: AgentMessage, *, include_thinking: bool = False) -> st
block_type = block.get("type")
if block_type == "text" and isinstance(block.get("text"), str):
fragments.append(cast(str, block["text"]))
elif include_thinking and block_type == "thinking" and isinstance(block.get("thinking"), str):
elif (
include_thinking
and block_type == "thinking"
and isinstance(block.get("thinking"), str)
):
fragments.append(cast(str, block["thinking"]))
return "".join(fragments) or None
@@ -1024,7 +1116,9 @@ def message_text_with_thinking(message: AgentMessage) -> str | None:
return message_text(message, include_thinking=True)
def assistant_text(message: AgentMessage, *, include_thinking: bool = False) -> str | None:
def assistant_text(
message: AgentMessage, *, include_thinking: bool = False
) -> str | None:
if message.get("role") != "assistant":
return None
return message_text(message, include_thinking=include_thinking)
@@ -1048,7 +1142,8 @@ def parse_model_info(payload: JsonObject | None) -> ModelInfo | None:
provider=_require_str(payload, "provider"),
base_url=_require_str(payload, "baseUrl"),
reasoning=bool(payload.get("reasoning", False)),
input_modalities=_tuple_of_strings(payload.get("input"), field="model.input") or (),
input_modalities=_tuple_of_strings(payload.get("input"), field="model.input")
or (),
cost=ModelCost(
input=float(cost_payload.get("input", 0.0)),
output=float(cost_payload.get("output", 0.0)),
@@ -1057,22 +1152,39 @@ def parse_model_info(payload: JsonObject | None) -> ModelInfo | None:
),
context_window=int(payload.get("contextWindow", 0)),
max_tokens=int(payload.get("maxTokens", 0)),
headers=cast(dict[str, str] | None, _optional_json_object(headers_payload, field="model.headers")),
premium_multiplier=float(payload["premiumMultiplier"]) if "premiumMultiplier" in payload else None,
prefer_websockets=bool(payload["preferWebsockets"]) if "preferWebsockets" in payload else None,
headers=cast(
dict[str, str] | None,
_optional_json_object(headers_payload, field="model.headers"),
),
premium_multiplier=float(payload["premiumMultiplier"])
if "premiumMultiplier" in payload
else None,
prefer_websockets=bool(payload["preferWebsockets"])
if "preferWebsockets" in payload
else None,
context_promotion_target=(
str(payload["contextPromotionTarget"]) if "contextPromotionTarget" in payload else None
str(payload["contextPromotionTarget"])
if "contextPromotionTarget" in payload
else None
),
priority=int(payload["priority"]) if "priority" in payload else None,
thinking=(
ThinkingConfig(
min_level=cast(
ThinkingLevel,
_require_literal(thinking_payload.get("minLevel"), _THINKING_LEVEL_VALUES, field="model.thinking.minLevel"),
_require_literal(
thinking_payload.get("minLevel"),
_THINKING_LEVEL_VALUES,
field="model.thinking.minLevel",
),
),
max_level=cast(
ThinkingLevel,
_require_literal(thinking_payload.get("maxLevel"), _THINKING_LEVEL_VALUES, field="model.thinking.maxLevel"),
_require_literal(
thinking_payload.get("maxLevel"),
_THINKING_LEVEL_VALUES,
field="model.thinking.maxLevel",
),
),
mode=_require_str(cast(JsonObject, thinking_payload), "mode"),
)
@@ -1087,7 +1199,9 @@ def parse_tool_descriptor(payload: JsonObject) -> ToolDescriptor:
return ToolDescriptor(
name=_require_str(payload, "name"),
description=_require_str(payload, "description"),
parameters=_clone_json_value(payload.get("parameters"), field="tool.parameters"),
parameters=_clone_json_value(
payload.get("parameters"), field="tool.parameters"
),
)
@@ -1097,7 +1211,11 @@ def parse_todo_item(payload: JsonObject) -> TodoItem:
content=_require_str(payload, "content"),
status=cast(
TodoStatus,
_require_literal(payload.get("status", "pending"), _TODO_STATUS_VALUES, field="todo.status"),
_require_literal(
payload.get("status", "pending"),
_TODO_STATUS_VALUES,
field="todo.status",
),
),
notes=_optional_str(payload, "notes"),
details=_optional_str(payload, "details"),
@@ -1111,7 +1229,10 @@ def parse_todo_phase(payload: JsonObject) -> TodoPhase:
else:
if not isinstance(raw_tasks, list):
raise ValueError("tasks must be a list")
tasks = tuple(parse_todo_item(_clone_json_object(item, field="tasks[]")) for item in raw_tasks)
tasks = tuple(
parse_todo_item(_clone_json_object(item, field="tasks[]"))
for item in raw_tasks
)
return TodoPhase(
id=str(payload.get("id", "")),
name=_require_str(payload, "name"),
@@ -1127,27 +1248,44 @@ def parse_todo_phases(payload: JsonValue | None) -> tuple[TodoPhase, ...]:
def parse_session_state(payload: JsonObject) -> SessionState:
dump_tools = tuple(
parse_tool_descriptor(_clone_json_object(item, field="dumpTools[]")) for item in cast(list[Any], payload.get("dumpTools") or [])
parse_tool_descriptor(_clone_json_object(item, field="dumpTools[]"))
for item in cast(list[Any], payload.get("dumpTools") or [])
)
return SessionState(
model=parse_model_info(cast(JsonObject | None, payload.get("model"))),
thinking_level=cast(
ThinkingLevel | None,
_optional_literal(payload.get("thinkingLevel"), _THINKING_LEVEL_VALUES, field="thinkingLevel"),
_optional_literal(
payload.get("thinkingLevel"),
_THINKING_LEVEL_VALUES,
field="thinkingLevel",
),
),
is_streaming=bool(payload.get("isStreaming", False)),
is_compacting=bool(payload.get("isCompacting", False)),
steering_mode=cast(
SteeringMode,
_require_literal(payload.get("steeringMode", "one-at-a-time"), _STEERING_MODE_VALUES, field="steeringMode"),
_require_literal(
payload.get("steeringMode", "one-at-a-time"),
_STEERING_MODE_VALUES,
field="steeringMode",
),
),
follow_up_mode=cast(
SteeringMode,
_require_literal(payload.get("followUpMode", "one-at-a-time"), _STEERING_MODE_VALUES, field="followUpMode"),
_require_literal(
payload.get("followUpMode", "one-at-a-time"),
_STEERING_MODE_VALUES,
field="followUpMode",
),
),
interrupt_mode=cast(
InterruptMode,
_require_literal(payload.get("interruptMode", "immediate"), _INTERRUPT_MODE_VALUES, field="interruptMode"),
_require_literal(
payload.get("interruptMode", "immediate"),
_INTERRUPT_MODE_VALUES,
field="interruptMode",
),
),
session_file=_optional_str(payload, "sessionFile"),
session_id=_require_str(payload, "sessionId"),
@@ -1155,7 +1293,9 @@ def parse_session_state(payload: JsonObject) -> SessionState:
auto_compaction_enabled=bool(payload.get("autoCompactionEnabled", False)),
message_count=int(payload.get("messageCount", 0)),
queued_message_count=int(payload.get("queuedMessageCount", 0)),
todo_phases=parse_todo_phases(cast(JsonValue | None, payload.get("todoPhases"))),
todo_phases=parse_todo_phases(
cast(JsonValue | None, payload.get("todoPhases"))
),
system_prompt=_optional_str_list(payload, "systemPrompt"),
dump_tools=dump_tools,
)
@@ -1181,8 +1321,12 @@ def parse_compaction_result(payload: JsonObject) -> CompactionResult:
short_summary=_optional_str(payload, "shortSummary"),
first_kept_entry_id=str(payload.get("firstKeptEntryId", "")),
tokens_before=int(payload.get("tokensBefore", 0)),
details=_clone_json_value(payload.get("details"), field="compaction.details") if "details" in payload else None,
preserve_data=_optional_json_object(payload.get("preserveData"), field="compaction.preserveData"),
details=_clone_json_value(payload.get("details"), field="compaction.details")
if "details" in payload
else None,
preserve_data=_optional_json_object(
payload.get("preserveData"), field="compaction.preserveData"
),
)
@@ -1199,7 +1343,9 @@ def parse_model_cycle_result(payload: JsonObject | None) -> ModelCycleResult | N
)
def parse_thinking_level_cycle_result(payload: JsonObject | None) -> ThinkingLevelCycleResult | None:
def parse_thinking_level_cycle_result(
payload: JsonObject | None,
) -> ThinkingLevelCycleResult | None:
if payload is None or payload.get("level") is None:
return None
return ThinkingLevelCycleResult(level=cast(ThinkingLevel, payload["level"]))
@@ -1211,7 +1357,10 @@ def parse_cancellation_result(payload: JsonObject | None) -> CancellationResult:
def parse_branch_result(payload: JsonObject | None) -> BranchResult:
payload = payload or {}
return BranchResult(text=str(payload.get("text", "")), cancelled=bool(payload.get("cancelled", False)))
return BranchResult(
text=str(payload.get("text", "")),
cancelled=bool(payload.get("cancelled", False)),
)
def parse_branch_messages(payload: JsonObject | None) -> tuple[BranchMessage, ...]:
@@ -1220,7 +1369,9 @@ def parse_branch_messages(payload: JsonObject | None) -> tuple[BranchMessage, ..
raise ValueError("messages must be a list")
return tuple(
BranchMessage(
entry_id=str(_clone_json_object(item, field="messages[]").get("entryId", "")),
entry_id=str(
_clone_json_object(item, field="messages[]").get("entryId", "")
),
text=str(_clone_json_object(item, field="messages[]").get("text", "")),
)
for item in messages
@@ -1228,7 +1379,9 @@ def parse_branch_messages(payload: JsonObject | None) -> tuple[BranchMessage, ..
def parse_session_stats(payload: JsonObject) -> SessionStats:
tokens_payload = _optional_json_object(payload.get("tokens"), field="sessionStats.tokens") or {}
tokens_payload = (
_optional_json_object(payload.get("tokens"), field="sessionStats.tokens") or {}
)
return SessionStats(
session_file=_optional_str(payload, "sessionFile"),
session_id=str(payload.get("sessionId", "")),
@@ -1254,10 +1407,16 @@ def parse_extension_ui_request(payload: JsonObject) -> ExtensionUiRequest:
id=_require_str(payload, "id"),
method=cast(
ExtensionUiMethod,
_require_literal(payload.get("method"), _EXTENSION_UI_METHOD_VALUES, field="extension_ui_request.method"),
_require_literal(
payload.get("method"),
_EXTENSION_UI_METHOD_VALUES,
field="extension_ui_request.method",
),
),
title=_optional_str(payload, "title"),
options=_tuple_of_strings(payload.get("options"), field="extension_ui_request.options"),
options=_tuple_of_strings(
payload.get("options"), field="extension_ui_request.options"
),
message=_optional_str(payload, "message"),
placeholder=_optional_str(payload, "placeholder"),
prefill=_optional_str(payload, "prefill"),
@@ -1266,12 +1425,18 @@ def parse_extension_ui_request(payload: JsonObject) -> ExtensionUiRequest:
target_id=_optional_str(payload, "targetId"),
notify_type=cast(
NotifyType | None,
_optional_literal(payload.get("notifyType"), _NOTIFY_TYPE_VALUES, field="extension_ui_request.notifyType"),
_optional_literal(
payload.get("notifyType"),
_NOTIFY_TYPE_VALUES,
field="extension_ui_request.notifyType",
),
),
status_key=_optional_str(payload, "statusKey"),
status_text=_optional_str(payload, "statusText"),
widget_key=_optional_str(payload, "widgetKey"),
widget_lines=_tuple_of_strings(payload.get("widgetLines"), field="extension_ui_request.widgetLines"),
widget_lines=_tuple_of_strings(
payload.get("widgetLines"), field="extension_ui_request.widgetLines"
),
widget_placement=cast(
WidgetPlacement | None,
_optional_literal(
@@ -1303,7 +1468,11 @@ def parse_notification(payload: JsonObject) -> RpcNotification:
if event_type == "agent_start":
return AgentStartEvent()
if event_type == "agent_end":
return AgentEndEvent(messages=parse_agent_messages(cast(JsonValue | None, payload.get("messages"))))
return AgentEndEvent(
messages=parse_agent_messages(
cast(JsonValue | None, payload.get("messages"))
)
)
if event_type == "turn_start":
return TurnStartEvent()
if event_type == "turn_end":
@@ -1313,25 +1482,35 @@ def parse_notification(payload: JsonObject) -> RpcNotification:
field="turn_end.message",
),
tool_results=tuple(
_parse_tool_result_message(_clone_json_object(item, field="turn_end.toolResults[]"), field="turn_end.toolResults[]")
_parse_tool_result_message(
_clone_json_object(item, field="turn_end.toolResults[]"),
field="turn_end.toolResults[]",
)
for item in cast(list[Any], payload.get("toolResults") or [])
),
)
if event_type == "message_start":
return MessageStartEvent(
message=_parse_agent_message(
_clone_json_object(payload.get("message"), field="message_start.message"),
_clone_json_object(
payload.get("message"), field="message_start.message"
),
field="message_start.message",
)
)
if event_type == "message_update":
return MessageUpdateEvent(
message=_parse_agent_message(
_clone_json_object(payload.get("message"), field="message_update.message"),
_clone_json_object(
payload.get("message"), field="message_update.message"
),
field="message_update.message",
),
assistant_message_event=parse_assistant_message_event(
_clone_json_object(payload.get("assistantMessageEvent"), field="message_update.assistantMessageEvent")
_clone_json_object(
payload.get("assistantMessageEvent"),
field="message_update.assistantMessageEvent",
)
),
)
if event_type == "message_end":
@@ -1345,16 +1524,27 @@ def parse_notification(payload: JsonObject) -> RpcNotification:
return ToolExecutionStartEvent(
tool_call_id=str(payload.get("toolCallId", "")),
tool_name=str(payload.get("toolName", "")),
args=_clone_json_value(payload.get("args"), field="tool_execution_start.args") if "args" in payload else None,
args=_clone_json_value(
payload.get("args"), field="tool_execution_start.args"
)
if "args" in payload
else None,
intent=_optional_str(payload, "intent"),
)
if event_type == "tool_execution_update":
return ToolExecutionUpdateEvent(
tool_call_id=str(payload.get("toolCallId", "")),
tool_name=str(payload.get("toolName", "")),
args=_clone_json_value(payload.get("args"), field="tool_execution_update.args") if "args" in payload else None,
args=_clone_json_value(
payload.get("args"), field="tool_execution_update.args"
)
if "args" in payload
else None,
partial_result=(
_clone_json_value(payload.get("partialResult"), field="tool_execution_update.partialResult")
_clone_json_value(
payload.get("partialResult"),
field="tool_execution_update.partialResult",
)
if "partialResult" in payload
else None
),
@@ -1363,18 +1553,30 @@ def parse_notification(payload: JsonObject) -> RpcNotification:
return ToolExecutionEndEvent(
tool_call_id=str(payload.get("toolCallId", "")),
tool_name=str(payload.get("toolName", "")),
result=_clone_json_value(payload.get("result"), field="tool_execution_end.result") if "result" in payload else None,
result=_clone_json_value(
payload.get("result"), field="tool_execution_end.result"
)
if "result" in payload
else None,
is_error=_optional_bool(payload, "isError"),
)
if event_type == "auto_compaction_start":
return AutoCompactionStartEvent(
reason=cast(
Literal["threshold", "overflow", "idle"],
_require_literal(payload.get("reason", "threshold"), _AUTO_COMPACTION_REASON_VALUES, field="auto_compaction_start.reason"),
_require_literal(
payload.get("reason", "threshold"),
_AUTO_COMPACTION_REASON_VALUES,
field="auto_compaction_start.reason",
),
),
action=cast(
Literal["context-full", "handoff"],
_require_literal(payload.get("action", "context-full"), _AUTO_COMPACTION_ACTION_VALUES, field="auto_compaction_start.action"),
_require_literal(
payload.get("action", "context-full"),
_AUTO_COMPACTION_ACTION_VALUES,
field="auto_compaction_start.action",
),
),
)
if event_type == "auto_compaction_end":
@@ -1382,10 +1584,18 @@ def parse_notification(payload: JsonObject) -> RpcNotification:
return AutoCompactionEndEvent(
action=cast(
Literal["context-full", "handoff"],
_require_literal(payload.get("action", "context-full"), _AUTO_COMPACTION_ACTION_VALUES, field="auto_compaction_end.action"),
_require_literal(
payload.get("action", "context-full"),
_AUTO_COMPACTION_ACTION_VALUES,
field="auto_compaction_end.action",
),
),
result=(
parse_compaction_result(_clone_json_object(result_payload, field="auto_compaction_end.result"))
parse_compaction_result(
_clone_json_object(
result_payload, field="auto_compaction_end.result"
)
)
if result_payload is not None
else None
),
@@ -1414,9 +1624,15 @@ def parse_notification(payload: JsonObject) -> RpcNotification:
role=str(payload.get("role", "")),
)
if event_type == "retry_fallback_succeeded":
return RetryFallbackSucceededEvent(model=str(payload.get("model", "")), role=str(payload.get("role", "")))
return RetryFallbackSucceededEvent(
model=str(payload.get("model", "")), role=str(payload.get("role", ""))
)
if event_type == "ttsr_triggered":
return TtsrTriggeredEvent(rules=_clone_json_objects(payload.get("rules"), field="ttsr_triggered.rules"))
return TtsrTriggeredEvent(
rules=_clone_json_objects(
payload.get("rules"), field="ttsr_triggered.rules"
)
)
if event_type == "todo_reminder":
return TodoReminderEvent(
todos=tuple(
@@ -1428,4 +1644,6 @@ def parse_notification(payload: JsonObject) -> RpcNotification:
)
if event_type == "todo_auto_clear":
return TodoAutoClearEvent()
return UnknownNotification(payload=_clone_json_object(payload, field="notification"))
return UnknownNotification(
payload=_clone_json_object(payload, field="notification")
)