feat(omp-rpc): introduced host tool execution framework with custom tool registration
- Added host tool execution framework with HostTool, HostToolContext, and host_tool() factory for custom tool integration. - Added RpcConcurrencyError exception and _PromptLifecycleCoordinator to enforce single-flight constraint on prompt lifecycle methods. - Enhanced JSON parsing with 10 validation helpers and enum frozensets for safe field extraction with detailed error messages. - Replaced manual event/error list management with _BoundedHistory for bounded-size history with offset tracking. - Added deep JSON cloning to prevent external mutations of stored payloads and improved UTF-8 error handling in subprocess stderr. - Added custom_tools parameter to RpcClient and set_custom_tools() method for runtime tool registration.
This commit is contained in:
+632
-124
@@ -2,15 +2,18 @@ from __future__ import annotations
|
||||
|
||||
import sys
|
||||
import textwrap
|
||||
import threading
|
||||
import time
|
||||
import unittest
|
||||
|
||||
from omp_rpc import RpcClient, RpcCommandError, RpcError
|
||||
from omp_rpc import RpcClient, RpcCommandError, RpcConcurrencyError, RpcError, host_tool
|
||||
|
||||
|
||||
FAKE_SERVER = textwrap.dedent(
|
||||
"""
|
||||
import json
|
||||
import sys
|
||||
import time
|
||||
|
||||
def usage():
|
||||
return {
|
||||
@@ -28,20 +31,195 @@ FAKE_SERVER = textwrap.dedent(
|
||||
},
|
||||
}
|
||||
|
||||
def model_info(model_id: str, provider: str = "anthropic"):
|
||||
return {
|
||||
"id": model_id,
|
||||
"name": f"Model {model_id}",
|
||||
"api": "anthropic-messages",
|
||||
"provider": provider,
|
||||
"baseUrl": "https://api.anthropic.com",
|
||||
"reasoning": True,
|
||||
"input": ["text"],
|
||||
"cost": {
|
||||
"input": 1.0,
|
||||
"output": 2.0,
|
||||
"cacheRead": 0.0,
|
||||
"cacheWrite": 0.0,
|
||||
},
|
||||
"contextWindow": 200000,
|
||||
"maxTokens": 8192,
|
||||
}
|
||||
|
||||
def assistant_message(text: str):
|
||||
return {
|
||||
"role": "assistant",
|
||||
"content": [{"type": "text", "text": text}],
|
||||
"api": "anthropic-messages",
|
||||
"provider": "anthropic",
|
||||
"model": "claude-sonnet-4-5",
|
||||
"provider": model_provider,
|
||||
"model": model_id,
|
||||
"usage": usage(),
|
||||
"stopReason": "stop",
|
||||
"timestamp": 1,
|
||||
}
|
||||
|
||||
registered_host_tools = []
|
||||
|
||||
def current_state():
|
||||
return {
|
||||
"model": model_info(model_id, model_provider),
|
||||
"thinkingLevel": thinking_level,
|
||||
"isStreaming": False,
|
||||
"isCompacting": False,
|
||||
"steeringMode": steering_mode,
|
||||
"followUpMode": follow_up_mode,
|
||||
"interruptMode": interrupt_mode,
|
||||
"sessionId": "fake-session",
|
||||
"sessionName": session_name,
|
||||
"autoCompactionEnabled": auto_compaction_enabled,
|
||||
"messageCount": len(messages),
|
||||
"queuedMessageCount": 0,
|
||||
"todoPhases": todo_phases,
|
||||
"dumpTools": [{"name": "read", "description": "Read files", "parameters": {"type": "object"}}] + registered_host_tools,
|
||||
}
|
||||
|
||||
def emit_prompt_turn(text: str, delay: float = 0.0, include_extra_events: bool = False):
|
||||
global last_assistant_text, messages
|
||||
print(json.dumps({"type": "agent_start"}), flush=True)
|
||||
print(json.dumps({"type": "turn_start"}), flush=True)
|
||||
partial = assistant_message("")
|
||||
print(json.dumps({"type": "message_start", "message": partial}), flush=True)
|
||||
print(
|
||||
json.dumps(
|
||||
{
|
||||
"type": "message_update",
|
||||
"message": partial,
|
||||
"assistantMessageEvent": {
|
||||
"type": "text_delta",
|
||||
"contentIndex": 0,
|
||||
"delta": text,
|
||||
"partial": partial,
|
||||
},
|
||||
}
|
||||
),
|
||||
flush=True,
|
||||
)
|
||||
|
||||
if delay:
|
||||
time.sleep(delay)
|
||||
|
||||
if include_extra_events:
|
||||
print(
|
||||
json.dumps(
|
||||
{
|
||||
"type": "tool_execution_start",
|
||||
"toolCallId": "tool-1",
|
||||
"toolName": "read",
|
||||
"args": {"path": "README.md"},
|
||||
"intent": "Inspect docs",
|
||||
}
|
||||
),
|
||||
flush=True,
|
||||
)
|
||||
print(
|
||||
json.dumps(
|
||||
{
|
||||
"type": "tool_execution_update",
|
||||
"toolCallId": "tool-1",
|
||||
"toolName": "read",
|
||||
"args": {"path": "README.md"},
|
||||
"partialResult": {"bytes": 12},
|
||||
}
|
||||
),
|
||||
flush=True,
|
||||
)
|
||||
print(
|
||||
json.dumps(
|
||||
{
|
||||
"type": "tool_execution_end",
|
||||
"toolCallId": "tool-1",
|
||||
"toolName": "read",
|
||||
"result": {"text": "docs"},
|
||||
"isError": False,
|
||||
}
|
||||
),
|
||||
flush=True,
|
||||
)
|
||||
print(json.dumps({"type": "auto_compaction_start", "reason": "threshold", "action": "context-full"}), flush=True)
|
||||
print(
|
||||
json.dumps(
|
||||
{
|
||||
"type": "auto_compaction_end",
|
||||
"action": "context-full",
|
||||
"result": {
|
||||
"summary": "trimmed",
|
||||
"shortSummary": "trimmed",
|
||||
"firstKeptEntryId": "entry-1",
|
||||
"tokensBefore": 123,
|
||||
},
|
||||
"aborted": False,
|
||||
"willRetry": False,
|
||||
}
|
||||
),
|
||||
flush=True,
|
||||
)
|
||||
print(
|
||||
json.dumps(
|
||||
{
|
||||
"type": "auto_retry_start",
|
||||
"attempt": 1,
|
||||
"maxAttempts": 3,
|
||||
"delayMs": 25,
|
||||
"errorMessage": "retrying",
|
||||
}
|
||||
),
|
||||
flush=True,
|
||||
)
|
||||
print(json.dumps({"type": "auto_retry_end", "success": True, "attempt": 1}), flush=True)
|
||||
print(json.dumps({"type": "retry_fallback_applied", "from": "a", "to": "b", "role": "primary"}), flush=True)
|
||||
print(json.dumps({"type": "retry_fallback_succeeded", "model": "b", "role": "primary"}), flush=True)
|
||||
print(json.dumps({"type": "ttsr_triggered", "rules": [{"id": "rule-1"}]}), flush=True)
|
||||
print(
|
||||
json.dumps(
|
||||
{
|
||||
"type": "todo_reminder",
|
||||
"attempt": 1,
|
||||
"maxAttempts": 2,
|
||||
"todos": [{"id": "task-1", "content": "Map tools", "status": "pending"}],
|
||||
}
|
||||
),
|
||||
flush=True,
|
||||
)
|
||||
print(json.dumps({"type": "todo_auto_clear"}), flush=True)
|
||||
|
||||
assistant = assistant_message(text)
|
||||
print(json.dumps({"type": "message_end", "message": assistant}), flush=True)
|
||||
print(json.dumps({"type": "turn_end", "message": assistant, "toolResults": []}), flush=True)
|
||||
print(json.dumps({"type": "agent_end", "messages": [assistant]}), flush=True)
|
||||
last_assistant_text = text
|
||||
messages = [assistant]
|
||||
|
||||
def respond(request_id, command, data=None, success=True, error=None):
|
||||
payload = {"id": request_id, "type": "response", "command": command, "success": success}
|
||||
if success and data is not None:
|
||||
payload["data"] = data
|
||||
if not success:
|
||||
payload["error"] = error
|
||||
print(json.dumps(payload), flush=True)
|
||||
|
||||
print(json.dumps({"type": "ready"}), flush=True)
|
||||
todo_phases = []
|
||||
messages = []
|
||||
branch_messages = [{"entryId": "entry-1", "text": "branch message"}]
|
||||
model_provider = "anthropic"
|
||||
model_id = "claude-sonnet-4-5"
|
||||
thinking_level = "medium"
|
||||
steering_mode = "one-at-a-time"
|
||||
follow_up_mode = "one-at-a-time"
|
||||
interrupt_mode = "immediate"
|
||||
auto_compaction_enabled = True
|
||||
auto_retry_enabled = True
|
||||
session_name = "Scratchpad"
|
||||
last_assistant_text = None
|
||||
|
||||
for raw_line in sys.stdin:
|
||||
raw_line = raw_line.strip()
|
||||
@@ -53,151 +231,185 @@ FAKE_SERVER = textwrap.dedent(
|
||||
request_id = command.get("id")
|
||||
|
||||
if command_type == "extension_ui_response":
|
||||
print(json.dumps({"type": "agent_end", "messages": [assistant_message("ui acknowledged")]}), flush=True)
|
||||
emit_prompt_turn("ui acknowledged")
|
||||
continue
|
||||
|
||||
if command_type == "get_state":
|
||||
print(
|
||||
json.dumps(
|
||||
{
|
||||
"id": request_id,
|
||||
"type": "response",
|
||||
"command": "get_state",
|
||||
"success": True,
|
||||
"data": {
|
||||
"model": {
|
||||
"id": "claude-sonnet-4-5",
|
||||
"name": "Claude Sonnet 4.5",
|
||||
"api": "anthropic-messages",
|
||||
"provider": "anthropic",
|
||||
"baseUrl": "https://api.anthropic.com",
|
||||
"reasoning": True,
|
||||
"input": ["text"],
|
||||
"cost": {
|
||||
"input": 1.0,
|
||||
"output": 2.0,
|
||||
"cacheRead": 0.0,
|
||||
"cacheWrite": 0.0,
|
||||
},
|
||||
"contextWindow": 200000,
|
||||
"maxTokens": 8192,
|
||||
},
|
||||
"thinkingLevel": "medium",
|
||||
"isStreaming": False,
|
||||
"isCompacting": False,
|
||||
"steeringMode": "one-at-a-time",
|
||||
"followUpMode": "one-at-a-time",
|
||||
"interruptMode": "immediate",
|
||||
"sessionId": "fake-session",
|
||||
"autoCompactionEnabled": True,
|
||||
"messageCount": 0,
|
||||
"queuedMessageCount": 0,
|
||||
"todoPhases": todo_phases,
|
||||
},
|
||||
}
|
||||
),
|
||||
flush=True,
|
||||
respond(request_id, "get_state", current_state())
|
||||
elif command_type == "set_host_tools":
|
||||
registered_host_tools = command.get("tools", [])
|
||||
respond(
|
||||
request_id,
|
||||
"set_host_tools",
|
||||
{"toolNames": [tool.get("name", "") for tool in registered_host_tools]},
|
||||
)
|
||||
elif command_type == "set_todos":
|
||||
todo_phases = command.get("phases", [])
|
||||
print(
|
||||
json.dumps(
|
||||
{
|
||||
"id": request_id,
|
||||
"type": "response",
|
||||
"command": "set_todos",
|
||||
"success": True,
|
||||
"data": {
|
||||
"todoPhases": todo_phases,
|
||||
},
|
||||
}
|
||||
),
|
||||
flush=True,
|
||||
respond(request_id, "set_todos", {"todoPhases": todo_phases})
|
||||
elif command_type == "get_messages":
|
||||
respond(request_id, "get_messages", {"messages": messages})
|
||||
elif command_type == "set_host_tools":
|
||||
tool_names = [tool.get("name", "") for tool in command.get("tools", [])]
|
||||
respond(request_id, "set_host_tools", {"toolNames": tool_names})
|
||||
elif command_type == "set_model":
|
||||
model_provider = command["provider"]
|
||||
model_id = command["modelId"]
|
||||
respond(request_id, "set_model", model_info(model_id, model_provider))
|
||||
elif command_type == "cycle_model":
|
||||
model_id = "claude-sonnet-4-6" if model_id == "claude-sonnet-4-5" else "claude-sonnet-4-5"
|
||||
respond(request_id, "cycle_model", {"model": model_info(model_id, model_provider), "thinkingLevel": thinking_level, "isScoped": False})
|
||||
elif command_type == "get_available_models":
|
||||
respond(
|
||||
request_id,
|
||||
"get_available_models",
|
||||
{
|
||||
"models": [
|
||||
model_info("claude-sonnet-4-5", "anthropic"),
|
||||
model_info("claude-sonnet-4-6", "anthropic"),
|
||||
]
|
||||
},
|
||||
)
|
||||
elif command_type == "set_thinking_level":
|
||||
thinking_level = command["level"]
|
||||
respond(request_id, "set_thinking_level", {})
|
||||
elif command_type == "cycle_thinking_level":
|
||||
thinking_level = "high" if thinking_level != "high" else "low"
|
||||
respond(request_id, "cycle_thinking_level", {"level": thinking_level})
|
||||
elif command_type == "set_steering_mode":
|
||||
steering_mode = command["mode"]
|
||||
respond(request_id, "set_steering_mode", {})
|
||||
elif command_type == "set_follow_up_mode":
|
||||
follow_up_mode = command["mode"]
|
||||
respond(request_id, "set_follow_up_mode", {})
|
||||
elif command_type == "set_interrupt_mode":
|
||||
interrupt_mode = command["mode"]
|
||||
respond(request_id, "set_interrupt_mode", {})
|
||||
elif command_type == "compact":
|
||||
respond(
|
||||
request_id,
|
||||
"compact",
|
||||
{"summary": "trimmed", "shortSummary": "trimmed", "firstKeptEntryId": "entry-1", "tokensBefore": 123},
|
||||
)
|
||||
elif command_type == "set_auto_compaction":
|
||||
auto_compaction_enabled = command["enabled"]
|
||||
respond(request_id, "set_auto_compaction", {})
|
||||
elif command_type == "set_auto_retry":
|
||||
auto_retry_enabled = command["enabled"]
|
||||
respond(request_id, "set_auto_retry", {})
|
||||
elif command_type == "abort_retry":
|
||||
respond(request_id, "abort_retry", {})
|
||||
elif command_type == "bash":
|
||||
print(
|
||||
json.dumps(
|
||||
{
|
||||
"id": request_id,
|
||||
"type": "response",
|
||||
"command": "bash",
|
||||
"success": True,
|
||||
"data": {
|
||||
"output": "hello\\n",
|
||||
"exitCode": 0,
|
||||
"cancelled": False,
|
||||
"truncated": False,
|
||||
"totalLines": 1,
|
||||
"totalBytes": 6,
|
||||
"outputLines": 1,
|
||||
"outputBytes": 6,
|
||||
},
|
||||
}
|
||||
),
|
||||
flush=True,
|
||||
respond(
|
||||
request_id,
|
||||
"bash",
|
||||
{
|
||||
"output": "hello\\n",
|
||||
"exitCode": 0,
|
||||
"cancelled": False,
|
||||
"truncated": False,
|
||||
"totalLines": 1,
|
||||
"totalBytes": 6,
|
||||
"outputLines": 1,
|
||||
"outputBytes": 6,
|
||||
},
|
||||
)
|
||||
elif command_type == "prompt":
|
||||
print(
|
||||
json.dumps(
|
||||
{
|
||||
"id": request_id,
|
||||
"type": "response",
|
||||
"command": "prompt",
|
||||
"success": True,
|
||||
}
|
||||
),
|
||||
flush=True,
|
||||
elif command_type == "abort_bash":
|
||||
respond(request_id, "abort_bash", {})
|
||||
elif command_type == "get_session_stats":
|
||||
respond(
|
||||
request_id,
|
||||
"get_session_stats",
|
||||
{
|
||||
"sessionFile": "/tmp/fake-session.jsonl",
|
||||
"sessionId": "fake-session",
|
||||
"userMessages": 1,
|
||||
"assistantMessages": len(messages),
|
||||
"toolCalls": 1,
|
||||
"toolResults": 1,
|
||||
"totalMessages": len(messages) + 1,
|
||||
"tokens": {"input": 10, "output": 5, "cacheRead": 0, "cacheWrite": 0, "total": 15},
|
||||
"premiumRequests": 0,
|
||||
"cost": 0.0,
|
||||
},
|
||||
)
|
||||
if command["message"] == "needs ui":
|
||||
elif command_type == "export_html":
|
||||
respond(request_id, "export_html", {"path": command.get("outputPath") or "/tmp/session.html"})
|
||||
elif command_type == "new_session":
|
||||
respond(request_id, "new_session", {"cancelled": False})
|
||||
elif command_type == "switch_session":
|
||||
respond(request_id, "switch_session", {"cancelled": False})
|
||||
elif command_type == "branch":
|
||||
branch_messages = [{"entryId": command["entryId"], "text": "branch message"}]
|
||||
respond(request_id, "branch", {"text": "branch created", "cancelled": False})
|
||||
elif command_type == "get_branch_messages":
|
||||
respond(request_id, "get_branch_messages", {"messages": branch_messages})
|
||||
elif command_type == "get_last_assistant_text":
|
||||
respond(request_id, "get_last_assistant_text", {"text": last_assistant_text})
|
||||
elif command_type == "set_session_name":
|
||||
session_name = command["name"]
|
||||
respond(request_id, "set_session_name", {})
|
||||
elif command_type in {"steer", "follow_up", "abort"}:
|
||||
respond(request_id, command_type, {})
|
||||
elif command_type in {"prompt", "abort_and_prompt"}:
|
||||
respond(request_id, command_type, {})
|
||||
message = command["message"]
|
||||
if message == "needs ui":
|
||||
print(json.dumps({"type": "extension_ui_request", "id": "ui-1", "method": "input", "title": "Need input", "placeholder": "value"}), flush=True)
|
||||
continue
|
||||
if message == "needs confirm":
|
||||
print(json.dumps({"type": "extension_ui_request", "id": "ui-2", "method": "confirm", "title": "Confirm", "message": "Continue?"}), flush=True)
|
||||
continue
|
||||
if message == "needs cancel":
|
||||
print(json.dumps({"type": "extension_ui_request", "id": "ui-3", "method": "editor", "title": "Edit", "placeholder": "value"}), flush=True)
|
||||
continue
|
||||
if message == "needs host tool":
|
||||
print(json.dumps({"type": "agent_start"}), flush=True)
|
||||
print(
|
||||
json.dumps(
|
||||
{
|
||||
"type": "extension_ui_request",
|
||||
"id": "ui-1",
|
||||
"method": "input",
|
||||
"title": "Need input",
|
||||
"placeholder": "value",
|
||||
"type": "host_tool_call",
|
||||
"id": "host-call-1",
|
||||
"toolCallId": "toolu_host_1",
|
||||
"toolName": "echo_host",
|
||||
"arguments": {"message": "hello"},
|
||||
}
|
||||
),
|
||||
flush=True,
|
||||
)
|
||||
continue
|
||||
|
||||
print(json.dumps({"type": "agent_start"}), flush=True)
|
||||
print(json.dumps({"type": "turn_start"}), flush=True)
|
||||
partial = assistant_message("")
|
||||
if message == "notifications":
|
||||
print(json.dumps({"type": "extension_error", "extensionPath": "/tmp/ext.py", "event": "run", "error": "boom"}), flush=True)
|
||||
print(json.dumps({"type": "unknown_future_event", "value": 1}), flush=True)
|
||||
emit_prompt_turn("pong", delay=0.3 if message == "slow" else 0.0, include_extra_events=message == "all events")
|
||||
elif command_type == "host_tool_update":
|
||||
print(
|
||||
json.dumps(
|
||||
{
|
||||
"type": "message_update",
|
||||
"message": partial,
|
||||
"assistantMessageEvent": {
|
||||
"type": "text_delta",
|
||||
"contentIndex": 0,
|
||||
"delta": "pong",
|
||||
"partial": partial,
|
||||
},
|
||||
"type": "tool_execution_update",
|
||||
"toolCallId": "toolu_host_1",
|
||||
"toolName": "echo_host",
|
||||
"args": {"message": "hello"},
|
||||
"partialResult": command["partialResult"],
|
||||
}
|
||||
),
|
||||
flush=True,
|
||||
)
|
||||
assistant = assistant_message("pong")
|
||||
print(json.dumps({"type": "message_end", "message": assistant}), flush=True)
|
||||
print(json.dumps({"type": "turn_end", "message": assistant, "toolResults": []}), flush=True)
|
||||
print(json.dumps({"type": "agent_end", "messages": [assistant]}), flush=True)
|
||||
elif command_type == "host_tool_result":
|
||||
print(
|
||||
json.dumps(
|
||||
{
|
||||
"type": "tool_execution_end",
|
||||
"toolCallId": "toolu_host_1",
|
||||
"toolName": "echo_host",
|
||||
"result": command["result"],
|
||||
"isError": command.get("isError", False),
|
||||
}
|
||||
),
|
||||
flush=True,
|
||||
)
|
||||
print(json.dumps({"type": "agent_end", "messages": []}), flush=True)
|
||||
else:
|
||||
print(
|
||||
json.dumps(
|
||||
{
|
||||
"id": request_id,
|
||||
"type": "response",
|
||||
"command": command_type,
|
||||
"success": False,
|
||||
"error": f"unsupported: {command_type}",
|
||||
}
|
||||
),
|
||||
flush=True,
|
||||
)
|
||||
respond(request_id, command_type, success=False, error=f"unsupported: {command_type}")
|
||||
"""
|
||||
)
|
||||
|
||||
@@ -214,6 +426,20 @@ IDLESS_ERROR_SERVER = textwrap.dedent(
|
||||
continue
|
||||
|
||||
command = json.loads(raw_line)
|
||||
if command["type"] == "set_host_tools":
|
||||
print(
|
||||
json.dumps(
|
||||
{
|
||||
"id": command.get("id"),
|
||||
"type": "response",
|
||||
"command": "set_host_tools",
|
||||
"success": True,
|
||||
"data": {"toolNames": []},
|
||||
}
|
||||
),
|
||||
flush=True,
|
||||
)
|
||||
continue
|
||||
print(
|
||||
json.dumps(
|
||||
{
|
||||
@@ -242,6 +468,20 @@ LATE_PROMPT_FAILURE_SERVER = textwrap.dedent(
|
||||
|
||||
command = json.loads(raw_line)
|
||||
request_id = command.get("id")
|
||||
if command["type"] == "set_host_tools":
|
||||
print(
|
||||
json.dumps(
|
||||
{
|
||||
"id": request_id,
|
||||
"type": "response",
|
||||
"command": "set_host_tools",
|
||||
"success": True,
|
||||
"data": {"toolNames": []},
|
||||
}
|
||||
),
|
||||
flush=True,
|
||||
)
|
||||
continue
|
||||
if command["type"] == "prompt":
|
||||
print(
|
||||
json.dumps(
|
||||
@@ -292,8 +532,44 @@ STDERR_SERVER = textwrap.dedent(
|
||||
sys.stderr.flush()
|
||||
print(json.dumps({"type": "ready"}), flush=True)
|
||||
|
||||
for _ in sys.stdin:
|
||||
pass
|
||||
for raw_line in sys.stdin:
|
||||
raw_line = raw_line.strip()
|
||||
if not raw_line:
|
||||
continue
|
||||
command = json.loads(raw_line)
|
||||
if command["type"] == "set_host_tools":
|
||||
print(
|
||||
json.dumps(
|
||||
{
|
||||
"id": command.get("id"),
|
||||
"type": "response",
|
||||
"command": "set_host_tools",
|
||||
"success": True,
|
||||
"data": {"toolNames": []},
|
||||
}
|
||||
),
|
||||
flush=True,
|
||||
)
|
||||
"""
|
||||
)
|
||||
|
||||
INVALID_JSON_SERVER = textwrap.dedent(
|
||||
"""
|
||||
import sys
|
||||
|
||||
sys.stdout.write('{"type":"ready"}\\n')
|
||||
sys.stdout.flush()
|
||||
sys.stdout.write('{"type":"broken"\\n')
|
||||
sys.stdout.flush()
|
||||
"""
|
||||
)
|
||||
|
||||
BROKEN_STARTUP_SERVER = textwrap.dedent(
|
||||
"""
|
||||
import sys
|
||||
|
||||
sys.stdout.write('not-json\\n')
|
||||
sys.stdout.flush()
|
||||
"""
|
||||
)
|
||||
|
||||
@@ -358,6 +634,38 @@ class RpcClientTests(unittest.TestCase):
|
||||
self.assertEqual(turn.require_assistant_text(), "pong")
|
||||
self.assertGreaterEqual(len(turn.events), 3)
|
||||
|
||||
def test_custom_tools_are_registered_and_executed_via_rpc(self) -> None:
|
||||
def echo_host(args: dict[str, str], context) -> str:
|
||||
context.send_update(f"working:{args['message']}")
|
||||
return f"host:{args['message']}"
|
||||
|
||||
with self.make_client(
|
||||
custom_tools=(
|
||||
host_tool(
|
||||
name="echo_host",
|
||||
description="Echo from the Python host process",
|
||||
parameters={
|
||||
"type": "object",
|
||||
"properties": {"message": {"type": "string"}},
|
||||
"required": ["message"],
|
||||
"additionalProperties": False,
|
||||
},
|
||||
execute=echo_host,
|
||||
),
|
||||
)
|
||||
) as client:
|
||||
state = client.get_state()
|
||||
self.assertEqual(state.dump_tools[-1].name, "echo_host")
|
||||
|
||||
turn = client.prompt_and_wait("needs host tool", timeout=2.0)
|
||||
update_events = [event for event in turn.events if getattr(event, "type", None) == "tool_execution_update"]
|
||||
end_events = [event for event in turn.events if getattr(event, "type", None) == "tool_execution_end"]
|
||||
|
||||
self.assertEqual(len(update_events), 1)
|
||||
self.assertEqual(update_events[0].partial_result["content"][0]["text"], "working:hello")
|
||||
self.assertEqual(len(end_events), 1)
|
||||
self.assertEqual(end_events[0].result["content"][0]["text"], "host:hello")
|
||||
|
||||
def test_extension_ui_round_trip(self) -> None:
|
||||
with self.make_client() as client:
|
||||
client.prompt("needs ui")
|
||||
@@ -411,6 +719,198 @@ class RpcClientTests(unittest.TestCase):
|
||||
state = client.get_state()
|
||||
self.assertEqual(state.todo_phases[0].tasks[1].content, "Exercise edits")
|
||||
|
||||
def test_model_mode_and_session_commands(self) -> None:
|
||||
with self.make_client() as client:
|
||||
model = client.set_model("anthropic", "claude-sonnet-4-6")
|
||||
self.assertEqual(model.id, "claude-sonnet-4-6")
|
||||
|
||||
cycled = client.cycle_model()
|
||||
self.assertIsNotNone(cycled)
|
||||
self.assertEqual(cycled.model.id, "claude-sonnet-4-5")
|
||||
|
||||
available = client.get_available_models()
|
||||
self.assertEqual([item.id for item in available], ["claude-sonnet-4-5", "claude-sonnet-4-6"])
|
||||
|
||||
client.set_thinking_level("high")
|
||||
self.assertEqual(client.get_state().thinking_level, "high")
|
||||
|
||||
cycled_level = client.cycle_thinking_level()
|
||||
self.assertIsNotNone(cycled_level)
|
||||
self.assertEqual(cycled_level.level, "low")
|
||||
|
||||
client.set_steering_mode("all")
|
||||
client.set_follow_up_mode("all")
|
||||
client.set_interrupt_mode("wait")
|
||||
client.set_auto_compaction(False)
|
||||
client.set_auto_retry(False)
|
||||
client.set_session_name("Renamed")
|
||||
|
||||
state = client.get_state()
|
||||
self.assertEqual(state.steering_mode, "all")
|
||||
self.assertEqual(state.follow_up_mode, "all")
|
||||
self.assertEqual(state.interrupt_mode, "wait")
|
||||
self.assertFalse(state.auto_compaction_enabled)
|
||||
self.assertEqual(state.session_name, "Renamed")
|
||||
|
||||
compacted = client.compact()
|
||||
self.assertEqual(compacted.summary, "trimmed")
|
||||
|
||||
stats = client.get_session_stats()
|
||||
self.assertEqual(stats.session_id, "fake-session")
|
||||
self.assertEqual(stats.tokens.total, 15)
|
||||
|
||||
exported = client.export_html("/tmp/custom.html")
|
||||
self.assertEqual(str(exported), "/tmp/custom.html")
|
||||
|
||||
new_session = client.new_session()
|
||||
switched = client.switch_session("/tmp/session.jsonl")
|
||||
self.assertFalse(new_session.cancelled)
|
||||
self.assertFalse(switched.cancelled)
|
||||
|
||||
branch = client.branch("entry-9")
|
||||
self.assertEqual(branch.text, "branch created")
|
||||
branch_messages = client.get_branch_messages()
|
||||
self.assertEqual(branch_messages[0].entry_id, "entry-9")
|
||||
|
||||
def test_message_and_control_commands(self) -> None:
|
||||
with self.make_client() as client:
|
||||
turn = client.prompt_and_wait("say hello", timeout=2.0)
|
||||
self.assertEqual(turn.require_assistant_text(), "pong")
|
||||
self.assertEqual(client.get_last_assistant_text(), "pong")
|
||||
|
||||
messages = client.get_messages()
|
||||
self.assertEqual(len(messages), 1)
|
||||
self.assertEqual(messages[0]["role"], "assistant")
|
||||
|
||||
client.clear_todos()
|
||||
self.assertEqual(client.get_todos(), ())
|
||||
|
||||
client.steer("nudge")
|
||||
client.follow_up("later")
|
||||
client.abort()
|
||||
client.abort_retry()
|
||||
client.abort_bash()
|
||||
|
||||
client.abort_and_prompt("say hello")
|
||||
client.wait_for_idle(timeout=2.0)
|
||||
self.assertEqual(client.get_last_assistant_text(), "pong")
|
||||
|
||||
def test_collect_events_returns_turn_events(self) -> None:
|
||||
with self.make_client() as client:
|
||||
client.prompt("slow")
|
||||
events = client.collect_events(timeout=2.0)
|
||||
|
||||
self.assertGreaterEqual(len(events), 1)
|
||||
self.assertEqual(events[-1].type, "agent_end")
|
||||
|
||||
def test_all_typed_event_listeners_receive_eventful_prompt(self) -> None:
|
||||
seen: list[str] = []
|
||||
|
||||
with self.make_client() as client:
|
||||
client.on_event(lambda event: seen.append(f"event:{event.type}"))
|
||||
client.on_agent_start(lambda event: seen.append(event.type))
|
||||
client.on_turn_end(lambda event: seen.append(event.type))
|
||||
client.on_message_start(lambda event: seen.append(event.type))
|
||||
client.on_message_end(lambda event: seen.append(event.type))
|
||||
client.on_tool_execution_start(lambda event: seen.append(event.type))
|
||||
client.on_tool_execution_update(lambda event: seen.append(event.type))
|
||||
client.on_tool_execution_end(lambda event: seen.append(event.type))
|
||||
client.on_auto_compaction_start(lambda event: seen.append(event.type))
|
||||
client.on_auto_compaction_end(lambda event: seen.append(event.type))
|
||||
client.on_auto_retry_start(lambda event: seen.append(event.type))
|
||||
client.on_auto_retry_end(lambda event: seen.append(event.type))
|
||||
client.on_retry_fallback_applied(lambda event: seen.append(event.type))
|
||||
client.on_retry_fallback_succeeded(lambda event: seen.append(event.type))
|
||||
client.on_ttsr_triggered(lambda event: seen.append(event.type))
|
||||
client.on_todo_reminder(lambda event: seen.append(event.type))
|
||||
client.on_todo_auto_clear(lambda event: seen.append(event.type))
|
||||
|
||||
turn = client.prompt_and_wait("all events", timeout=2.0)
|
||||
|
||||
self.assertEqual(turn.require_assistant_text(), "pong")
|
||||
for expected in [
|
||||
"agent_start",
|
||||
"message_start",
|
||||
"message_end",
|
||||
"turn_end",
|
||||
"tool_execution_start",
|
||||
"tool_execution_update",
|
||||
"tool_execution_end",
|
||||
"auto_compaction_start",
|
||||
"auto_compaction_end",
|
||||
"auto_retry_start",
|
||||
"auto_retry_end",
|
||||
"retry_fallback_applied",
|
||||
"retry_fallback_succeeded",
|
||||
"ttsr_triggered",
|
||||
"todo_reminder",
|
||||
"todo_auto_clear",
|
||||
]:
|
||||
self.assertIn(expected, seen)
|
||||
|
||||
def test_extension_and_unknown_notification_listeners(self) -> None:
|
||||
seen_extension_errors: list[str] = []
|
||||
seen_unknown: list[str] = []
|
||||
|
||||
with self.make_client() as client:
|
||||
client.on_extension_error(lambda event: seen_extension_errors.append(event.error))
|
||||
client.on_unknown_notification(lambda event: seen_unknown.append(str(event.payload.get("type"))))
|
||||
client.prompt_and_wait("notifications", timeout=2.0)
|
||||
|
||||
self.assertEqual(seen_extension_errors, ["boom"])
|
||||
self.assertEqual(seen_unknown, ["unknown_future_event"])
|
||||
|
||||
def test_ui_confirmation_and_cancel_round_trip(self) -> None:
|
||||
with self.make_client() as client:
|
||||
client.prompt("needs confirm")
|
||||
confirm_request = client.next_ui_request(timeout=2.0)
|
||||
self.assertEqual(confirm_request.method, "confirm")
|
||||
client.send_ui_confirmation(confirm_request.id, True)
|
||||
client.wait_for_idle(timeout=2.0)
|
||||
|
||||
client.prompt("needs cancel")
|
||||
editor_request = client.next_ui_request(timeout=2.0)
|
||||
self.assertEqual(editor_request.method, "editor")
|
||||
client.cancel_ui_request(editor_request.id)
|
||||
client.wait_for_idle(timeout=2.0)
|
||||
|
||||
def test_prompt_lifecycle_collectors_are_single_flight(self) -> None:
|
||||
results: list[str] = []
|
||||
errors: list[BaseException] = []
|
||||
|
||||
with self.make_client() as client:
|
||||
def run_prompt() -> None:
|
||||
try:
|
||||
results.append(client.prompt_and_wait("slow", timeout=2.0).require_assistant_text())
|
||||
except BaseException as exc: # pragma: no cover - defensive thread capture
|
||||
errors.append(exc)
|
||||
|
||||
thread = threading.Thread(target=run_prompt)
|
||||
thread.start()
|
||||
|
||||
deadline = time.time() + 1.0
|
||||
while client._prompt_lifecycle.active_operation != "prompt_and_wait" and time.time() < deadline:
|
||||
time.sleep(0.01)
|
||||
|
||||
self.assertEqual(client._prompt_lifecycle.active_operation, "prompt_and_wait")
|
||||
with self.assertRaises(RpcConcurrencyError):
|
||||
client.collect_events(timeout=1.0)
|
||||
|
||||
thread.join(timeout=2.0)
|
||||
self.assertFalse(thread.is_alive())
|
||||
|
||||
self.assertEqual(errors, [])
|
||||
self.assertEqual(results, ["pong"])
|
||||
|
||||
def test_listener_mutation_does_not_change_retained_turn(self) -> None:
|
||||
with self.make_client() as client:
|
||||
client.on_message_end(lambda event: event.message["content"].__setitem__(0, {"type": "text", "text": "mutated"}))
|
||||
turn = client.prompt_and_wait("say hello", timeout=2.0)
|
||||
messages = client.get_messages()
|
||||
|
||||
self.assertEqual(turn.require_assistant_text(), "pong")
|
||||
self.assertEqual(messages[0]["content"][0]["text"], "pong")
|
||||
|
||||
def test_id_less_error_responses_are_correlated(self) -> None:
|
||||
with self.make_client(server=IDLESS_ERROR_SERVER) as client:
|
||||
with self.assertRaises(RpcCommandError) as ctx:
|
||||
@@ -470,6 +970,14 @@ class RpcClientTests(unittest.TestCase):
|
||||
|
||||
self.assertEqual(client.stderr, "second\n")
|
||||
|
||||
def test_broken_startup_frame_is_reported(self) -> None:
|
||||
client = self.make_client(server=BROKEN_STARTUP_SERVER)
|
||||
|
||||
with self.assertRaises(RpcError) as ctx:
|
||||
client.start()
|
||||
|
||||
self.assertIn("Frame: 'not-json'", str(ctx.exception))
|
||||
|
||||
def test_event_history_limit_reports_overflow(self) -> None:
|
||||
with self.make_client(max_event_history=2) as client:
|
||||
with self.assertRaises(RpcError) as ctx:
|
||||
|
||||
@@ -170,6 +170,116 @@ class ProtocolParsingTests(unittest.TestCase):
|
||||
self.assertEqual(assistant_text(message), "visible")
|
||||
self.assertEqual(assistant_text_with_thinking(message), "internalvisible")
|
||||
|
||||
def test_parse_session_state_rejects_invalid_thinking_level(self) -> None:
|
||||
with self.assertRaises(ValueError):
|
||||
parse_session_state(
|
||||
{
|
||||
"sessionId": "session-123",
|
||||
"thinkingLevel": "extreme",
|
||||
"steeringMode": "one-at-a-time",
|
||||
"followUpMode": "one-at-a-time",
|
||||
"interruptMode": "immediate",
|
||||
}
|
||||
)
|
||||
|
||||
def test_parse_extension_ui_request_rejects_invalid_method(self) -> None:
|
||||
with self.assertRaises(ValueError):
|
||||
parse_notification({"type": "extension_ui_request", "id": "ui-1", "method": "launch"})
|
||||
|
||||
def test_parse_message_update_rejects_invalid_assistant_done_reason(self) -> None:
|
||||
with self.assertRaises(ValueError):
|
||||
parse_notification(
|
||||
{
|
||||
"type": "message_update",
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": [{"type": "text", "text": "hello"}],
|
||||
"api": "anthropic-messages",
|
||||
"provider": "anthropic",
|
||||
"model": "claude-sonnet-4-5",
|
||||
"usage": {
|
||||
"input": 1,
|
||||
"output": 1,
|
||||
"cacheRead": 0,
|
||||
"cacheWrite": 0,
|
||||
"totalTokens": 2,
|
||||
"cost": {
|
||||
"input": 0.0,
|
||||
"output": 0.0,
|
||||
"cacheRead": 0.0,
|
||||
"cacheWrite": 0.0,
|
||||
"total": 0.0,
|
||||
},
|
||||
},
|
||||
"stopReason": "stop",
|
||||
"timestamp": 1,
|
||||
},
|
||||
"assistantMessageEvent": {
|
||||
"type": "done",
|
||||
"reason": "error",
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": [{"type": "text", "text": "hello"}],
|
||||
"api": "anthropic-messages",
|
||||
"provider": "anthropic",
|
||||
"model": "claude-sonnet-4-5",
|
||||
"usage": {
|
||||
"input": 1,
|
||||
"output": 1,
|
||||
"cacheRead": 0,
|
||||
"cacheWrite": 0,
|
||||
"totalTokens": 2,
|
||||
"cost": {
|
||||
"input": 0.0,
|
||||
"output": 0.0,
|
||||
"cacheRead": 0.0,
|
||||
"cacheWrite": 0.0,
|
||||
"total": 0.0,
|
||||
},
|
||||
},
|
||||
"stopReason": "stop",
|
||||
"timestamp": 1,
|
||||
},
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
def test_parse_notification_deep_clones_nested_messages(self) -> None:
|
||||
payload = {
|
||||
"type": "agent_end",
|
||||
"messages": [
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [{"type": "text", "text": "hello"}],
|
||||
"api": "anthropic-messages",
|
||||
"provider": "anthropic",
|
||||
"model": "claude-sonnet-4-5",
|
||||
"usage": {
|
||||
"input": 1,
|
||||
"output": 1,
|
||||
"cacheRead": 0,
|
||||
"cacheWrite": 0,
|
||||
"totalTokens": 2,
|
||||
"cost": {
|
||||
"input": 0.0,
|
||||
"output": 0.0,
|
||||
"cacheRead": 0.0,
|
||||
"cacheWrite": 0.0,
|
||||
"total": 0.0,
|
||||
},
|
||||
},
|
||||
"stopReason": "stop",
|
||||
"timestamp": 1,
|
||||
}
|
||||
],
|
||||
}
|
||||
|
||||
notification = parse_notification(payload)
|
||||
payload["messages"][0]["content"][0]["text"] = "mutated"
|
||||
|
||||
self.assertIsInstance(notification, AgentEndEvent)
|
||||
self.assertEqual(notification.messages[0]["content"][0]["text"], "hello")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
Reference in New Issue
Block a user