feat(omp-rpc): introduced RpcClient with subprocess communication and event listeners
- Added RpcClient class with subprocess-based RPC communication, thread-safe request/response handling, and 30+ public methods for agent operations. - Defined 65+ public API symbols including protocol types, dataclasses, parser functions, and custom exception types for RPC communication. - Implemented event listener system with on_event, on_ui_request, and on_extension_error callbacks for asynchronous RPC event handling. - Added comprehensive test suite with integration tests for RpcClient and unit tests for protocol parsing functions. - Configured Python package with setuptools, PEP 561 type hints support, and documentation for usage and API reference.
This commit is contained in:
@@ -0,0 +1,70 @@
|
||||
# omp-rpc
|
||||
|
||||
Typed Python bindings for the `omp --mode rpc` protocol used by the coding agent.
|
||||
|
||||
This package wraps the newline-delimited JSON RPC transport exposed by the CLI and
|
||||
provides:
|
||||
|
||||
- typed command methods for the stable RPC surface
|
||||
- typed protocol models for state, bash results, compaction, and session stats
|
||||
- a process-backed client that manages request correlation over stdio
|
||||
- helpers for collecting prompt runs and handling extension UI requests
|
||||
|
||||
## Basic Usage
|
||||
|
||||
```python
|
||||
from omp_rpc import RpcClient
|
||||
|
||||
with RpcClient(provider="anthropic", model="claude-sonnet-4-5") as client:
|
||||
state = client.get_state()
|
||||
print(state.model.id if state.model else "no model")
|
||||
|
||||
turn = client.prompt_and_wait("Reply with just the word hello")
|
||||
print(turn.require_assistant_text())
|
||||
```
|
||||
|
||||
By default the client runs:
|
||||
|
||||
```bash
|
||||
omp --mode rpc
|
||||
```
|
||||
|
||||
You can also point it at a custom command, which is useful inside this repo while
|
||||
developing against the Bun entrypoint:
|
||||
|
||||
```python
|
||||
from omp_rpc import RpcClient
|
||||
|
||||
with RpcClient(
|
||||
command=[
|
||||
"bun",
|
||||
"packages/coding-agent/src/cli.ts",
|
||||
"--mode",
|
||||
"rpc",
|
||||
"--provider",
|
||||
"anthropic",
|
||||
"--model",
|
||||
"claude-sonnet-4-5",
|
||||
],
|
||||
) as client:
|
||||
print(client.get_state().session_id)
|
||||
```
|
||||
|
||||
## Extension UI Requests
|
||||
|
||||
Extensions in RPC mode can ask the host for input. Those requests are available as
|
||||
typed `ExtensionUiRequest` instances:
|
||||
|
||||
```python
|
||||
request = client.next_ui_request(timeout=5.0)
|
||||
|
||||
if request.method == "confirm":
|
||||
client.send_ui_confirmation(request.id, True)
|
||||
elif request.method in {"input", "editor"}:
|
||||
client.send_ui_value(request.id, "approved")
|
||||
```
|
||||
|
||||
## Protocol Reference
|
||||
|
||||
The canonical wire protocol still lives in the repo at
|
||||
[`docs/rpc.md`](../../docs/rpc.md).
|
||||
@@ -0,0 +1,36 @@
|
||||
[build-system]
|
||||
requires = ["setuptools>=69"]
|
||||
build-backend = "setuptools.build_meta"
|
||||
|
||||
[project]
|
||||
name = "omp-rpc"
|
||||
version = "0.1.0"
|
||||
description = "Typed Python client for the omp coding-agent RPC protocol"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.11"
|
||||
license = { text = "MIT" }
|
||||
authors = [{ name = "OpenAI Codex" }]
|
||||
keywords = ["omp", "rpc", "agent", "coding-agent", "jsonl", "stdio"]
|
||||
classifiers = [
|
||||
"Development Status :: 3 - Alpha",
|
||||
"Intended Audience :: Developers",
|
||||
"License :: OSI Approved :: MIT License",
|
||||
"Programming Language :: Python :: 3",
|
||||
"Programming Language :: Python :: 3.11",
|
||||
"Programming Language :: Python :: 3.12",
|
||||
"Topic :: Software Development :: Libraries :: Python Modules",
|
||||
"Topic :: Software Development :: Libraries",
|
||||
]
|
||||
|
||||
[project.urls]
|
||||
Repository = "https://github.com/can1357/oh-my-pi"
|
||||
Documentation = "https://github.com/can1357/oh-my-pi/blob/main/docs/rpc.md"
|
||||
|
||||
[tool.setuptools]
|
||||
package-dir = { "" = "src" }
|
||||
|
||||
[tool.setuptools.packages.find]
|
||||
where = ["src"]
|
||||
|
||||
[tool.setuptools.package-data]
|
||||
omp_rpc = ["py.typed"]
|
||||
@@ -0,0 +1,129 @@
|
||||
from .client import (
|
||||
AgentEventListener,
|
||||
ExtensionErrorListener,
|
||||
PromptTurn,
|
||||
RpcClient,
|
||||
RpcCommandError,
|
||||
RpcError,
|
||||
RpcProcessExitError,
|
||||
RpcTimeoutError,
|
||||
UiRequestListener,
|
||||
)
|
||||
from .protocol import (
|
||||
AgentEndEvent,
|
||||
AgentMessage,
|
||||
AssistantMessage,
|
||||
AssistantMessageEvent,
|
||||
AutoCompactionEndEvent,
|
||||
AutoCompactionStartEvent,
|
||||
AutoRetryEndEvent,
|
||||
AutoRetryStartEvent,
|
||||
BashResult,
|
||||
BranchMessage,
|
||||
BranchResult,
|
||||
CancellationResult,
|
||||
CompactionResult,
|
||||
CompactionSummaryMessage,
|
||||
CustomMessage,
|
||||
DeveloperMessage,
|
||||
ExtensionError,
|
||||
ExtensionUiRequest,
|
||||
FileMentionMessage,
|
||||
HookMessage,
|
||||
ImageContent,
|
||||
MessageEndEvent,
|
||||
MessageStartEvent,
|
||||
MessageUpdateEvent,
|
||||
ModelCycleResult,
|
||||
ModelInfo,
|
||||
ModelCost,
|
||||
PythonExecutionMessage,
|
||||
ReadyEvent,
|
||||
RetryFallbackAppliedEvent,
|
||||
RetryFallbackSucceededEvent,
|
||||
RpcAgentEvent,
|
||||
RpcNotification,
|
||||
SessionState,
|
||||
SessionStats,
|
||||
ThinkingConfig,
|
||||
ThinkingLevel,
|
||||
ThinkingLevelCycleResult,
|
||||
ToolDescriptor,
|
||||
ToolExecutionEndEvent,
|
||||
ToolExecutionStartEvent,
|
||||
ToolExecutionUpdateEvent,
|
||||
ToolResultMessage,
|
||||
TurnEndEvent,
|
||||
TurnStartEvent,
|
||||
UnknownNotification,
|
||||
UserMessage,
|
||||
assistant_text,
|
||||
image_from_path,
|
||||
message_text,
|
||||
parse_notification,
|
||||
parse_session_state,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"AgentEventListener",
|
||||
"AgentEndEvent",
|
||||
"AgentMessage",
|
||||
"AssistantMessage",
|
||||
"AssistantMessageEvent",
|
||||
"AutoCompactionEndEvent",
|
||||
"AutoCompactionStartEvent",
|
||||
"AutoRetryEndEvent",
|
||||
"AutoRetryStartEvent",
|
||||
"BashResult",
|
||||
"BranchMessage",
|
||||
"BranchResult",
|
||||
"CancellationResult",
|
||||
"CompactionResult",
|
||||
"CompactionSummaryMessage",
|
||||
"CustomMessage",
|
||||
"DeveloperMessage",
|
||||
"ExtensionError",
|
||||
"ExtensionErrorListener",
|
||||
"ExtensionUiRequest",
|
||||
"FileMentionMessage",
|
||||
"HookMessage",
|
||||
"ImageContent",
|
||||
"MessageEndEvent",
|
||||
"MessageStartEvent",
|
||||
"MessageUpdateEvent",
|
||||
"ModelCost",
|
||||
"ModelCycleResult",
|
||||
"ModelInfo",
|
||||
"PromptTurn",
|
||||
"PythonExecutionMessage",
|
||||
"ReadyEvent",
|
||||
"RetryFallbackAppliedEvent",
|
||||
"RetryFallbackSucceededEvent",
|
||||
"RpcAgentEvent",
|
||||
"RpcClient",
|
||||
"RpcCommandError",
|
||||
"RpcError",
|
||||
"RpcNotification",
|
||||
"RpcProcessExitError",
|
||||
"RpcTimeoutError",
|
||||
"SessionState",
|
||||
"SessionStats",
|
||||
"ThinkingConfig",
|
||||
"ThinkingLevel",
|
||||
"ThinkingLevelCycleResult",
|
||||
"ToolDescriptor",
|
||||
"ToolExecutionEndEvent",
|
||||
"ToolExecutionStartEvent",
|
||||
"ToolExecutionUpdateEvent",
|
||||
"ToolResultMessage",
|
||||
"TurnEndEvent",
|
||||
"TurnStartEvent",
|
||||
"UiRequestListener",
|
||||
"UnknownNotification",
|
||||
"UserMessage",
|
||||
"assistant_text",
|
||||
"image_from_path",
|
||||
"message_text",
|
||||
"parse_notification",
|
||||
"parse_session_state",
|
||||
]
|
||||
@@ -0,0 +1,600 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
import queue
|
||||
import subprocess
|
||||
import threading
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any, Callable, Mapping, Sequence, TypeVar, cast
|
||||
|
||||
from .protocol import (
|
||||
AgentEndEvent,
|
||||
AgentMessage,
|
||||
AssistantMessage,
|
||||
BashResult,
|
||||
BranchMessage,
|
||||
BranchResult,
|
||||
CancellationResult,
|
||||
CompactionResult,
|
||||
ExtensionError,
|
||||
ExtensionUiRequest,
|
||||
ImageContent,
|
||||
InterruptMode,
|
||||
JsonObject,
|
||||
JsonValue,
|
||||
ModelCycleResult,
|
||||
ModelInfo,
|
||||
ReadyEvent,
|
||||
RpcAgentEvent,
|
||||
SessionState,
|
||||
SessionStats,
|
||||
SteeringMode,
|
||||
StreamingBehavior,
|
||||
ThinkingLevel,
|
||||
ThinkingLevelCycleResult,
|
||||
assistant_text,
|
||||
parse_bash_result,
|
||||
parse_branch_messages,
|
||||
parse_branch_result,
|
||||
parse_cancellation_result,
|
||||
parse_compaction_result,
|
||||
parse_model_cycle_result,
|
||||
parse_model_info,
|
||||
parse_notification,
|
||||
parse_session_state,
|
||||
parse_session_stats,
|
||||
parse_thinking_level_cycle_result,
|
||||
)
|
||||
|
||||
AgentEventListener = Callable[[RpcAgentEvent], None]
|
||||
UiRequestListener = Callable[[ExtensionUiRequest], None]
|
||||
ExtensionErrorListener = Callable[[ExtensionError], None]
|
||||
TListener = TypeVar("TListener")
|
||||
|
||||
|
||||
class RpcError(RuntimeError):
|
||||
"""Base exception for the Python RPC client."""
|
||||
|
||||
|
||||
class RpcTimeoutError(RpcError):
|
||||
"""Raised when the server does not respond before a timeout."""
|
||||
|
||||
|
||||
class RpcProcessExitError(RpcError):
|
||||
"""Raised when the RPC process exits while a request is pending."""
|
||||
|
||||
|
||||
class RpcCommandError(RpcError):
|
||||
"""Raised when the RPC server returns `success: false`."""
|
||||
|
||||
def __init__(self, command: str, error: str):
|
||||
super().__init__(f"{command}: {error}")
|
||||
self.command = command
|
||||
self.error = error
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class PromptTurn:
|
||||
events: tuple[RpcAgentEvent, ...]
|
||||
messages: tuple[AgentMessage, ...]
|
||||
assistant_message: AssistantMessage | None
|
||||
assistant_text: str | None
|
||||
|
||||
def require_assistant_text(self) -> str:
|
||||
if self.assistant_text is None:
|
||||
raise RpcError("Prompt completed without a text assistant message")
|
||||
return self.assistant_text
|
||||
|
||||
|
||||
class RpcClient:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
command: Sequence[str] | None = None,
|
||||
executable: str = "omp",
|
||||
provider: str | None = None,
|
||||
model: str | None = None,
|
||||
session_dir: str | Path | None = None,
|
||||
cwd: str | Path | None = None,
|
||||
env: Mapping[str, str] | None = None,
|
||||
extra_args: Sequence[str] = (),
|
||||
startup_timeout: float = 30.0,
|
||||
request_timeout: float = 30.0,
|
||||
) -> None:
|
||||
self._command = tuple(command) if command is not None else None
|
||||
self._executable = executable
|
||||
self._provider = provider
|
||||
self._model = model
|
||||
self._session_dir = Path(session_dir) if session_dir is not None else None
|
||||
self._cwd = Path(cwd) if cwd is not None else None
|
||||
self._env = dict(env or {})
|
||||
self._extra_args = tuple(extra_args)
|
||||
self._startup_timeout = startup_timeout
|
||||
self._request_timeout = request_timeout
|
||||
|
||||
self._process: subprocess.Popen[str] | None = None
|
||||
self._stdout_thread: threading.Thread | None = None
|
||||
self._stderr_thread: threading.Thread | None = None
|
||||
self._ready = threading.Event()
|
||||
self._write_lock = threading.Lock()
|
||||
self._state_lock = threading.Lock()
|
||||
self._event_condition = threading.Condition()
|
||||
self._pending: dict[str, queue.Queue[JsonObject | BaseException]] = {}
|
||||
self._request_id = 0
|
||||
self._events: list[RpcAgentEvent] = []
|
||||
self._ui_requests: queue.Queue[ExtensionUiRequest] = queue.Queue()
|
||||
self._stderr_chunks: list[str] = []
|
||||
self._closed_error: BaseException | None = None
|
||||
self._stopping = False
|
||||
|
||||
self._event_listeners: list[AgentEventListener] = []
|
||||
self._ui_request_listeners: list[UiRequestListener] = []
|
||||
self._extension_error_listeners: list[ExtensionErrorListener] = []
|
||||
|
||||
def __enter__(self) -> RpcClient:
|
||||
return self.start()
|
||||
|
||||
def __exit__(self, _exc_type: object, _exc: object, _tb: object) -> None:
|
||||
self.stop()
|
||||
|
||||
@property
|
||||
def stderr(self) -> str:
|
||||
return "".join(self._stderr_chunks)
|
||||
|
||||
@property
|
||||
def command(self) -> tuple[str, ...]:
|
||||
return self._build_command()
|
||||
|
||||
def start(self) -> RpcClient:
|
||||
if self._process is not None:
|
||||
raise RpcError("RPC client is already started")
|
||||
|
||||
self._ready.clear()
|
||||
self._stopping = False
|
||||
self._closed_error = None
|
||||
self._events.clear()
|
||||
self._ui_requests = queue.Queue()
|
||||
self._stderr_chunks.clear()
|
||||
|
||||
process = subprocess.Popen(
|
||||
list(self._build_command()),
|
||||
cwd=str(self._cwd) if self._cwd is not None else None,
|
||||
env={**os.environ, **self._env},
|
||||
stdin=subprocess.PIPE,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.PIPE,
|
||||
text=True,
|
||||
encoding="utf-8",
|
||||
bufsize=1,
|
||||
)
|
||||
self._process = process
|
||||
|
||||
self._stdout_thread = threading.Thread(target=self._read_stdout_loop, name="omp-rpc-stdout", daemon=True)
|
||||
self._stderr_thread = threading.Thread(target=self._read_stderr_loop, name="omp-rpc-stderr", daemon=True)
|
||||
self._stdout_thread.start()
|
||||
self._stderr_thread.start()
|
||||
|
||||
if not self._ready.wait(self._startup_timeout):
|
||||
stderr = self.stderr
|
||||
self.stop()
|
||||
raise RpcTimeoutError(f"Timed out waiting for RPC ready signal. Stderr: {stderr}")
|
||||
|
||||
return self
|
||||
|
||||
def stop(self) -> None:
|
||||
process = self._process
|
||||
if process is None:
|
||||
return
|
||||
|
||||
self._stopping = True
|
||||
|
||||
try:
|
||||
if process.stdin is not None:
|
||||
try:
|
||||
process.stdin.close()
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
if process.poll() is None:
|
||||
process.terminate()
|
||||
try:
|
||||
process.wait(timeout=1.0)
|
||||
except subprocess.TimeoutExpired:
|
||||
process.kill()
|
||||
process.wait(timeout=1.0)
|
||||
finally:
|
||||
if process.stdout is not None:
|
||||
try:
|
||||
process.stdout.close()
|
||||
except OSError:
|
||||
pass
|
||||
if process.stderr is not None:
|
||||
try:
|
||||
process.stderr.close()
|
||||
except OSError:
|
||||
pass
|
||||
self._fail_pending(RpcProcessExitError("RPC process stopped"))
|
||||
self._process = None
|
||||
self._ready.set()
|
||||
with self._event_condition:
|
||||
self._event_condition.notify_all()
|
||||
if self._stdout_thread is not None:
|
||||
self._stdout_thread.join(timeout=1.0)
|
||||
if self._stderr_thread is not None:
|
||||
self._stderr_thread.join(timeout=1.0)
|
||||
self._stdout_thread = None
|
||||
self._stderr_thread = None
|
||||
|
||||
def on_event(self, listener: AgentEventListener) -> Callable[[], None]:
|
||||
self._event_listeners.append(listener)
|
||||
return lambda: self._remove_listener(self._event_listeners, listener)
|
||||
|
||||
def on_ui_request(self, listener: UiRequestListener) -> Callable[[], None]:
|
||||
self._ui_request_listeners.append(listener)
|
||||
return lambda: self._remove_listener(self._ui_request_listeners, listener)
|
||||
|
||||
def on_extension_error(self, listener: ExtensionErrorListener) -> Callable[[], None]:
|
||||
self._extension_error_listeners.append(listener)
|
||||
return lambda: self._remove_listener(self._extension_error_listeners, listener)
|
||||
|
||||
def next_ui_request(self, timeout: float | None = None) -> ExtensionUiRequest:
|
||||
try:
|
||||
return self._ui_requests.get(timeout=timeout)
|
||||
except queue.Empty as exc:
|
||||
raise RpcTimeoutError("Timed out waiting for an extension UI request") from exc
|
||||
|
||||
def send_ui_value(self, request_id: str, value: str) -> None:
|
||||
self._send_notification({"type": "extension_ui_response", "id": request_id, "value": value})
|
||||
|
||||
def send_ui_confirmation(self, request_id: str, confirmed: bool) -> None:
|
||||
self._send_notification({"type": "extension_ui_response", "id": request_id, "confirmed": confirmed})
|
||||
|
||||
def cancel_ui_request(self, request_id: str, *, timed_out: bool = False) -> None:
|
||||
payload: JsonObject = {"type": "extension_ui_response", "id": request_id, "cancelled": True}
|
||||
if timed_out:
|
||||
payload["timedOut"] = True
|
||||
self._send_notification(payload)
|
||||
|
||||
def get_state(self) -> SessionState:
|
||||
payload = self._request("get_state")
|
||||
return parse_session_state(payload)
|
||||
|
||||
def set_model(self, provider: str, model_id: str) -> ModelInfo:
|
||||
payload = self._request("set_model", provider=provider, modelId=model_id)
|
||||
model = parse_model_info(payload)
|
||||
if model is None:
|
||||
raise RpcError("set_model returned an empty payload")
|
||||
return model
|
||||
|
||||
def cycle_model(self) -> ModelCycleResult | None:
|
||||
return parse_model_cycle_result(self._request("cycle_model"))
|
||||
|
||||
def get_available_models(self) -> tuple[ModelInfo, ...]:
|
||||
payload = self._request("get_available_models")
|
||||
models = cast(list[JsonObject], payload.get("models") or [])
|
||||
return tuple(filter(None, (parse_model_info(model) for model in models)))
|
||||
|
||||
def set_thinking_level(self, level: ThinkingLevel) -> None:
|
||||
self._request("set_thinking_level", level=level)
|
||||
|
||||
def cycle_thinking_level(self) -> ThinkingLevelCycleResult | None:
|
||||
return parse_thinking_level_cycle_result(self._request("cycle_thinking_level"))
|
||||
|
||||
def set_steering_mode(self, mode: SteeringMode) -> None:
|
||||
self._request("set_steering_mode", mode=mode)
|
||||
|
||||
def set_follow_up_mode(self, mode: SteeringMode) -> None:
|
||||
self._request("set_follow_up_mode", mode=mode)
|
||||
|
||||
def set_interrupt_mode(self, mode: InterruptMode) -> None:
|
||||
self._request("set_interrupt_mode", mode=mode)
|
||||
|
||||
def compact(self, custom_instructions: str | None = None) -> CompactionResult:
|
||||
payload = self._request("compact", customInstructions=custom_instructions)
|
||||
return parse_compaction_result(payload)
|
||||
|
||||
def set_auto_compaction(self, enabled: bool) -> None:
|
||||
self._request("set_auto_compaction", enabled=enabled)
|
||||
|
||||
def set_auto_retry(self, enabled: bool) -> None:
|
||||
self._request("set_auto_retry", enabled=enabled)
|
||||
|
||||
def abort_retry(self) -> None:
|
||||
self._request("abort_retry")
|
||||
|
||||
def bash(self, command: str) -> BashResult:
|
||||
payload = self._request("bash", command=command)
|
||||
return parse_bash_result(payload)
|
||||
|
||||
def abort_bash(self) -> None:
|
||||
self._request("abort_bash")
|
||||
|
||||
def get_session_stats(self) -> SessionStats:
|
||||
payload = self._request("get_session_stats")
|
||||
return parse_session_stats(payload)
|
||||
|
||||
def export_html(self, output_path: str | Path | None = None) -> Path:
|
||||
payload = self._request("export_html", outputPath=str(output_path) if output_path is not None else None)
|
||||
return Path(str(payload["path"]))
|
||||
|
||||
def new_session(self, parent_session: str | None = None) -> CancellationResult:
|
||||
return parse_cancellation_result(self._request("new_session", parentSession=parent_session))
|
||||
|
||||
def switch_session(self, session_path: str | Path) -> CancellationResult:
|
||||
return parse_cancellation_result(self._request("switch_session", sessionPath=str(session_path)))
|
||||
|
||||
def branch(self, entry_id: str) -> BranchResult:
|
||||
return parse_branch_result(self._request("branch", entryId=entry_id))
|
||||
|
||||
def get_branch_messages(self) -> tuple[BranchMessage, ...]:
|
||||
return parse_branch_messages(self._request("get_branch_messages"))
|
||||
|
||||
def get_last_assistant_text(self) -> str | None:
|
||||
payload = self._request("get_last_assistant_text")
|
||||
value = payload.get("text")
|
||||
return str(value) if isinstance(value, str) else None
|
||||
|
||||
def set_session_name(self, name: str) -> None:
|
||||
self._request("set_session_name", name=name)
|
||||
|
||||
def get_messages(self) -> tuple[AgentMessage, ...]:
|
||||
payload = self._request("get_messages")
|
||||
return tuple(cast(list[AgentMessage], payload.get("messages") or []))
|
||||
|
||||
def prompt(
|
||||
self,
|
||||
message: str,
|
||||
*,
|
||||
images: Sequence[ImageContent] | None = None,
|
||||
streaming_behavior: StreamingBehavior | None = None,
|
||||
) -> None:
|
||||
self._request(
|
||||
"prompt",
|
||||
message=message,
|
||||
images=list(images) if images is not None else None,
|
||||
streamingBehavior=streaming_behavior,
|
||||
)
|
||||
|
||||
def steer(self, message: str, *, images: Sequence[ImageContent] | None = None) -> None:
|
||||
self._request("steer", message=message, images=list(images) if images is not None else None)
|
||||
|
||||
def follow_up(self, message: str, *, images: Sequence[ImageContent] | None = None) -> None:
|
||||
self._request("follow_up", message=message, images=list(images) if images is not None else None)
|
||||
|
||||
def abort(self) -> None:
|
||||
self._request("abort")
|
||||
|
||||
def abort_and_prompt(self, message: str, *, images: Sequence[ImageContent] | None = None) -> None:
|
||||
self._request("abort_and_prompt", message=message, images=list(images) if images is not None else None)
|
||||
|
||||
def prompt_and_wait(
|
||||
self,
|
||||
message: str,
|
||||
*,
|
||||
images: Sequence[ImageContent] | None = None,
|
||||
streaming_behavior: StreamingBehavior | None = None,
|
||||
timeout: float | None = None,
|
||||
) -> PromptTurn:
|
||||
start_index = self._current_event_index()
|
||||
self.prompt(message, images=images, streaming_behavior=streaming_behavior)
|
||||
events = self._wait_for_agent_end(start_index, timeout=timeout)
|
||||
return self._build_prompt_turn(events)
|
||||
|
||||
def wait_for_idle(self, timeout: float | None = None) -> None:
|
||||
start_index = self._current_event_index()
|
||||
self._wait_for_agent_end(start_index, timeout=timeout)
|
||||
|
||||
def collect_events(self, timeout: float | None = None) -> tuple[RpcAgentEvent, ...]:
|
||||
start_index = self._current_event_index()
|
||||
return self._wait_for_agent_end(start_index, timeout=timeout)
|
||||
|
||||
def request_raw(self, command_type: str, **payload: JsonValue) -> JsonObject:
|
||||
return self._request(command_type, **payload)
|
||||
|
||||
def _current_event_index(self) -> int:
|
||||
with self._event_condition:
|
||||
return len(self._events)
|
||||
|
||||
def _build_prompt_turn(self, events: tuple[RpcAgentEvent, ...]) -> PromptTurn:
|
||||
final_messages: tuple[AgentMessage, ...] = ()
|
||||
for event in reversed(events):
|
||||
if isinstance(event, AgentEndEvent):
|
||||
final_messages = event.messages
|
||||
break
|
||||
|
||||
assistant_message: AssistantMessage | None = None
|
||||
for message in reversed(final_messages):
|
||||
if message.get("role") == "assistant":
|
||||
assistant_message = cast(AssistantMessage, message)
|
||||
break
|
||||
|
||||
if assistant_message is None:
|
||||
for event in reversed(events):
|
||||
if hasattr(event, "message"):
|
||||
message = cast(AgentMessage | None, getattr(event, "message", None))
|
||||
if isinstance(message, dict) and message.get("role") == "assistant":
|
||||
assistant_message = cast(AssistantMessage, message)
|
||||
break
|
||||
|
||||
return PromptTurn(
|
||||
events=events,
|
||||
messages=final_messages,
|
||||
assistant_message=assistant_message,
|
||||
assistant_text=assistant_text(assistant_message) if assistant_message is not None else None,
|
||||
)
|
||||
|
||||
def _wait_for_agent_end(self, start_index: int, timeout: float | None = None) -> tuple[RpcAgentEvent, ...]:
|
||||
deadline = time.monotonic() + (timeout if timeout is not None else 60.0)
|
||||
with self._event_condition:
|
||||
while True:
|
||||
if self._closed_error is not None:
|
||||
raise RpcProcessExitError(str(self._closed_error))
|
||||
|
||||
events = tuple(self._events[start_index:])
|
||||
if any(isinstance(event, AgentEndEvent) for event in events):
|
||||
return events
|
||||
|
||||
remaining = deadline - time.monotonic()
|
||||
if remaining <= 0:
|
||||
raise RpcTimeoutError(f"Timed out waiting for agent_end. Stderr: {self.stderr}")
|
||||
self._event_condition.wait(remaining)
|
||||
|
||||
def _request(self, command_type: str, **payload: JsonValue) -> JsonObject:
|
||||
process = self._require_process()
|
||||
request_id = self._next_request_id()
|
||||
envelope: JsonObject = {"id": request_id, "type": command_type}
|
||||
for key, value in payload.items():
|
||||
if value is not None:
|
||||
envelope[key] = value
|
||||
|
||||
response_queue: queue.Queue[JsonObject | BaseException] = queue.Queue(maxsize=1)
|
||||
with self._state_lock:
|
||||
self._pending[request_id] = response_queue
|
||||
|
||||
self._write_json(process, envelope)
|
||||
|
||||
try:
|
||||
response = response_queue.get(timeout=self._request_timeout)
|
||||
except queue.Empty as exc:
|
||||
with self._state_lock:
|
||||
self._pending.pop(request_id, None)
|
||||
raise RpcTimeoutError(f"Timed out waiting for response to {command_type}. Stderr: {self.stderr}") from exc
|
||||
|
||||
if isinstance(response, BaseException):
|
||||
raise response
|
||||
|
||||
if not bool(response.get("success", False)):
|
||||
raise RpcCommandError(command=str(response.get("command", command_type)), error=str(response.get("error", "")))
|
||||
|
||||
data = response.get("data")
|
||||
return dict(cast(JsonObject, data or {}))
|
||||
|
||||
def _send_notification(self, payload: JsonObject) -> None:
|
||||
process = self._require_process()
|
||||
self._write_json(process, payload)
|
||||
|
||||
def _build_command(self) -> tuple[str, ...]:
|
||||
if self._command is not None:
|
||||
return self._command
|
||||
|
||||
command: list[str] = [self._executable, "--mode", "rpc"]
|
||||
if self._provider:
|
||||
command.extend(["--provider", self._provider])
|
||||
if self._model:
|
||||
command.extend(["--model", self._model])
|
||||
if self._session_dir is not None:
|
||||
command.extend(["--session-dir", str(self._session_dir)])
|
||||
command.extend(self._extra_args)
|
||||
return tuple(command)
|
||||
|
||||
def _next_request_id(self) -> str:
|
||||
with self._state_lock:
|
||||
self._request_id += 1
|
||||
return f"req_{self._request_id}"
|
||||
|
||||
def _require_process(self) -> subprocess.Popen[str]:
|
||||
if self._process is None:
|
||||
raise RpcError("RPC client is not started")
|
||||
return self._process
|
||||
|
||||
def _write_json(self, process: subprocess.Popen[str], payload: JsonObject) -> None:
|
||||
if process.stdin is None:
|
||||
raise RpcProcessExitError("RPC process stdin is unavailable")
|
||||
with self._write_lock:
|
||||
try:
|
||||
process.stdin.write(json.dumps(payload))
|
||||
process.stdin.write("\n")
|
||||
process.stdin.flush()
|
||||
except (BrokenPipeError, OSError) as exc:
|
||||
raise RpcProcessExitError(f"Failed to write RPC command: {exc}") from exc
|
||||
|
||||
def _read_stdout_loop(self) -> None:
|
||||
process = self._process
|
||||
if process is None or process.stdout is None:
|
||||
return
|
||||
|
||||
try:
|
||||
for line in process.stdout:
|
||||
stripped = line.strip()
|
||||
if not stripped:
|
||||
continue
|
||||
|
||||
payload = cast(JsonObject, json.loads(stripped))
|
||||
if payload.get("type") == "response":
|
||||
request_id = payload.get("id")
|
||||
if isinstance(request_id, str):
|
||||
with self._state_lock:
|
||||
pending = self._pending.pop(request_id, None)
|
||||
if pending is not None:
|
||||
pending.put(payload)
|
||||
continue
|
||||
|
||||
notification = parse_notification(payload)
|
||||
if isinstance(notification, ReadyEvent):
|
||||
self._ready.set()
|
||||
continue
|
||||
|
||||
if isinstance(notification, ExtensionUiRequest):
|
||||
self._ui_requests.put(notification)
|
||||
for listener in list(self._ui_request_listeners):
|
||||
listener(notification)
|
||||
continue
|
||||
|
||||
if isinstance(notification, ExtensionError):
|
||||
for listener in list(self._extension_error_listeners):
|
||||
listener(notification)
|
||||
continue
|
||||
|
||||
if getattr(notification, "type", None) != "unknown":
|
||||
with self._event_condition:
|
||||
self._events.append(cast(RpcAgentEvent, notification))
|
||||
self._event_condition.notify_all()
|
||||
for listener in list(self._event_listeners):
|
||||
listener(cast(RpcAgentEvent, notification))
|
||||
except json.JSONDecodeError as exc:
|
||||
self._mark_closed(RpcError(f"Failed to decode RPC output: {exc}"))
|
||||
except Exception as exc:
|
||||
self._mark_closed(exc)
|
||||
else:
|
||||
if not self._stopping:
|
||||
exit_code = process.poll()
|
||||
if exit_code is None:
|
||||
try:
|
||||
exit_code = process.wait(timeout=1.0)
|
||||
except subprocess.TimeoutExpired:
|
||||
self._mark_closed(RpcProcessExitError("RPC process stdout closed before the process exited"))
|
||||
return
|
||||
self._mark_closed(RpcProcessExitError(f"RPC process exited with code {exit_code}. Stderr: {self.stderr}"))
|
||||
|
||||
def _read_stderr_loop(self) -> None:
|
||||
process = self._process
|
||||
if process is None or process.stderr is None:
|
||||
return
|
||||
for chunk in process.stderr:
|
||||
self._stderr_chunks.append(chunk)
|
||||
|
||||
def _mark_closed(self, error: BaseException) -> None:
|
||||
if self._closed_error is not None:
|
||||
return
|
||||
self._closed_error = error
|
||||
self._ready.set()
|
||||
self._fail_pending(error)
|
||||
with self._event_condition:
|
||||
self._event_condition.notify_all()
|
||||
|
||||
def _fail_pending(self, error: BaseException) -> None:
|
||||
with self._state_lock:
|
||||
pending = list(self._pending.values())
|
||||
self._pending.clear()
|
||||
for response_queue in pending:
|
||||
response_queue.put(error)
|
||||
|
||||
@staticmethod
|
||||
def _remove_listener(listeners: list[TListener], listener: TListener) -> None:
|
||||
try:
|
||||
listeners.remove(listener)
|
||||
except ValueError:
|
||||
pass
|
||||
@@ -0,0 +1,963 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import mimetypes
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any, Literal, NotRequired, TypedDict, TypeAlias, cast
|
||||
|
||||
JsonPrimitive: TypeAlias = str | int | float | bool | None
|
||||
JsonValue: TypeAlias = JsonPrimitive | list["JsonValue"] | dict[str, "JsonValue"]
|
||||
JsonObject: TypeAlias = dict[str, JsonValue]
|
||||
|
||||
Attribution: TypeAlias = Literal["user", "agent"]
|
||||
ThinkingLevel: TypeAlias = Literal["off", "minimal", "low", "medium", "high", "xhigh"]
|
||||
StreamingBehavior: TypeAlias = Literal["steer", "followUp"]
|
||||
SteeringMode: TypeAlias = Literal["all", "one-at-a-time"]
|
||||
InterruptMode: TypeAlias = Literal["immediate", "wait"]
|
||||
StopReason: TypeAlias = Literal["stop", "length", "toolUse", "error", "aborted"]
|
||||
NotifyType: TypeAlias = Literal["info", "warning", "error"]
|
||||
WidgetPlacement: TypeAlias = Literal["aboveEditor", "belowEditor"]
|
||||
|
||||
|
||||
class TextContent(TypedDict, total=False):
|
||||
type: Literal["text"]
|
||||
text: str
|
||||
textSignature: NotRequired[str]
|
||||
|
||||
|
||||
class ThinkingContent(TypedDict, total=False):
|
||||
type: Literal["thinking"]
|
||||
thinking: str
|
||||
thinkingSignature: NotRequired[str]
|
||||
|
||||
|
||||
class RedactedThinkingContent(TypedDict, total=False):
|
||||
type: Literal["redactedThinking"]
|
||||
data: str
|
||||
|
||||
|
||||
class ImageContent(TypedDict, total=False):
|
||||
type: Literal["image"]
|
||||
data: str
|
||||
mimeType: str
|
||||
|
||||
|
||||
class ToolCall(TypedDict, total=False):
|
||||
type: Literal["toolCall"]
|
||||
id: str
|
||||
name: str
|
||||
arguments: dict[str, Any]
|
||||
thoughtSignature: NotRequired[str]
|
||||
intent: NotRequired[str]
|
||||
|
||||
|
||||
class UsageCost(TypedDict):
|
||||
input: float
|
||||
output: float
|
||||
cacheRead: float
|
||||
cacheWrite: float
|
||||
total: float
|
||||
|
||||
|
||||
class Usage(TypedDict, total=False):
|
||||
input: int
|
||||
output: int
|
||||
cacheRead: int
|
||||
cacheWrite: int
|
||||
totalTokens: int
|
||||
premiumRequests: NotRequired[int]
|
||||
cost: UsageCost
|
||||
|
||||
|
||||
class UserMessage(TypedDict, total=False):
|
||||
role: Literal["user"]
|
||||
content: str | list[TextContent | ImageContent]
|
||||
synthetic: NotRequired[bool]
|
||||
attribution: NotRequired[Attribution]
|
||||
providerPayload: NotRequired[JsonObject]
|
||||
timestamp: int
|
||||
|
||||
|
||||
class DeveloperMessage(TypedDict, total=False):
|
||||
role: Literal["developer"]
|
||||
content: str | list[TextContent | ImageContent]
|
||||
attribution: NotRequired[Attribution]
|
||||
providerPayload: NotRequired[JsonObject]
|
||||
timestamp: int
|
||||
|
||||
|
||||
class AssistantMessage(TypedDict, total=False):
|
||||
role: Literal["assistant"]
|
||||
content: list[TextContent | ThinkingContent | RedactedThinkingContent | ToolCall]
|
||||
api: str
|
||||
provider: str
|
||||
model: str
|
||||
responseId: NotRequired[str]
|
||||
usage: Usage
|
||||
stopReason: StopReason
|
||||
errorMessage: NotRequired[str]
|
||||
providerPayload: NotRequired[JsonObject]
|
||||
timestamp: int
|
||||
duration: NotRequired[int]
|
||||
ttft: NotRequired[int]
|
||||
|
||||
|
||||
class ToolResultMessage(TypedDict, total=False):
|
||||
role: Literal["toolResult"]
|
||||
toolCallId: str
|
||||
toolName: str
|
||||
content: list[TextContent | ImageContent]
|
||||
details: NotRequired[JsonValue]
|
||||
isError: bool
|
||||
attribution: NotRequired[Attribution]
|
||||
prunedAt: NotRequired[int]
|
||||
timestamp: int
|
||||
|
||||
|
||||
class BashExecutionMessage(TypedDict, total=False):
|
||||
role: Literal["bashExecution"]
|
||||
command: str
|
||||
output: str
|
||||
exitCode: int | None
|
||||
cancelled: bool
|
||||
truncated: bool
|
||||
meta: NotRequired[JsonObject]
|
||||
timestamp: int
|
||||
excludeFromContext: NotRequired[bool]
|
||||
|
||||
|
||||
class PythonExecutionMessage(TypedDict, total=False):
|
||||
role: Literal["pythonExecution"]
|
||||
code: str
|
||||
output: str
|
||||
exitCode: int | None
|
||||
cancelled: bool
|
||||
truncated: bool
|
||||
meta: NotRequired[JsonObject]
|
||||
timestamp: int
|
||||
excludeFromContext: NotRequired[bool]
|
||||
|
||||
|
||||
class CustomMessage(TypedDict, total=False):
|
||||
role: Literal["custom"]
|
||||
customType: str
|
||||
content: str | list[TextContent | ImageContent]
|
||||
display: bool
|
||||
details: NotRequired[JsonValue]
|
||||
attribution: NotRequired[Attribution]
|
||||
timestamp: int
|
||||
|
||||
|
||||
class HookMessage(TypedDict, total=False):
|
||||
role: Literal["hookMessage"]
|
||||
customType: str
|
||||
content: str | list[TextContent | ImageContent]
|
||||
display: bool
|
||||
details: NotRequired[JsonValue]
|
||||
attribution: NotRequired[Attribution]
|
||||
timestamp: int
|
||||
|
||||
|
||||
class BranchSummaryMessage(TypedDict, total=False):
|
||||
role: Literal["branchSummary"]
|
||||
summary: str
|
||||
fromId: str
|
||||
timestamp: int
|
||||
|
||||
|
||||
class CompactionSummaryMessage(TypedDict, total=False):
|
||||
role: Literal["compactionSummary"]
|
||||
summary: str
|
||||
shortSummary: NotRequired[str]
|
||||
tokensBefore: int
|
||||
providerPayload: NotRequired[JsonObject]
|
||||
timestamp: int
|
||||
|
||||
|
||||
class FileMentionItem(TypedDict, total=False):
|
||||
path: str
|
||||
content: str
|
||||
lineCount: NotRequired[int]
|
||||
byteSize: NotRequired[int]
|
||||
skippedReason: NotRequired[Literal["tooLarge"]]
|
||||
image: NotRequired[ImageContent]
|
||||
|
||||
|
||||
class FileMentionMessage(TypedDict, total=False):
|
||||
role: Literal["fileMention"]
|
||||
files: list[FileMentionItem]
|
||||
timestamp: int
|
||||
|
||||
|
||||
AgentMessage: TypeAlias = (
|
||||
UserMessage
|
||||
| DeveloperMessage
|
||||
| AssistantMessage
|
||||
| ToolResultMessage
|
||||
| BashExecutionMessage
|
||||
| PythonExecutionMessage
|
||||
| CustomMessage
|
||||
| HookMessage
|
||||
| BranchSummaryMessage
|
||||
| CompactionSummaryMessage
|
||||
| FileMentionMessage
|
||||
)
|
||||
|
||||
|
||||
class AssistantMessageStartEvent(TypedDict):
|
||||
type: Literal["start"]
|
||||
partial: AssistantMessage
|
||||
|
||||
|
||||
class AssistantTextStartEvent(TypedDict):
|
||||
type: Literal["text_start"]
|
||||
contentIndex: int
|
||||
partial: AssistantMessage
|
||||
|
||||
|
||||
class AssistantTextDeltaEvent(TypedDict):
|
||||
type: Literal["text_delta"]
|
||||
contentIndex: int
|
||||
delta: str
|
||||
partial: AssistantMessage
|
||||
|
||||
|
||||
class AssistantTextEndEvent(TypedDict):
|
||||
type: Literal["text_end"]
|
||||
contentIndex: int
|
||||
content: str
|
||||
partial: AssistantMessage
|
||||
|
||||
|
||||
class AssistantThinkingStartEvent(TypedDict):
|
||||
type: Literal["thinking_start"]
|
||||
contentIndex: int
|
||||
partial: AssistantMessage
|
||||
|
||||
|
||||
class AssistantThinkingDeltaEvent(TypedDict):
|
||||
type: Literal["thinking_delta"]
|
||||
contentIndex: int
|
||||
delta: str
|
||||
partial: AssistantMessage
|
||||
|
||||
|
||||
class AssistantThinkingEndEvent(TypedDict):
|
||||
type: Literal["thinking_end"]
|
||||
contentIndex: int
|
||||
content: str
|
||||
partial: AssistantMessage
|
||||
|
||||
|
||||
class AssistantToolCallStartEvent(TypedDict):
|
||||
type: Literal["toolcall_start"]
|
||||
contentIndex: int
|
||||
partial: AssistantMessage
|
||||
|
||||
|
||||
class AssistantToolCallDeltaEvent(TypedDict):
|
||||
type: Literal["toolcall_delta"]
|
||||
contentIndex: int
|
||||
delta: str
|
||||
partial: AssistantMessage
|
||||
|
||||
|
||||
class AssistantToolCallEndEvent(TypedDict):
|
||||
type: Literal["toolcall_end"]
|
||||
contentIndex: int
|
||||
toolCall: ToolCall
|
||||
partial: AssistantMessage
|
||||
|
||||
|
||||
class AssistantDoneEvent(TypedDict):
|
||||
type: Literal["done"]
|
||||
reason: Literal["stop", "length", "toolUse"]
|
||||
message: AssistantMessage
|
||||
|
||||
|
||||
class AssistantErrorEvent(TypedDict):
|
||||
type: Literal["error"]
|
||||
reason: Literal["aborted", "error"]
|
||||
error: AssistantMessage
|
||||
|
||||
|
||||
AssistantMessageEvent: TypeAlias = (
|
||||
AssistantMessageStartEvent
|
||||
| AssistantTextStartEvent
|
||||
| AssistantTextDeltaEvent
|
||||
| AssistantTextEndEvent
|
||||
| AssistantThinkingStartEvent
|
||||
| AssistantThinkingDeltaEvent
|
||||
| AssistantThinkingEndEvent
|
||||
| AssistantToolCallStartEvent
|
||||
| AssistantToolCallDeltaEvent
|
||||
| AssistantToolCallEndEvent
|
||||
| AssistantDoneEvent
|
||||
| AssistantErrorEvent
|
||||
)
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class ModelCost:
|
||||
input: float
|
||||
output: float
|
||||
cache_read: float
|
||||
cache_write: float
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class ThinkingConfig:
|
||||
min_level: ThinkingLevel
|
||||
max_level: ThinkingLevel
|
||||
mode: str
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class ModelInfo:
|
||||
id: str
|
||||
name: str
|
||||
api: str
|
||||
provider: str
|
||||
base_url: str
|
||||
reasoning: bool
|
||||
input_modalities: tuple[str, ...]
|
||||
cost: ModelCost
|
||||
context_window: int
|
||||
max_tokens: int
|
||||
headers: dict[str, str] | None = None
|
||||
premium_multiplier: float | None = None
|
||||
prefer_websockets: bool | None = None
|
||||
context_promotion_target: str | None = None
|
||||
priority: int | None = None
|
||||
thinking: ThinkingConfig | None = None
|
||||
compat: JsonObject | None = None
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class ToolDescriptor:
|
||||
name: str
|
||||
description: str
|
||||
parameters: JsonValue
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class SessionState:
|
||||
model: ModelInfo | None
|
||||
thinking_level: ThinkingLevel | None
|
||||
is_streaming: bool
|
||||
is_compacting: bool
|
||||
steering_mode: SteeringMode
|
||||
follow_up_mode: SteeringMode
|
||||
interrupt_mode: InterruptMode
|
||||
session_file: str | None
|
||||
session_id: str
|
||||
session_name: str | None
|
||||
auto_compaction_enabled: bool
|
||||
message_count: int
|
||||
queued_message_count: int
|
||||
system_prompt: str | None = None
|
||||
dump_tools: tuple[ToolDescriptor, ...] = ()
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class BashResult:
|
||||
output: str
|
||||
exit_code: int | None
|
||||
cancelled: bool
|
||||
truncated: bool
|
||||
total_lines: int
|
||||
total_bytes: int
|
||||
output_lines: int
|
||||
output_bytes: int
|
||||
artifact_id: str | None = None
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class CompactionResult:
|
||||
summary: str
|
||||
first_kept_entry_id: str
|
||||
tokens_before: int
|
||||
short_summary: str | None = None
|
||||
details: JsonValue | None = None
|
||||
preserve_data: JsonObject | None = None
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class ModelCycleResult:
|
||||
model: ModelInfo
|
||||
thinking_level: ThinkingLevel | None
|
||||
is_scoped: bool
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class ThinkingLevelCycleResult:
|
||||
level: ThinkingLevel
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class CancellationResult:
|
||||
cancelled: bool
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class BranchMessage:
|
||||
entry_id: str
|
||||
text: str
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class BranchResult:
|
||||
text: str
|
||||
cancelled: bool
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class TokenUsage:
|
||||
input: int
|
||||
output: int
|
||||
cache_read: int
|
||||
cache_write: int
|
||||
total: int
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class SessionStats:
|
||||
session_file: str | None
|
||||
session_id: str
|
||||
user_messages: int
|
||||
assistant_messages: int
|
||||
tool_calls: int
|
||||
tool_results: int
|
||||
total_messages: int
|
||||
tokens: TokenUsage
|
||||
premium_requests: int
|
||||
cost: float
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class ReadyEvent:
|
||||
type: Literal["ready"] = "ready"
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class ExtensionUiRequest:
|
||||
id: str
|
||||
method: str
|
||||
title: str | None = None
|
||||
options: tuple[str, ...] | None = None
|
||||
message: str | None = None
|
||||
placeholder: str | None = None
|
||||
prefill: str | None = None
|
||||
timeout: int | None = None
|
||||
prompt_style: bool | None = None
|
||||
target_id: str | None = None
|
||||
notify_type: NotifyType | None = None
|
||||
status_key: str | None = None
|
||||
status_text: str | None = None
|
||||
widget_key: str | None = None
|
||||
widget_lines: tuple[str, ...] | None = None
|
||||
widget_placement: WidgetPlacement | None = None
|
||||
text: str | None = None
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class ExtensionError:
|
||||
extension_path: str
|
||||
event: str
|
||||
error: str
|
||||
type: Literal["extension_error"] = "extension_error"
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class AgentStartEvent:
|
||||
type: Literal["agent_start"] = "agent_start"
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class AgentEndEvent:
|
||||
messages: tuple[AgentMessage, ...]
|
||||
type: Literal["agent_end"] = "agent_end"
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class TurnStartEvent:
|
||||
type: Literal["turn_start"] = "turn_start"
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class TurnEndEvent:
|
||||
message: AgentMessage
|
||||
tool_results: tuple[ToolResultMessage, ...]
|
||||
type: Literal["turn_end"] = "turn_end"
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class MessageStartEvent:
|
||||
message: AgentMessage
|
||||
type: Literal["message_start"] = "message_start"
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class MessageUpdateEvent:
|
||||
message: AgentMessage
|
||||
assistant_message_event: AssistantMessageEvent
|
||||
type: Literal["message_update"] = "message_update"
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class MessageEndEvent:
|
||||
message: AgentMessage
|
||||
type: Literal["message_end"] = "message_end"
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class ToolExecutionStartEvent:
|
||||
tool_call_id: str
|
||||
tool_name: str
|
||||
args: JsonValue
|
||||
intent: str | None = None
|
||||
type: Literal["tool_execution_start"] = "tool_execution_start"
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class ToolExecutionUpdateEvent:
|
||||
tool_call_id: str
|
||||
tool_name: str
|
||||
args: JsonValue
|
||||
partial_result: JsonValue
|
||||
type: Literal["tool_execution_update"] = "tool_execution_update"
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class ToolExecutionEndEvent:
|
||||
tool_call_id: str
|
||||
tool_name: str
|
||||
result: JsonValue
|
||||
is_error: bool | None = None
|
||||
type: Literal["tool_execution_end"] = "tool_execution_end"
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class AutoCompactionStartEvent:
|
||||
reason: Literal["threshold", "overflow", "idle"]
|
||||
action: Literal["context-full", "handoff"]
|
||||
type: Literal["auto_compaction_start"] = "auto_compaction_start"
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class AutoCompactionEndEvent:
|
||||
action: Literal["context-full", "handoff"]
|
||||
result: CompactionResult | None
|
||||
aborted: bool
|
||||
will_retry: bool
|
||||
error_message: str | None = None
|
||||
skipped: bool | None = None
|
||||
type: Literal["auto_compaction_end"] = "auto_compaction_end"
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class AutoRetryStartEvent:
|
||||
attempt: int
|
||||
max_attempts: int
|
||||
delay_ms: int
|
||||
error_message: str
|
||||
type: Literal["auto_retry_start"] = "auto_retry_start"
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class AutoRetryEndEvent:
|
||||
success: bool
|
||||
attempt: int
|
||||
final_error: str | None = None
|
||||
type: Literal["auto_retry_end"] = "auto_retry_end"
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class RetryFallbackAppliedEvent:
|
||||
from_model: str
|
||||
to_model: str
|
||||
role: str
|
||||
type: Literal["retry_fallback_applied"] = "retry_fallback_applied"
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class RetryFallbackSucceededEvent:
|
||||
model: str
|
||||
role: str
|
||||
type: Literal["retry_fallback_succeeded"] = "retry_fallback_succeeded"
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class TtsrTriggeredEvent:
|
||||
rules: tuple[JsonObject, ...]
|
||||
type: Literal["ttsr_triggered"] = "ttsr_triggered"
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class TodoReminderEvent:
|
||||
todos: tuple[JsonObject, ...]
|
||||
attempt: int
|
||||
max_attempts: int
|
||||
type: Literal["todo_reminder"] = "todo_reminder"
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class TodoAutoClearEvent:
|
||||
type: Literal["todo_auto_clear"] = "todo_auto_clear"
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class UnknownNotification:
|
||||
payload: JsonObject
|
||||
type: Literal["unknown"] = "unknown"
|
||||
|
||||
|
||||
RpcAgentEvent: TypeAlias = (
|
||||
AgentStartEvent
|
||||
| AgentEndEvent
|
||||
| TurnStartEvent
|
||||
| TurnEndEvent
|
||||
| MessageStartEvent
|
||||
| MessageUpdateEvent
|
||||
| MessageEndEvent
|
||||
| ToolExecutionStartEvent
|
||||
| ToolExecutionUpdateEvent
|
||||
| ToolExecutionEndEvent
|
||||
| AutoCompactionStartEvent
|
||||
| AutoCompactionEndEvent
|
||||
| AutoRetryStartEvent
|
||||
| AutoRetryEndEvent
|
||||
| RetryFallbackAppliedEvent
|
||||
| RetryFallbackSucceededEvent
|
||||
| TtsrTriggeredEvent
|
||||
| TodoReminderEvent
|
||||
| TodoAutoClearEvent
|
||||
)
|
||||
|
||||
RpcNotification: TypeAlias = ReadyEvent | ExtensionUiRequest | ExtensionError | RpcAgentEvent | UnknownNotification
|
||||
|
||||
|
||||
def image_from_path(path: str | Path, mime_type: str | None = None) -> ImageContent:
|
||||
file_path = Path(path)
|
||||
resolved_mime_type = mime_type or mimetypes.guess_type(file_path.name)[0] or "application/octet-stream"
|
||||
return {
|
||||
"type": "image",
|
||||
"mimeType": resolved_mime_type,
|
||||
"data": base64.b64encode(file_path.read_bytes()).decode("ascii"),
|
||||
}
|
||||
|
||||
|
||||
def message_text(message: AgentMessage) -> str | None:
|
||||
role = message.get("role")
|
||||
if role not in {"user", "developer", "assistant", "toolResult", "custom", "hookMessage"}:
|
||||
return None
|
||||
|
||||
content = message.get("content")
|
||||
if isinstance(content, str):
|
||||
return content
|
||||
if not isinstance(content, list):
|
||||
return None
|
||||
|
||||
fragments: list[str] = []
|
||||
for block in content:
|
||||
if not isinstance(block, dict):
|
||||
continue
|
||||
block_type = block.get("type")
|
||||
if block_type == "text" and isinstance(block.get("text"), str):
|
||||
fragments.append(cast(str, block["text"]))
|
||||
elif block_type == "thinking" and isinstance(block.get("thinking"), str):
|
||||
fragments.append(cast(str, block["thinking"]))
|
||||
return "".join(fragments) or None
|
||||
|
||||
|
||||
def assistant_text(message: AgentMessage) -> str | None:
|
||||
if message.get("role") != "assistant":
|
||||
return None
|
||||
return message_text(message)
|
||||
|
||||
|
||||
def parse_model_info(payload: JsonObject | None) -> ModelInfo | None:
|
||||
if payload is None:
|
||||
return None
|
||||
cost_payload = cast(dict[str, Any], payload.get("cost") or {})
|
||||
thinking_payload = payload.get("thinking")
|
||||
headers_payload = payload.get("headers")
|
||||
compat_payload = payload.get("compat")
|
||||
return ModelInfo(
|
||||
id=str(payload["id"]),
|
||||
name=str(payload["name"]),
|
||||
api=str(payload["api"]),
|
||||
provider=str(payload["provider"]),
|
||||
base_url=str(payload["baseUrl"]),
|
||||
reasoning=bool(payload.get("reasoning", False)),
|
||||
input_modalities=tuple(str(item) for item in cast(list[Any], payload.get("input") or [])),
|
||||
cost=ModelCost(
|
||||
input=float(cost_payload.get("input", 0.0)),
|
||||
output=float(cost_payload.get("output", 0.0)),
|
||||
cache_read=float(cost_payload.get("cacheRead", 0.0)),
|
||||
cache_write=float(cost_payload.get("cacheWrite", 0.0)),
|
||||
),
|
||||
context_window=int(payload.get("contextWindow", 0)),
|
||||
max_tokens=int(payload.get("maxTokens", 0)),
|
||||
headers=dict(cast(dict[str, str], headers_payload)) if isinstance(headers_payload, dict) else None,
|
||||
premium_multiplier=float(payload["premiumMultiplier"]) if "premiumMultiplier" in payload else None,
|
||||
prefer_websockets=bool(payload["preferWebsockets"]) if "preferWebsockets" in payload else None,
|
||||
context_promotion_target=(
|
||||
str(payload["contextPromotionTarget"]) if "contextPromotionTarget" in payload else None
|
||||
),
|
||||
priority=int(payload["priority"]) if "priority" in payload else None,
|
||||
thinking=(
|
||||
ThinkingConfig(
|
||||
min_level=cast(ThinkingLevel, thinking_payload["minLevel"]),
|
||||
max_level=cast(ThinkingLevel, thinking_payload["maxLevel"]),
|
||||
mode=str(thinking_payload["mode"]),
|
||||
)
|
||||
if isinstance(thinking_payload, dict)
|
||||
else None
|
||||
),
|
||||
compat=dict(cast(dict[str, JsonValue], compat_payload)) if isinstance(compat_payload, dict) else None,
|
||||
)
|
||||
|
||||
|
||||
def parse_tool_descriptor(payload: JsonObject) -> ToolDescriptor:
|
||||
return ToolDescriptor(
|
||||
name=str(payload["name"]),
|
||||
description=str(payload["description"]),
|
||||
parameters=cast(JsonValue, payload.get("parameters")),
|
||||
)
|
||||
|
||||
|
||||
def parse_session_state(payload: JsonObject) -> SessionState:
|
||||
dump_tools = tuple(
|
||||
parse_tool_descriptor(cast(JsonObject, item)) for item in cast(list[Any], payload.get("dumpTools") or [])
|
||||
)
|
||||
return SessionState(
|
||||
model=parse_model_info(cast(JsonObject | None, payload.get("model"))),
|
||||
thinking_level=cast(ThinkingLevel | None, payload.get("thinkingLevel")),
|
||||
is_streaming=bool(payload.get("isStreaming", False)),
|
||||
is_compacting=bool(payload.get("isCompacting", False)),
|
||||
steering_mode=cast(SteeringMode, payload.get("steeringMode", "one-at-a-time")),
|
||||
follow_up_mode=cast(SteeringMode, payload.get("followUpMode", "one-at-a-time")),
|
||||
interrupt_mode=cast(InterruptMode, payload.get("interruptMode", "immediate")),
|
||||
session_file=str(payload["sessionFile"]) if payload.get("sessionFile") is not None else None,
|
||||
session_id=str(payload["sessionId"]),
|
||||
session_name=str(payload["sessionName"]) if payload.get("sessionName") is not None else None,
|
||||
auto_compaction_enabled=bool(payload.get("autoCompactionEnabled", False)),
|
||||
message_count=int(payload.get("messageCount", 0)),
|
||||
queued_message_count=int(payload.get("queuedMessageCount", 0)),
|
||||
system_prompt=str(payload["systemPrompt"]) if payload.get("systemPrompt") is not None else None,
|
||||
dump_tools=dump_tools,
|
||||
)
|
||||
|
||||
|
||||
def parse_bash_result(payload: JsonObject) -> BashResult:
|
||||
return BashResult(
|
||||
output=str(payload.get("output", "")),
|
||||
exit_code=int(payload["exitCode"]) if payload.get("exitCode") is not None else None,
|
||||
cancelled=bool(payload.get("cancelled", False)),
|
||||
truncated=bool(payload.get("truncated", False)),
|
||||
total_lines=int(payload.get("totalLines", 0)),
|
||||
total_bytes=int(payload.get("totalBytes", 0)),
|
||||
output_lines=int(payload.get("outputLines", 0)),
|
||||
output_bytes=int(payload.get("outputBytes", 0)),
|
||||
artifact_id=str(payload["artifactId"]) if payload.get("artifactId") is not None else None,
|
||||
)
|
||||
|
||||
|
||||
def parse_compaction_result(payload: JsonObject) -> CompactionResult:
|
||||
return CompactionResult(
|
||||
summary=str(payload.get("summary", "")),
|
||||
short_summary=str(payload["shortSummary"]) if payload.get("shortSummary") is not None else None,
|
||||
first_kept_entry_id=str(payload.get("firstKeptEntryId", "")),
|
||||
tokens_before=int(payload.get("tokensBefore", 0)),
|
||||
details=cast(JsonValue | None, payload.get("details")),
|
||||
preserve_data=cast(JsonObject | None, payload.get("preserveData")),
|
||||
)
|
||||
|
||||
|
||||
def parse_model_cycle_result(payload: JsonObject | None) -> ModelCycleResult | None:
|
||||
if payload is None:
|
||||
return None
|
||||
model = parse_model_info(cast(JsonObject, payload.get("model")))
|
||||
if model is None:
|
||||
raise ValueError("cycle_model response did not include a model")
|
||||
return ModelCycleResult(
|
||||
model=model,
|
||||
thinking_level=cast(ThinkingLevel | None, payload.get("thinkingLevel")),
|
||||
is_scoped=bool(payload.get("isScoped", False)),
|
||||
)
|
||||
|
||||
|
||||
def parse_thinking_level_cycle_result(payload: JsonObject | None) -> ThinkingLevelCycleResult | None:
|
||||
if payload is None or payload.get("level") is None:
|
||||
return None
|
||||
return ThinkingLevelCycleResult(level=cast(ThinkingLevel, payload["level"]))
|
||||
|
||||
|
||||
def parse_cancellation_result(payload: JsonObject | None) -> CancellationResult:
|
||||
return CancellationResult(cancelled=bool((payload or {}).get("cancelled", False)))
|
||||
|
||||
|
||||
def parse_branch_result(payload: JsonObject | None) -> BranchResult:
|
||||
payload = payload or {}
|
||||
return BranchResult(text=str(payload.get("text", "")), cancelled=bool(payload.get("cancelled", False)))
|
||||
|
||||
|
||||
def parse_branch_messages(payload: JsonObject | None) -> tuple[BranchMessage, ...]:
|
||||
messages = cast(list[Any], (payload or {}).get("messages") or [])
|
||||
return tuple(
|
||||
BranchMessage(entry_id=str(item.get("entryId", "")), text=str(item.get("text", ""))) for item in messages
|
||||
)
|
||||
|
||||
|
||||
def parse_session_stats(payload: JsonObject) -> SessionStats:
|
||||
tokens_payload = cast(dict[str, Any], payload.get("tokens") or {})
|
||||
return SessionStats(
|
||||
session_file=str(payload["sessionFile"]) if payload.get("sessionFile") is not None else None,
|
||||
session_id=str(payload.get("sessionId", "")),
|
||||
user_messages=int(payload.get("userMessages", 0)),
|
||||
assistant_messages=int(payload.get("assistantMessages", 0)),
|
||||
tool_calls=int(payload.get("toolCalls", 0)),
|
||||
tool_results=int(payload.get("toolResults", 0)),
|
||||
total_messages=int(payload.get("totalMessages", 0)),
|
||||
tokens=TokenUsage(
|
||||
input=int(tokens_payload.get("input", 0)),
|
||||
output=int(tokens_payload.get("output", 0)),
|
||||
cache_read=int(tokens_payload.get("cacheRead", 0)),
|
||||
cache_write=int(tokens_payload.get("cacheWrite", 0)),
|
||||
total=int(tokens_payload.get("total", 0)),
|
||||
),
|
||||
premium_requests=int(payload.get("premiumRequests", 0)),
|
||||
cost=float(payload.get("cost", 0.0)),
|
||||
)
|
||||
|
||||
|
||||
def parse_extension_ui_request(payload: JsonObject) -> ExtensionUiRequest:
|
||||
return ExtensionUiRequest(
|
||||
id=str(payload["id"]),
|
||||
method=str(payload["method"]),
|
||||
title=str(payload["title"]) if payload.get("title") is not None else None,
|
||||
options=tuple(str(item) for item in cast(list[Any], payload.get("options") or [])) or None,
|
||||
message=str(payload["message"]) if payload.get("message") is not None else None,
|
||||
placeholder=str(payload["placeholder"]) if payload.get("placeholder") is not None else None,
|
||||
prefill=str(payload["prefill"]) if payload.get("prefill") is not None else None,
|
||||
timeout=int(payload["timeout"]) if payload.get("timeout") is not None else None,
|
||||
prompt_style=bool(payload["promptStyle"]) if "promptStyle" in payload else None,
|
||||
target_id=str(payload["targetId"]) if payload.get("targetId") is not None else None,
|
||||
notify_type=cast(NotifyType | None, payload.get("notifyType")),
|
||||
status_key=str(payload["statusKey"]) if payload.get("statusKey") is not None else None,
|
||||
status_text=str(payload["statusText"]) if payload.get("statusText") is not None else None,
|
||||
widget_key=str(payload["widgetKey"]) if payload.get("widgetKey") is not None else None,
|
||||
widget_lines=tuple(str(item) for item in cast(list[Any], payload.get("widgetLines") or [])) or None,
|
||||
widget_placement=cast(WidgetPlacement | None, payload.get("widgetPlacement")),
|
||||
text=str(payload["text"]) if payload.get("text") is not None else None,
|
||||
)
|
||||
|
||||
|
||||
def parse_extension_error(payload: JsonObject) -> ExtensionError:
|
||||
return ExtensionError(
|
||||
extension_path=str(payload.get("extensionPath", "")),
|
||||
event=str(payload.get("event", "")),
|
||||
error=str(payload.get("error", "")),
|
||||
)
|
||||
|
||||
|
||||
def parse_notification(payload: JsonObject) -> RpcNotification:
|
||||
event_type = payload.get("type")
|
||||
if event_type == "ready":
|
||||
return ReadyEvent()
|
||||
if event_type == "extension_ui_request":
|
||||
return parse_extension_ui_request(payload)
|
||||
if event_type == "extension_error":
|
||||
return parse_extension_error(payload)
|
||||
if event_type == "agent_start":
|
||||
return AgentStartEvent()
|
||||
if event_type == "agent_end":
|
||||
return AgentEndEvent(messages=tuple(cast(list[AgentMessage], payload.get("messages") or [])))
|
||||
if event_type == "turn_start":
|
||||
return TurnStartEvent()
|
||||
if event_type == "turn_end":
|
||||
return TurnEndEvent(
|
||||
message=cast(AgentMessage, payload["message"]),
|
||||
tool_results=tuple(cast(list[ToolResultMessage], payload.get("toolResults") or [])),
|
||||
)
|
||||
if event_type == "message_start":
|
||||
return MessageStartEvent(message=cast(AgentMessage, payload["message"]))
|
||||
if event_type == "message_update":
|
||||
return MessageUpdateEvent(
|
||||
message=cast(AgentMessage, payload["message"]),
|
||||
assistant_message_event=cast(AssistantMessageEvent, payload["assistantMessageEvent"]),
|
||||
)
|
||||
if event_type == "message_end":
|
||||
return MessageEndEvent(message=cast(AgentMessage, payload["message"]))
|
||||
if event_type == "tool_execution_start":
|
||||
return ToolExecutionStartEvent(
|
||||
tool_call_id=str(payload.get("toolCallId", "")),
|
||||
tool_name=str(payload.get("toolName", "")),
|
||||
args=cast(JsonValue, payload.get("args")),
|
||||
intent=str(payload["intent"]) if payload.get("intent") is not None else None,
|
||||
)
|
||||
if event_type == "tool_execution_update":
|
||||
return ToolExecutionUpdateEvent(
|
||||
tool_call_id=str(payload.get("toolCallId", "")),
|
||||
tool_name=str(payload.get("toolName", "")),
|
||||
args=cast(JsonValue, payload.get("args")),
|
||||
partial_result=cast(JsonValue, payload.get("partialResult")),
|
||||
)
|
||||
if event_type == "tool_execution_end":
|
||||
return ToolExecutionEndEvent(
|
||||
tool_call_id=str(payload.get("toolCallId", "")),
|
||||
tool_name=str(payload.get("toolName", "")),
|
||||
result=cast(JsonValue, payload.get("result")),
|
||||
is_error=bool(payload["isError"]) if "isError" in payload else None,
|
||||
)
|
||||
if event_type == "auto_compaction_start":
|
||||
return AutoCompactionStartEvent(
|
||||
reason=cast(Literal["threshold", "overflow", "idle"], payload.get("reason", "threshold")),
|
||||
action=cast(Literal["context-full", "handoff"], payload.get("action", "context-full")),
|
||||
)
|
||||
if event_type == "auto_compaction_end":
|
||||
result_payload = payload.get("result")
|
||||
return AutoCompactionEndEvent(
|
||||
action=cast(Literal["context-full", "handoff"], payload.get("action", "context-full")),
|
||||
result=(
|
||||
parse_compaction_result(cast(JsonObject, result_payload)) if isinstance(result_payload, dict) else None
|
||||
),
|
||||
aborted=bool(payload.get("aborted", False)),
|
||||
will_retry=bool(payload.get("willRetry", False)),
|
||||
error_message=str(payload["errorMessage"]) if payload.get("errorMessage") is not None else None,
|
||||
skipped=bool(payload["skipped"]) if "skipped" in payload else None,
|
||||
)
|
||||
if event_type == "auto_retry_start":
|
||||
return AutoRetryStartEvent(
|
||||
attempt=int(payload.get("attempt", 0)),
|
||||
max_attempts=int(payload.get("maxAttempts", 0)),
|
||||
delay_ms=int(payload.get("delayMs", 0)),
|
||||
error_message=str(payload.get("errorMessage", "")),
|
||||
)
|
||||
if event_type == "auto_retry_end":
|
||||
return AutoRetryEndEvent(
|
||||
success=bool(payload.get("success", False)),
|
||||
attempt=int(payload.get("attempt", 0)),
|
||||
final_error=str(payload["finalError"]) if payload.get("finalError") is not None else None,
|
||||
)
|
||||
if event_type == "retry_fallback_applied":
|
||||
return RetryFallbackAppliedEvent(
|
||||
from_model=str(payload.get("from", "")),
|
||||
to_model=str(payload.get("to", "")),
|
||||
role=str(payload.get("role", "")),
|
||||
)
|
||||
if event_type == "retry_fallback_succeeded":
|
||||
return RetryFallbackSucceededEvent(model=str(payload.get("model", "")), role=str(payload.get("role", "")))
|
||||
if event_type == "ttsr_triggered":
|
||||
return TtsrTriggeredEvent(rules=tuple(cast(list[JsonObject], payload.get("rules") or [])))
|
||||
if event_type == "todo_reminder":
|
||||
return TodoReminderEvent(
|
||||
todos=tuple(cast(list[JsonObject], payload.get("todos") or [])),
|
||||
attempt=int(payload.get("attempt", 0)),
|
||||
max_attempts=int(payload.get("maxAttempts", 0)),
|
||||
)
|
||||
if event_type == "todo_auto_clear":
|
||||
return TodoAutoClearEvent()
|
||||
return UnknownNotification(payload=dict(payload))
|
||||
@@ -0,0 +1 @@
|
||||
|
||||
@@ -0,0 +1,6 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src"))
|
||||
@@ -0,0 +1,220 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
import textwrap
|
||||
import unittest
|
||||
|
||||
from omp_rpc import RpcClient
|
||||
|
||||
|
||||
FAKE_SERVER = textwrap.dedent(
|
||||
"""
|
||||
import json
|
||||
import sys
|
||||
|
||||
def usage():
|
||||
return {
|
||||
"input": 1,
|
||||
"output": 1,
|
||||
"cacheRead": 0,
|
||||
"cacheWrite": 0,
|
||||
"totalTokens": 2,
|
||||
"cost": {
|
||||
"input": 0.0,
|
||||
"output": 0.0,
|
||||
"cacheRead": 0.0,
|
||||
"cacheWrite": 0.0,
|
||||
"total": 0.0,
|
||||
},
|
||||
}
|
||||
|
||||
def assistant_message(text: str):
|
||||
return {
|
||||
"role": "assistant",
|
||||
"content": [{"type": "text", "text": text}],
|
||||
"api": "anthropic-messages",
|
||||
"provider": "anthropic",
|
||||
"model": "claude-sonnet-4-5",
|
||||
"usage": usage(),
|
||||
"stopReason": "stop",
|
||||
"timestamp": 1,
|
||||
}
|
||||
|
||||
print(json.dumps({"type": "ready"}), flush=True)
|
||||
|
||||
for raw_line in sys.stdin:
|
||||
raw_line = raw_line.strip()
|
||||
if not raw_line:
|
||||
continue
|
||||
|
||||
command = json.loads(raw_line)
|
||||
command_type = command["type"]
|
||||
request_id = command.get("id")
|
||||
|
||||
if command_type == "extension_ui_response":
|
||||
print(json.dumps({"type": "agent_end", "messages": [assistant_message("ui acknowledged")]}), flush=True)
|
||||
continue
|
||||
|
||||
if command_type == "get_state":
|
||||
print(
|
||||
json.dumps(
|
||||
{
|
||||
"id": request_id,
|
||||
"type": "response",
|
||||
"command": "get_state",
|
||||
"success": True,
|
||||
"data": {
|
||||
"model": {
|
||||
"id": "claude-sonnet-4-5",
|
||||
"name": "Claude Sonnet 4.5",
|
||||
"api": "anthropic-messages",
|
||||
"provider": "anthropic",
|
||||
"baseUrl": "https://api.anthropic.com",
|
||||
"reasoning": True,
|
||||
"input": ["text"],
|
||||
"cost": {
|
||||
"input": 1.0,
|
||||
"output": 2.0,
|
||||
"cacheRead": 0.0,
|
||||
"cacheWrite": 0.0,
|
||||
},
|
||||
"contextWindow": 200000,
|
||||
"maxTokens": 8192,
|
||||
},
|
||||
"thinkingLevel": "medium",
|
||||
"isStreaming": False,
|
||||
"isCompacting": False,
|
||||
"steeringMode": "one-at-a-time",
|
||||
"followUpMode": "one-at-a-time",
|
||||
"interruptMode": "immediate",
|
||||
"sessionId": "fake-session",
|
||||
"autoCompactionEnabled": True,
|
||||
"messageCount": 0,
|
||||
"queuedMessageCount": 0,
|
||||
},
|
||||
}
|
||||
),
|
||||
flush=True,
|
||||
)
|
||||
elif command_type == "bash":
|
||||
print(
|
||||
json.dumps(
|
||||
{
|
||||
"id": request_id,
|
||||
"type": "response",
|
||||
"command": "bash",
|
||||
"success": True,
|
||||
"data": {
|
||||
"output": "hello\\n",
|
||||
"exitCode": 0,
|
||||
"cancelled": False,
|
||||
"truncated": False,
|
||||
"totalLines": 1,
|
||||
"totalBytes": 6,
|
||||
"outputLines": 1,
|
||||
"outputBytes": 6,
|
||||
},
|
||||
}
|
||||
),
|
||||
flush=True,
|
||||
)
|
||||
elif command_type == "prompt":
|
||||
print(
|
||||
json.dumps(
|
||||
{
|
||||
"id": request_id,
|
||||
"type": "response",
|
||||
"command": "prompt",
|
||||
"success": True,
|
||||
}
|
||||
),
|
||||
flush=True,
|
||||
)
|
||||
if command["message"] == "needs ui":
|
||||
print(
|
||||
json.dumps(
|
||||
{
|
||||
"type": "extension_ui_request",
|
||||
"id": "ui-1",
|
||||
"method": "input",
|
||||
"title": "Need input",
|
||||
"placeholder": "value",
|
||||
}
|
||||
),
|
||||
flush=True,
|
||||
)
|
||||
continue
|
||||
|
||||
print(json.dumps({"type": "agent_start"}), flush=True)
|
||||
partial = assistant_message("")
|
||||
print(
|
||||
json.dumps(
|
||||
{
|
||||
"type": "message_update",
|
||||
"message": partial,
|
||||
"assistantMessageEvent": {
|
||||
"type": "text_delta",
|
||||
"contentIndex": 0,
|
||||
"delta": "pong",
|
||||
"partial": partial,
|
||||
},
|
||||
}
|
||||
),
|
||||
flush=True,
|
||||
)
|
||||
assistant = assistant_message("pong")
|
||||
print(json.dumps({"type": "message_end", "message": assistant}), flush=True)
|
||||
print(json.dumps({"type": "agent_end", "messages": [assistant]}), flush=True)
|
||||
else:
|
||||
print(
|
||||
json.dumps(
|
||||
{
|
||||
"id": request_id,
|
||||
"type": "response",
|
||||
"command": command_type,
|
||||
"success": False,
|
||||
"error": f"unsupported: {command_type}",
|
||||
}
|
||||
),
|
||||
flush=True,
|
||||
)
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
class RpcClientTests(unittest.TestCase):
|
||||
def make_client(self) -> RpcClient:
|
||||
return RpcClient(
|
||||
command=[sys.executable, "-u", "-c", FAKE_SERVER],
|
||||
startup_timeout=2.0,
|
||||
request_timeout=2.0,
|
||||
)
|
||||
|
||||
def test_get_state_and_bash(self) -> None:
|
||||
with self.make_client() as client:
|
||||
state = client.get_state()
|
||||
self.assertEqual(state.session_id, "fake-session")
|
||||
self.assertEqual(state.model.id if state.model else None, "claude-sonnet-4-5")
|
||||
|
||||
result = client.bash("echo hello")
|
||||
self.assertEqual(result.output, "hello\n")
|
||||
self.assertEqual(result.exit_code, 0)
|
||||
|
||||
def test_prompt_and_wait_returns_assistant_text(self) -> None:
|
||||
with self.make_client() as client:
|
||||
turn = client.prompt_and_wait("say hello", timeout=2.0)
|
||||
self.assertEqual(turn.require_assistant_text(), "pong")
|
||||
self.assertGreaterEqual(len(turn.events), 3)
|
||||
|
||||
def test_extension_ui_round_trip(self) -> None:
|
||||
with self.make_client() as client:
|
||||
client.prompt("needs ui")
|
||||
request = client.next_ui_request(timeout=2.0)
|
||||
self.assertEqual(request.method, "input")
|
||||
|
||||
client.send_ui_value(request.id, "approved")
|
||||
client.wait_for_idle(timeout=2.0)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,123 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import unittest
|
||||
|
||||
from omp_rpc import (
|
||||
AgentEndEvent,
|
||||
ExtensionUiRequest,
|
||||
SessionState,
|
||||
assistant_text,
|
||||
parse_notification,
|
||||
parse_session_state,
|
||||
)
|
||||
|
||||
|
||||
class ProtocolParsingTests(unittest.TestCase):
|
||||
def test_parse_session_state(self) -> None:
|
||||
state = parse_session_state(
|
||||
{
|
||||
"model": {
|
||||
"id": "claude-sonnet-4-5",
|
||||
"name": "Claude Sonnet 4.5",
|
||||
"api": "anthropic-messages",
|
||||
"provider": "anthropic",
|
||||
"baseUrl": "https://api.anthropic.com",
|
||||
"reasoning": True,
|
||||
"input": ["text", "image"],
|
||||
"cost": {
|
||||
"input": 1.0,
|
||||
"output": 2.0,
|
||||
"cacheRead": 0.1,
|
||||
"cacheWrite": 0.2,
|
||||
},
|
||||
"contextWindow": 200000,
|
||||
"maxTokens": 8192,
|
||||
"thinking": {
|
||||
"minLevel": "minimal",
|
||||
"maxLevel": "high",
|
||||
"mode": "effort",
|
||||
},
|
||||
},
|
||||
"thinkingLevel": "medium",
|
||||
"isStreaming": False,
|
||||
"isCompacting": False,
|
||||
"steeringMode": "one-at-a-time",
|
||||
"followUpMode": "all",
|
||||
"interruptMode": "immediate",
|
||||
"sessionFile": "/tmp/test.jsonl",
|
||||
"sessionId": "session-123",
|
||||
"sessionName": "Scratchpad",
|
||||
"autoCompactionEnabled": True,
|
||||
"messageCount": 4,
|
||||
"queuedMessageCount": 1,
|
||||
"systemPrompt": "You are useful.",
|
||||
"dumpTools": [
|
||||
{
|
||||
"name": "read",
|
||||
"description": "Read files",
|
||||
"parameters": {"type": "object"},
|
||||
}
|
||||
],
|
||||
}
|
||||
)
|
||||
|
||||
self.assertIsInstance(state, SessionState)
|
||||
self.assertEqual(state.session_id, "session-123")
|
||||
self.assertEqual(state.follow_up_mode, "all")
|
||||
self.assertEqual(state.model.id if state.model else None, "claude-sonnet-4-5")
|
||||
self.assertEqual(state.dump_tools[0].name, "read")
|
||||
|
||||
def test_parse_agent_end_notification(self) -> None:
|
||||
notification = parse_notification(
|
||||
{
|
||||
"type": "agent_end",
|
||||
"messages": [
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [{"type": "text", "text": "hello"}],
|
||||
"api": "anthropic-messages",
|
||||
"provider": "anthropic",
|
||||
"model": "claude-sonnet-4-5",
|
||||
"usage": {
|
||||
"input": 1,
|
||||
"output": 1,
|
||||
"cacheRead": 0,
|
||||
"cacheWrite": 0,
|
||||
"totalTokens": 2,
|
||||
"cost": {
|
||||
"input": 0.0,
|
||||
"output": 0.0,
|
||||
"cacheRead": 0.0,
|
||||
"cacheWrite": 0.0,
|
||||
"total": 0.0,
|
||||
},
|
||||
},
|
||||
"stopReason": "stop",
|
||||
"timestamp": 1,
|
||||
}
|
||||
],
|
||||
}
|
||||
)
|
||||
|
||||
self.assertIsInstance(notification, AgentEndEvent)
|
||||
self.assertEqual(assistant_text(notification.messages[0]), "hello")
|
||||
|
||||
def test_parse_extension_ui_request(self) -> None:
|
||||
notification = parse_notification(
|
||||
{
|
||||
"type": "extension_ui_request",
|
||||
"id": "ui-1",
|
||||
"method": "confirm",
|
||||
"title": "Confirm",
|
||||
"message": "Continue?",
|
||||
"timeout": 1000,
|
||||
}
|
||||
)
|
||||
|
||||
self.assertIsInstance(notification, ExtensionUiRequest)
|
||||
self.assertEqual(notification.method, "confirm")
|
||||
self.assertEqual(notification.message, "Continue?")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user