fix(rpc): harden Python protocol compatibility
This commit is contained in:
@@ -49,6 +49,7 @@ from .protocol import (
|
||||
CancellationResult,
|
||||
CompactionResult,
|
||||
CompactionSummaryMessage,
|
||||
ContextUsage,
|
||||
CustomMessage,
|
||||
DeveloperMessage,
|
||||
Effort,
|
||||
@@ -116,6 +117,7 @@ __all__ = [
|
||||
"CancellationResult",
|
||||
"CompactionResult",
|
||||
"CompactionSummaryMessage",
|
||||
"ContextUsage",
|
||||
"CustomMessage",
|
||||
"DeveloperMessage",
|
||||
"Effort",
|
||||
|
||||
@@ -1353,7 +1353,9 @@ class RpcClient:
|
||||
|
||||
event_payloads = self._events.snapshot_from(start_index)
|
||||
if any(
|
||||
payload.get("type") == "agent_end" for payload in event_payloads
|
||||
payload.get("type") == "agent_end"
|
||||
and payload.get("isTerminal") is not False
|
||||
for payload in event_payloads
|
||||
):
|
||||
events = tuple(
|
||||
cast(RpcAgentEvent, parse_notification(payload))
|
||||
@@ -1898,8 +1900,16 @@ class RpcClient:
|
||||
payload_type = payload.get("type")
|
||||
if payload_type in ("tool_execution_update", "tool_execution_end"):
|
||||
self._normalize_host_tool_event(payload)
|
||||
notification = parse_notification(payload)
|
||||
listener_notification = parse_notification(payload)
|
||||
try:
|
||||
notification = parse_notification(payload)
|
||||
listener_notification = parse_notification(payload)
|
||||
except (TypeError, ValueError) as exc:
|
||||
notification = UnknownNotification(
|
||||
_clone_json_object(payload), parse_error=str(exc)
|
||||
)
|
||||
listener_notification = UnknownNotification(
|
||||
_clone_json_object(payload), parse_error=str(exc)
|
||||
)
|
||||
self._dispatch_listeners(
|
||||
"notification",
|
||||
listener_notification.type,
|
||||
@@ -1949,7 +1959,10 @@ class RpcClient:
|
||||
|
||||
listener_event = cast(RpcAgentEvent, listener_notification)
|
||||
self._append_event(payload)
|
||||
if listener_event.type == "agent_end":
|
||||
if (
|
||||
isinstance(listener_event, AgentEndEvent)
|
||||
and listener_event.is_terminal is not False
|
||||
):
|
||||
self._mark_agent_run_completed()
|
||||
self._dispatch_listeners(
|
||||
"event", listener_event.type, self._event_listeners, listener_event
|
||||
|
||||
@@ -35,17 +35,30 @@ ExtensionUiMethod: TypeAlias = Literal[
|
||||
"setWidget",
|
||||
"setTitle",
|
||||
"set_editor_text",
|
||||
"open_url",
|
||||
]
|
||||
InteractiveExtensionUiMethod: TypeAlias = Literal[
|
||||
"select", "confirm", "input", "editor"
|
||||
]
|
||||
PassiveExtensionUiMethod: TypeAlias = Literal[
|
||||
"notify", "setStatus", "setWidget", "setTitle", "set_editor_text"
|
||||
"notify",
|
||||
"setStatus",
|
||||
"setWidget",
|
||||
"setTitle",
|
||||
"set_editor_text",
|
||||
"open_url",
|
||||
]
|
||||
ValueExtensionUiMethod: TypeAlias = Literal["select", "input", "editor"]
|
||||
|
||||
PASSIVE_EXTENSION_UI_METHODS: Final[frozenset[PassiveExtensionUiMethod]] = frozenset(
|
||||
{"notify", "setStatus", "setWidget", "setTitle", "set_editor_text"}
|
||||
{
|
||||
"notify",
|
||||
"setStatus",
|
||||
"setWidget",
|
||||
"setTitle",
|
||||
"set_editor_text",
|
||||
"open_url",
|
||||
}
|
||||
)
|
||||
INTERACTIVE_EXTENSION_UI_METHODS: Final[frozenset[InteractiveExtensionUiMethod]] = (
|
||||
frozenset({"select", "confirm", "input", "editor"})
|
||||
@@ -81,6 +94,7 @@ _EXTENSION_UI_METHOD_VALUES: Final[frozenset[str]] = frozenset(
|
||||
"setWidget",
|
||||
"setTitle",
|
||||
"set_editor_text",
|
||||
"open_url",
|
||||
}
|
||||
)
|
||||
_AGENT_MESSAGE_ROLE_VALUES: Final[frozenset[str]] = frozenset(
|
||||
@@ -119,10 +133,10 @@ _ASSISTANT_DONE_REASON_VALUES: Final[frozenset[str]] = frozenset(
|
||||
)
|
||||
_ASSISTANT_ERROR_REASON_VALUES: Final[frozenset[str]] = frozenset({"aborted", "error"})
|
||||
_AUTO_COMPACTION_REASON_VALUES: Final[frozenset[str]] = frozenset(
|
||||
{"threshold", "overflow", "idle"}
|
||||
{"threshold", "overflow", "idle", "incomplete"}
|
||||
)
|
||||
_AUTO_COMPACTION_ACTION_VALUES: Final[frozenset[str]] = frozenset(
|
||||
{"context-full", "handoff"}
|
||||
{"context-full", "handoff", "shake", "snapcompact"}
|
||||
)
|
||||
|
||||
|
||||
@@ -774,6 +788,13 @@ class TodoPhase:
|
||||
tasks: tuple[TodoItem, ...]
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class ContextUsage:
|
||||
tokens: int
|
||||
context_window: int
|
||||
percent: float
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class SessionState:
|
||||
model: ModelInfo | None
|
||||
@@ -795,6 +816,7 @@ class SessionState:
|
||||
fast_mode_enabled: bool = False
|
||||
fast_mode_active: bool = False
|
||||
tokens_per_second: float | None = None
|
||||
context_usage: ContextUsage | None = None
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
@@ -913,6 +935,9 @@ class ExtensionUiRequest:
|
||||
widget_lines: tuple[str, ...] | None = None
|
||||
widget_placement: WidgetPlacement | None = None
|
||||
text: str | None = None
|
||||
url: str | None = None
|
||||
launch_url: str | None = None
|
||||
instructions: str | None = None
|
||||
type: Literal["extension_ui_request"] = "extension_ui_request"
|
||||
|
||||
def is_passive(self) -> bool:
|
||||
@@ -946,6 +971,7 @@ class AgentEndEvent:
|
||||
messages: tuple[AgentMessage, ...]
|
||||
type: Literal["agent_end"] = "agent_end"
|
||||
message_count: int | None = field(default=None, kw_only=True)
|
||||
is_terminal: bool | None = field(default=None, kw_only=True)
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
@@ -1008,14 +1034,14 @@ class ToolExecutionEndEvent:
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class AutoCompactionStartEvent:
|
||||
reason: Literal["threshold", "overflow", "idle"]
|
||||
action: Literal["context-full", "handoff"]
|
||||
reason: Literal["threshold", "overflow", "idle", "incomplete"]
|
||||
action: Literal["context-full", "handoff", "shake", "snapcompact"]
|
||||
type: Literal["auto_compaction_start"] = "auto_compaction_start"
|
||||
|
||||
|
||||
@dataclass(slots=True, frozen=True)
|
||||
class AutoCompactionEndEvent:
|
||||
action: Literal["context-full", "handoff"]
|
||||
action: Literal["context-full", "handoff", "shake", "snapcompact"]
|
||||
result: CompactionResult | None
|
||||
aborted: bool
|
||||
will_retry: bool
|
||||
@@ -1079,6 +1105,7 @@ class TodoAutoClearEvent:
|
||||
class UnknownNotification:
|
||||
payload: JsonObject
|
||||
type: Literal["unknown"] = "unknown"
|
||||
parse_error: str | None = field(default=None, kw_only=True)
|
||||
|
||||
|
||||
RpcAgentEvent: TypeAlias = (
|
||||
@@ -1372,6 +1399,11 @@ def parse_session_state(payload: JsonObject) -> SessionState:
|
||||
fast_mode_enabled=bool(payload.get("fastModeEnabled", False)),
|
||||
fast_mode_active=bool(payload.get("fastModeActive", False)),
|
||||
tokens_per_second=_optional_float(payload, "tokensPerSecond"),
|
||||
context_usage=parse_context_usage(
|
||||
_optional_json_object(
|
||||
payload.get("contextUsage"), field="sessionState.contextUsage"
|
||||
)
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@@ -1483,6 +1515,16 @@ def parse_session_stats(payload: JsonObject) -> SessionStats:
|
||||
)
|
||||
|
||||
|
||||
def parse_context_usage(payload: JsonObject | None) -> ContextUsage | None:
|
||||
if payload is None:
|
||||
return None
|
||||
return ContextUsage(
|
||||
tokens=int(payload.get("tokens", 0)),
|
||||
context_window=int(payload.get("contextWindow", 0)),
|
||||
percent=float(payload.get("percent", 0.0)),
|
||||
)
|
||||
|
||||
|
||||
def parse_extension_ui_request(payload: JsonObject) -> ExtensionUiRequest:
|
||||
return ExtensionUiRequest(
|
||||
id=_require_str(payload, "id"),
|
||||
@@ -1527,6 +1569,9 @@ def parse_extension_ui_request(payload: JsonObject) -> ExtensionUiRequest:
|
||||
),
|
||||
),
|
||||
text=_optional_str(payload, "text"),
|
||||
url=_optional_str(payload, "url"),
|
||||
launch_url=_optional_str(payload, "launchUrl"),
|
||||
instructions=_optional_str(payload, "instructions"),
|
||||
)
|
||||
|
||||
|
||||
@@ -1570,6 +1615,7 @@ def parse_notification(payload: JsonObject) -> RpcNotification:
|
||||
cast(JsonValue | None, payload.get("messages"))
|
||||
),
|
||||
message_count=_optional_int(payload, "messageCount"),
|
||||
is_terminal=_optional_bool(payload, "isTerminal"),
|
||||
)
|
||||
if event_type == "turn_start":
|
||||
return TurnStartEvent()
|
||||
@@ -1661,7 +1707,7 @@ def parse_notification(payload: JsonObject) -> RpcNotification:
|
||||
if event_type == "auto_compaction_start":
|
||||
return AutoCompactionStartEvent(
|
||||
reason=cast(
|
||||
Literal["threshold", "overflow", "idle"],
|
||||
Literal["threshold", "overflow", "idle", "incomplete"],
|
||||
_require_literal(
|
||||
payload.get("reason", "threshold"),
|
||||
_AUTO_COMPACTION_REASON_VALUES,
|
||||
@@ -1669,7 +1715,7 @@ def parse_notification(payload: JsonObject) -> RpcNotification:
|
||||
),
|
||||
),
|
||||
action=cast(
|
||||
Literal["context-full", "handoff"],
|
||||
Literal["context-full", "handoff", "shake", "snapcompact"],
|
||||
_require_literal(
|
||||
payload.get("action", "context-full"),
|
||||
_AUTO_COMPACTION_ACTION_VALUES,
|
||||
@@ -1681,7 +1727,7 @@ def parse_notification(payload: JsonObject) -> RpcNotification:
|
||||
result_payload = payload.get("result")
|
||||
return AutoCompactionEndEvent(
|
||||
action=cast(
|
||||
Literal["context-full", "handoff"],
|
||||
Literal["context-full", "handoff", "shake", "snapcompact"],
|
||||
_require_literal(
|
||||
payload.get("action", "context-full"),
|
||||
_AUTO_COMPACTION_ACTION_VALUES,
|
||||
|
||||
@@ -12,7 +12,14 @@ import threading
|
||||
import time
|
||||
import unittest
|
||||
|
||||
from omp_rpc import RpcClient, RpcCommandError, RpcConcurrencyError, RpcError, host_tool
|
||||
from omp_rpc import (
|
||||
AgentEndEvent,
|
||||
RpcClient,
|
||||
RpcCommandError,
|
||||
RpcConcurrencyError,
|
||||
RpcError,
|
||||
host_tool,
|
||||
)
|
||||
from omp_rpc.client import _RpcFrameDecoder
|
||||
|
||||
|
||||
@@ -806,6 +813,54 @@ BROKEN_STARTUP_SERVER = textwrap.dedent(
|
||||
"""
|
||||
)
|
||||
|
||||
FORWARD_COMPAT_SERVER = textwrap.dedent(
|
||||
"""
|
||||
import json
|
||||
import sys
|
||||
import time
|
||||
|
||||
print(json.dumps({"type": "ready"}), flush=True)
|
||||
for raw_line in sys.stdin:
|
||||
command = json.loads(raw_line)
|
||||
print(
|
||||
json.dumps(
|
||||
{
|
||||
"id": command.get("id"),
|
||||
"type": "response",
|
||||
"command": command["type"],
|
||||
"success": True,
|
||||
}
|
||||
),
|
||||
flush=True,
|
||||
)
|
||||
if command["type"] != "prompt":
|
||||
continue
|
||||
print(
|
||||
json.dumps(
|
||||
{
|
||||
"type": "auto_compaction_start",
|
||||
"reason": "future_reason",
|
||||
"action": "future_action",
|
||||
}
|
||||
),
|
||||
flush=True,
|
||||
)
|
||||
print(
|
||||
json.dumps(
|
||||
{"type": "agent_end", "messages": [], "isTerminal": False}
|
||||
),
|
||||
flush=True,
|
||||
)
|
||||
time.sleep(0.15)
|
||||
print(
|
||||
json.dumps(
|
||||
{"type": "agent_end", "messages": [], "isTerminal": True}
|
||||
),
|
||||
flush=True,
|
||||
)
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
class RpcClientTests(unittest.TestCase):
|
||||
def make_client(self, server: str = FAKE_SERVER, **kwargs: object) -> RpcClient:
|
||||
@@ -1260,6 +1315,24 @@ class RpcClientTests(unittest.TestCase):
|
||||
self.assertEqual(seen_extension_errors, ["boom"])
|
||||
self.assertEqual(seen_unknown, ["unknown_future_event"])
|
||||
|
||||
def test_additive_notification_values_do_not_stop_the_reader(self) -> None:
|
||||
unknown_errors: list[str | None] = []
|
||||
|
||||
with self.make_client(server=FORWARD_COMPAT_SERVER) as client:
|
||||
client.on_unknown_notification(
|
||||
lambda event: unknown_errors.append(event.parse_error)
|
||||
)
|
||||
turn = client.prompt_and_wait("forward compatible", timeout=2.0)
|
||||
|
||||
terminal_events = [
|
||||
event for event in turn.events if isinstance(event, AgentEndEvent)
|
||||
]
|
||||
self.assertEqual(
|
||||
[event.is_terminal for event in terminal_events], [False, True]
|
||||
)
|
||||
self.assertEqual(len(unknown_errors), 1)
|
||||
self.assertIn("auto_compaction_start.reason", unknown_errors[0] or "")
|
||||
|
||||
def test_ui_confirmation_and_cancel_round_trip(self) -> None:
|
||||
with self.make_client() as client:
|
||||
client.prompt("needs confirm")
|
||||
|
||||
@@ -4,6 +4,8 @@ import unittest
|
||||
|
||||
from omp_rpc import (
|
||||
AgentEndEvent,
|
||||
AutoCompactionEndEvent,
|
||||
AutoCompactionStartEvent,
|
||||
ExtensionUiRequest,
|
||||
SessionState,
|
||||
TodoReminderEvent,
|
||||
@@ -79,6 +81,11 @@ class ProtocolParsingTests(unittest.TestCase):
|
||||
"parameters": {"type": "object"},
|
||||
}
|
||||
],
|
||||
"contextUsage": {
|
||||
"tokens": 12345,
|
||||
"contextWindow": 200000,
|
||||
"percent": 6.1725,
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
@@ -90,6 +97,10 @@ class ProtocolParsingTests(unittest.TestCase):
|
||||
# Legacy bare-string systemPrompt is accepted and wrapped to a tuple.
|
||||
self.assertEqual(state.system_prompt, ("You are useful.",))
|
||||
self.assertEqual(state.dump_tools[0].name, "read")
|
||||
assert state.context_usage is not None
|
||||
self.assertEqual(state.context_usage.tokens, 12345)
|
||||
self.assertEqual(state.context_usage.context_window, 200000)
|
||||
self.assertEqual(state.context_usage.percent, 6.1725)
|
||||
assert state.model is not None and state.model.thinking is not None
|
||||
self.assertEqual(
|
||||
state.model.thinking.efforts, ("minimal", "low", "medium", "high")
|
||||
@@ -161,16 +172,43 @@ class ProtocolParsingTests(unittest.TestCase):
|
||||
}
|
||||
],
|
||||
"messageCount": 1,
|
||||
"isTerminal": False,
|
||||
}
|
||||
)
|
||||
|
||||
self.assertIsInstance(notification, AgentEndEvent)
|
||||
self.assertEqual(assistant_text(notification.messages[0]), "hello")
|
||||
self.assertEqual(notification.message_count, 1)
|
||||
self.assertFalse(notification.is_terminal)
|
||||
|
||||
legacy = AgentEndEvent(notification.messages, "agent_end")
|
||||
self.assertEqual(legacy.type, "agent_end")
|
||||
self.assertIsNone(legacy.message_count)
|
||||
self.assertIsNone(legacy.is_terminal)
|
||||
|
||||
def test_parse_current_compaction_variants(self) -> None:
|
||||
start = parse_notification(
|
||||
{
|
||||
"type": "auto_compaction_start",
|
||||
"reason": "incomplete",
|
||||
"action": "snapcompact",
|
||||
}
|
||||
)
|
||||
end = parse_notification(
|
||||
{
|
||||
"type": "auto_compaction_end",
|
||||
"action": "shake",
|
||||
"result": None,
|
||||
"aborted": False,
|
||||
"willRetry": False,
|
||||
}
|
||||
)
|
||||
|
||||
self.assertIsInstance(start, AutoCompactionStartEvent)
|
||||
self.assertEqual(start.reason, "incomplete")
|
||||
self.assertEqual(start.action, "snapcompact")
|
||||
self.assertIsInstance(end, AutoCompactionEndEvent)
|
||||
self.assertEqual(end.action, "shake")
|
||||
|
||||
def test_parse_extension_ui_request(self) -> None:
|
||||
notification = parse_notification(
|
||||
@@ -191,6 +229,24 @@ class ProtocolParsingTests(unittest.TestCase):
|
||||
self.assertTrue(notification.requires_response())
|
||||
self.assertFalse(notification.is_passive())
|
||||
|
||||
def test_parse_open_url_request(self) -> None:
|
||||
notification = parse_notification(
|
||||
{
|
||||
"type": "extension_ui_request",
|
||||
"id": "ui-oauth",
|
||||
"method": "open_url",
|
||||
"url": "https://example.com/oauth",
|
||||
"launchUrl": "http://127.0.0.1:8123/redirect",
|
||||
"instructions": "Open this URL to continue.",
|
||||
}
|
||||
)
|
||||
|
||||
self.assertIsInstance(notification, ExtensionUiRequest)
|
||||
self.assertEqual(notification.method, "open_url")
|
||||
self.assertEqual(notification.url, "https://example.com/oauth")
|
||||
self.assertEqual(notification.launch_url, "http://127.0.0.1:8123/redirect")
|
||||
self.assertTrue(notification.is_passive())
|
||||
|
||||
def test_parse_todo_reminder_notification(self) -> None:
|
||||
notification = parse_notification(
|
||||
{
|
||||
|
||||
Reference in New Issue
Block a user