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
+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: