diff --git a/python/omp-rpc/README.md b/python/omp-rpc/README.md index f6bb448ea..bb1dbd730 100644 --- a/python/omp-rpc/README.md +++ b/python/omp-rpc/README.md @@ -125,6 +125,46 @@ That helper ignores passive UI notifications (`notify`, `setStatus`, `setWidget` `setTitle`, `set_editor_text`), answers `confirm` with `False`, and cancels `select`/`input`/`editor` requests unless you provide explicit values. +## Error Handling and Retained History + +The client now surfaces more of the transport edge cases that the wire protocol +allows: + +- id-less `parse` and unknown-command failures are correlated back to the + waiting request when they can be matched unambiguously +- late `prompt` / `abort_and_prompt` scheduling failures cause + `prompt_and_wait()` and `wait_for_idle()` to raise instead of timing out +- unmatched background error responses are exposed through + `client.protocol_errors` and `client.on_protocol_error(...)` +- listener exceptions no longer kill the stdout reader thread; they are exposed + through `client.listener_errors` and `client.on_listener_error(...)` + +For long-lived hosts, retained event and stderr history is bounded by default: + +```python +from omp_rpc import RpcClient + +with RpcClient(max_event_history=20_000, max_stderr_chunks=256) as client: + ... +``` + +If a single prompt streams more events than `max_event_history` allows, +`prompt_and_wait()` raises a clear error so hosts can increase the limit instead +of silently losing earlier events. + +## Text Helpers + +`assistant_text()` and `message_text()` now return visible text blocks only. +If a host explicitly needs reasoning text too, use the `*_with_thinking` +helpers: + +```python +from omp_rpc import assistant_text, assistant_text_with_thinking + +visible = assistant_text(message) +full = assistant_text_with_thinking(message) +``` + ## Protocol Reference The canonical wire protocol still lives in the repo at diff --git a/python/omp-rpc/src/omp_rpc/__init__.py b/python/omp-rpc/src/omp_rpc/__init__.py index 0be2cdae8..207eafeba 100644 --- a/python/omp-rpc/src/omp_rpc/__init__.py +++ b/python/omp-rpc/src/omp_rpc/__init__.py @@ -1,6 +1,8 @@ from .client import ( AgentEventListener, ExtensionErrorListener, + ListenerErrorEvent, + ListenerErrorListener, NotificationListener, PromptTurn, ReadyListener, @@ -8,7 +10,9 @@ from .client import ( RpcCommandError, RpcError, RpcProcessExitError, + RpcProtocolError, RpcTimeoutError, + ProtocolErrorListener, UiRequestListener, ) from .protocol import ( @@ -65,8 +69,10 @@ from .protocol import ( UnknownNotification, UserMessage, assistant_text, + assistant_text_with_thinking, image_from_path, message_text, + message_text_with_thinking, parse_notification, parse_session_state, parse_todo_phases, @@ -96,6 +102,8 @@ __all__ = [ "FileMentionMessage", "HookMessage", "ImageContent", + "ListenerErrorEvent", + "ListenerErrorListener", "MessageEndEvent", "MessageStartEvent", "MessageUpdateEvent", @@ -104,6 +112,7 @@ __all__ = [ "ModelInfo", "NotificationListener", "PromptTurn", + "ProtocolErrorListener", "PythonExecutionMessage", "ReadyEvent", "ReadyListener", @@ -115,6 +124,7 @@ __all__ = [ "RpcError", "RpcNotification", "RpcProcessExitError", + "RpcProtocolError", "RpcTimeoutError", "SessionState", "SessionStats", @@ -137,8 +147,10 @@ __all__ = [ "UnknownNotification", "UserMessage", "assistant_text", + "assistant_text_with_thinking", "image_from_path", "message_text", + "message_text_with_thinking", "parse_notification", "parse_session_state", "parse_todo_phases", diff --git a/python/omp-rpc/src/omp_rpc/client.py b/python/omp-rpc/src/omp_rpc/client.py index 344388610..f21d142ec 100644 --- a/python/omp-rpc/src/omp_rpc/client.py +++ b/python/omp-rpc/src/omp_rpc/client.py @@ -98,9 +98,14 @@ RetryFallbackSucceededListener = Callable[[RetryFallbackSucceededEvent], None] TtsrTriggeredListener = Callable[[TtsrTriggeredEvent], None] TodoReminderListener = Callable[[TodoReminderEvent], None] TodoAutoClearListener = Callable[[TodoAutoClearEvent], None] +ProtocolErrorListener = Callable[["RpcProtocolError"], None] +ListenerErrorListener = Callable[["ListenerErrorEvent"], None] TListener = TypeVar("TListener") TEventListener = TypeVar("TEventListener", bound=Callable[..., None]) +_ASYNC_COMMANDS = frozenset({"prompt", "abort_and_prompt"}) +_DEFAULT_ERROR_HISTORY_LIMIT = 128 + class RpcError(RuntimeError): """Base exception for the Python RPC client.""" @@ -123,6 +128,36 @@ class RpcCommandError(RpcError): self.error = error +class RpcProtocolError(RpcError): + """Raised or reported when the transport receives an unmatched RPC error response.""" + + def __init__(self, payload: JsonObject): + self.payload = dict(payload) + command = payload.get("command") + request_id = payload.get("id") + error = payload.get("error") + self.command = str(command) if isinstance(command, str) else None + self.request_id = str(request_id) if isinstance(request_id, str) else None + self.remote_error = str(error) if isinstance(error, str) else None + + fragments = ["Received unmatched RPC error response"] + if self.command: + fragments.append(f"for {self.command}") + if self.request_id: + fragments.append(f"(id={self.request_id})") + if self.remote_error: + fragments.append(f": {self.remote_error}") + super().__init__(" ".join(fragments)) + + +@dataclass(slots=True, frozen=True) +class ListenerErrorEvent: + listener_kind: str + source_type: str | None + listener: Callable[..., None] + error: BaseException + + @dataclass(slots=True, frozen=True) class PromptTurn: events: tuple[RpcAgentEvent, ...] @@ -140,6 +175,12 @@ TodoSeed = str | TodoItem | Mapping[str, object] TodoPhaseSeed = TodoPhase | Mapping[str, object] +@dataclass(slots=True) +class _PendingRequest: + command: str + response_queue: queue.Queue[JsonObject | BaseException] + + class RpcClient: def __init__( self, @@ -163,6 +204,8 @@ class RpcClient: extra_args: Sequence[str] = (), startup_timeout: float = 30.0, request_timeout: float = 30.0, + max_event_history: int | None = 10_000, + max_stderr_chunks: int | None = 512, ) -> None: self._command = tuple(command) if command is not None else None self._executable = executable @@ -183,6 +226,8 @@ class RpcClient: self._extra_args = tuple(extra_args) self._startup_timeout = startup_timeout self._request_timeout = request_timeout + self._max_event_history = self._validate_history_limit("max_event_history", max_event_history) + self._max_stderr_chunks = self._validate_history_limit("max_stderr_chunks", max_stderr_chunks) self._process: subprocess.Popen[str] | None = None self._stdout_thread: threading.Thread | None = None @@ -191,13 +236,18 @@ class RpcClient: self._write_lock = threading.Lock() self._state_lock = threading.Lock() self._event_condition = threading.Condition() - self._pending: dict[str, queue.Queue[JsonObject | BaseException]] = {} + self._pending: dict[str, _PendingRequest] = {} self._request_id = 0 self._events: list[RpcAgentEvent] = [] + self._event_offset = 0 + self._async_errors: list[BaseException] = [] + self._async_error_offset = 0 self._ui_requests: queue.Queue[ExtensionUiRequest] = queue.Queue() self._stderr_chunks: list[str] = [] self._closed_error: BaseException | None = None self._stopping = False + self._protocol_errors: list[RpcProtocolError] = [] + self._listener_errors: list[ListenerErrorEvent] = [] self._notification_listeners: list[NotificationListener] = [] self._event_listeners: list[AgentEventListener] = [] @@ -206,6 +256,8 @@ class RpcClient: self._unknown_notification_listeners: list[UnknownNotificationListener] = [] self._ui_request_listeners: list[UiRequestListener] = [] self._extension_error_listeners: list[ExtensionErrorListener] = [] + self._protocol_error_listeners: list[ProtocolErrorListener] = [] + self._listener_error_listeners: list[ListenerErrorListener] = [] def __enter__(self) -> RpcClient: return self.start() @@ -221,6 +273,16 @@ class RpcClient: def command(self) -> tuple[str, ...]: return self._build_command() + @property + def protocol_errors(self) -> tuple[RpcProtocolError, ...]: + with self._state_lock: + return tuple(self._protocol_errors) + + @property + def listener_errors(self) -> tuple[ListenerErrorEvent, ...]: + with self._state_lock: + return tuple(self._listener_errors) + def start(self) -> RpcClient: if self._process is not None: raise RpcError("RPC client is already started") @@ -229,8 +291,14 @@ class RpcClient: self._stopping = False self._closed_error = None self._events.clear() + self._event_offset = 0 + self._async_errors.clear() + self._async_error_offset = 0 self._ui_requests = queue.Queue() self._stderr_chunks.clear() + with self._state_lock: + self._protocol_errors.clear() + self._listener_errors.clear() process = subprocess.Popen( list(self._build_command()), @@ -378,6 +446,14 @@ class RpcClient: self._extension_error_listeners.append(listener) return lambda: self._remove_listener(self._extension_error_listeners, listener) + def on_protocol_error(self, listener: ProtocolErrorListener) -> Callable[[], None]: + self._protocol_error_listeners.append(listener) + return lambda: self._remove_listener(self._protocol_error_listeners, listener) + + def on_listener_error(self, listener: ListenerErrorListener) -> Callable[[], None]: + self._listener_error_listeners.append(listener) + return lambda: self._remove_listener(self._listener_error_listeners, listener) + def on_unknown_notification(self, listener: UnknownNotificationListener) -> Callable[[], None]: self._unknown_notification_listeners.append(listener) return lambda: self._remove_listener(self._unknown_notification_listeners, listener) @@ -400,7 +476,17 @@ class RpcClient: def handle(request: ExtensionUiRequest) -> None: if on_request is not None: - on_request(request) + try: + on_request(request) + except Exception as exc: + self._record_listener_error( + ListenerErrorEvent( + listener_kind="headless_ui_request", + source_type=request.type, + listener=on_request, + error=exc, + ) + ) if request.method == "cancel" or request.is_passive(): return @@ -577,24 +663,31 @@ class RpcClient: timeout: float | None = None, ) -> PromptTurn: start_index = self._current_event_index() + start_async_error_index = self._current_async_error_index() self.prompt(message, images=images, streaming_behavior=streaming_behavior) - events = self._wait_for_agent_end(start_index, timeout=timeout) + events = self._wait_for_agent_end(start_index, start_async_error_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) + start_async_error_index = self._current_async_error_index() + self._wait_for_agent_end(start_index, start_async_error_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) + start_async_error_index = self._current_async_error_index() + return self._wait_for_agent_end(start_index, start_async_error_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) + return self._event_offset + len(self._events) + + def _current_async_error_index(self) -> int: + with self._event_condition: + return self._async_error_offset + len(self._async_errors) def _build_prompt_turn(self, events: tuple[RpcAgentEvent, ...]) -> PromptTurn: final_messages: tuple[AgentMessage, ...] = () @@ -624,14 +717,36 @@ class RpcClient: 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, ...]: + def _wait_for_agent_end( + self, + start_index: int, + start_async_error_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 start_index < self._event_offset: + raise RpcError( + "Event history limit was exceeded while waiting for agent_end. " + "Increase max_event_history to retain more streamed events." + ) + + if start_async_error_index < self._async_error_offset: + raise RpcError( + "Async error history limit was exceeded while waiting for agent_end. " + "Increase max_event_history if your host needs to retain more background failures." + ) + + async_error_index = start_async_error_index - self._async_error_offset + if async_error_index < len(self._async_errors): + raise self._async_errors[async_error_index] + + event_index = start_index - self._event_offset + events = tuple(self._events[event_index:]) if any(isinstance(event, AgentEndEvent) for event in events): return events @@ -650,7 +765,7 @@ class RpcClient: response_queue: queue.Queue[JsonObject | BaseException] = queue.Queue(maxsize=1) with self._state_lock: - self._pending[request_id] = response_queue + self._pending[request_id] = _PendingRequest(command=command_type, response_queue=response_queue) self._write_json(process, envelope) @@ -827,48 +942,40 @@ class RpcClient: 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) + self._handle_response(payload) continue notification = parse_notification(payload) - for listener in list(self._notification_listeners): - listener(notification) + self._dispatch_listeners("notification", notification.type, self._notification_listeners, notification) if isinstance(notification, ReadyEvent): self._ready.set() - for listener in list(self._ready_listeners): - listener(notification) + self._dispatch_listeners("ready", notification.type, self._ready_listeners, notification) continue if isinstance(notification, ExtensionUiRequest): self._ui_requests.put(notification) - for listener in list(self._ui_request_listeners): - listener(notification) + self._dispatch_listeners("ui_request", notification.type, self._ui_request_listeners, notification) continue if isinstance(notification, ExtensionError): - for listener in list(self._extension_error_listeners): - listener(notification) + self._dispatch_listeners( + "extension_error", notification.type, self._extension_error_listeners, notification + ) continue if isinstance(notification, UnknownNotification): - for listener in list(self._unknown_notification_listeners): - listener(notification) + self._dispatch_listeners( + "unknown_notification", notification.type, self._unknown_notification_listeners, notification + ) continue event = cast(RpcAgentEvent, notification) - with self._event_condition: - self._events.append(event) - self._event_condition.notify_all() - for listener in list(self._event_listeners): - listener(event) - for listener in list(self._typed_event_listeners.get(event.type, [])): - listener(event) + self._append_event(event) + self._dispatch_listeners("event", event.type, self._event_listeners, event) + self._dispatch_listeners( + "typed_event", event.type, self._typed_event_listeners.get(event.type, []), event + ) except json.JSONDecodeError as exc: self._mark_closed(RpcError(f"Failed to decode RPC output: {exc}")) except Exception as exc: @@ -890,6 +997,9 @@ class RpcClient: return for chunk in process.stderr: self._stderr_chunks.append(chunk) + if self._max_stderr_chunks is not None and len(self._stderr_chunks) > self._max_stderr_chunks: + trim = len(self._stderr_chunks) - self._max_stderr_chunks + del self._stderr_chunks[:trim] def _mark_closed(self, error: BaseException) -> None: if self._closed_error is not None: @@ -902,11 +1012,130 @@ class RpcClient: def _fail_pending(self, error: BaseException) -> None: with self._state_lock: - pending = list(self._pending.values()) + pending = [pending.response_queue for pending in self._pending.values()] self._pending.clear() for response_queue in pending: response_queue.put(error) + def _handle_response(self, payload: JsonObject) -> None: + 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.response_queue.put(payload) + return + + if self._deliver_correlated_error_response(payload): + return + + protocol_error = self._build_protocol_error(payload) + if protocol_error is None: + return + + if protocol_error.command in _ASYNC_COMMANDS and protocol_error.remote_error is not None: + self._append_async_error(RpcCommandError(protocol_error.command, protocol_error.remote_error)) + + self._record_protocol_error(protocol_error) + + def _deliver_correlated_error_response(self, payload: JsonObject) -> bool: + if bool(payload.get("success", False)): + return False + + command = payload.get("command") + if not isinstance(command, str): + return False + + with self._state_lock: + matching_ids = [request_id for request_id, pending in self._pending.items() if pending.command == command] + target_id: str | None = None + if len(matching_ids) == 1: + target_id = matching_ids[0] + elif command == "parse" and len(self._pending) == 1: + target_id = next(iter(self._pending)) + + if target_id is None: + return False + + pending = self._pending.pop(target_id) + + pending.response_queue.put(payload) + return True + + def _build_protocol_error(self, payload: JsonObject) -> RpcProtocolError | None: + if payload.get("type") != "response": + return None + if bool(payload.get("success", False)): + return None + return RpcProtocolError(payload) + + def _append_event(self, event: RpcAgentEvent) -> None: + with self._event_condition: + self._events.append(event) + if self._max_event_history is not None and len(self._events) > self._max_event_history: + trim = len(self._events) - self._max_event_history + del self._events[:trim] + self._event_offset += trim + self._event_condition.notify_all() + + def _append_async_error(self, error: BaseException) -> None: + with self._event_condition: + self._async_errors.append(error) + if len(self._async_errors) > _DEFAULT_ERROR_HISTORY_LIMIT: + trim = len(self._async_errors) - _DEFAULT_ERROR_HISTORY_LIMIT + del self._async_errors[:trim] + self._async_error_offset += trim + self._event_condition.notify_all() + + def _record_protocol_error(self, error: RpcProtocolError) -> None: + with self._state_lock: + self._protocol_errors.append(error) + if len(self._protocol_errors) > _DEFAULT_ERROR_HISTORY_LIMIT: + trim = len(self._protocol_errors) - _DEFAULT_ERROR_HISTORY_LIMIT + del self._protocol_errors[:trim] + self._dispatch_listeners("protocol_error", error.command, self._protocol_error_listeners, error) + + def _record_listener_error(self, event: ListenerErrorEvent) -> None: + with self._state_lock: + self._listener_errors.append(event) + if len(self._listener_errors) > _DEFAULT_ERROR_HISTORY_LIMIT: + trim = len(self._listener_errors) - _DEFAULT_ERROR_HISTORY_LIMIT + del self._listener_errors[:trim] + + for listener in list(self._listener_error_listeners): + try: + listener(event) + except Exception: + continue + + def _dispatch_listeners( + self, + listener_kind: str, + source_type: str | None, + listeners: Sequence[Callable[[Any], None]], + payload: Any, + ) -> None: + for listener in list(listeners): + try: + listener(payload) + except Exception as exc: + self._record_listener_error( + ListenerErrorEvent( + listener_kind=listener_kind, + source_type=source_type, + listener=listener, + error=exc, + ) + ) + + @staticmethod + def _validate_history_limit(name: str, limit: int | None) -> int | None: + if limit is None: + return None + if limit <= 0: + raise ValueError(f"{name} must be greater than zero") + return limit + @staticmethod def _remove_listener(listeners: list[TListener], listener: TListener) -> None: try: diff --git a/python/omp-rpc/src/omp_rpc/protocol.py b/python/omp-rpc/src/omp_rpc/protocol.py index 27482d607..a8e345ce3 100644 --- a/python/omp-rpc/src/omp_rpc/protocol.py +++ b/python/omp-rpc/src/omp_rpc/protocol.py @@ -702,7 +702,7 @@ def image_from_path(path: str | Path, mime_type: str | None = None) -> ImageCont } -def message_text(message: AgentMessage) -> str | None: +def message_text(message: AgentMessage, *, include_thinking: bool = False) -> str | None: role = message.get("role") if role not in {"user", "developer", "assistant", "toolResult", "custom", "hookMessage"}: return None @@ -720,15 +720,23 @@ def message_text(message: AgentMessage) -> str | None: 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): + elif include_thinking and 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: +def message_text_with_thinking(message: AgentMessage) -> str | None: + return message_text(message, include_thinking=True) + + +def assistant_text(message: AgentMessage, *, include_thinking: bool = False) -> str | None: if message.get("role") != "assistant": return None - return message_text(message) + return message_text(message, include_thinking=include_thinking) + + +def assistant_text_with_thinking(message: AgentMessage) -> str | None: + return assistant_text(message, include_thinking=True) def parse_model_info(payload: JsonObject | None) -> ModelInfo | None: diff --git a/python/omp-rpc/tests/test_client.py b/python/omp-rpc/tests/test_client.py index 3a1c3f199..dac5d5f0c 100644 --- a/python/omp-rpc/tests/test_client.py +++ b/python/omp-rpc/tests/test_client.py @@ -4,7 +4,7 @@ import sys import textwrap import unittest -from omp_rpc import RpcClient +from omp_rpc import RpcClient, RpcCommandError, RpcError FAKE_SERVER = textwrap.dedent( @@ -201,14 +201,106 @@ FAKE_SERVER = textwrap.dedent( """ ) +IDLESS_ERROR_SERVER = textwrap.dedent( + """ + import json + import sys + + 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) + print( + json.dumps( + { + "type": "response", + "command": command["type"], + "success": False, + "error": f"unsupported: {command['type']}", + } + ), + flush=True, + ) + """ +) + +LATE_PROMPT_FAILURE_SERVER = textwrap.dedent( + """ + import json + import sys + + 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) + request_id = command.get("id") + if command["type"] == "prompt": + print( + json.dumps( + { + "id": request_id, + "type": "response", + "command": "prompt", + "success": True, + } + ), + flush=True, + ) + print( + json.dumps( + { + "id": request_id, + "type": "response", + "command": "prompt", + "success": False, + "error": "late failure", + } + ), + flush=True, + ) + else: + print( + json.dumps( + { + "id": request_id, + "type": "response", + "command": command["type"], + "success": True, + } + ), + flush=True, + ) + """ +) + +STDERR_SERVER = textwrap.dedent( + """ + import json + import sys + + sys.stderr.write("first\\n") + sys.stderr.flush() + sys.stderr.write("second\\n") + sys.stderr.flush() + print(json.dumps({"type": "ready"}), flush=True) + + for _ in sys.stdin: + pass + """ +) + 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 make_client(self, server: str = FAKE_SERVER, **kwargs: object) -> RpcClient: + return RpcClient(command=[sys.executable, "-u", "-c", server], startup_timeout=2.0, request_timeout=2.0, **kwargs) def test_command_builder_supports_common_rpc_options(self) -> None: client = RpcClient( @@ -319,6 +411,72 @@ class RpcClientTests(unittest.TestCase): state = client.get_state() self.assertEqual(state.todo_phases[0].tasks[1].content, "Exercise edits") + 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: + client.request_raw("unknown") + + self.assertEqual(ctx.exception.command, "unknown") + self.assertEqual(ctx.exception.error, "unsupported: unknown") + + def test_prompt_and_wait_raises_for_late_prompt_failure(self) -> None: + protocol_errors: list[str] = [] + client = self.make_client(server=LATE_PROMPT_FAILURE_SERVER) + client.on_protocol_error(lambda error: protocol_errors.append(str(error))) + + try: + client.start() + with self.assertRaises(RpcCommandError) as ctx: + client.prompt_and_wait("say hello", timeout=2.0) + finally: + client.stop() + + self.assertEqual(ctx.exception.command, "prompt") + self.assertEqual(ctx.exception.error, "late failure") + self.assertEqual(len(protocol_errors), 1) + self.assertIn("late failure", protocol_errors[0]) + self.assertEqual(len(client.protocol_errors), 1) + + def test_listener_exceptions_are_reported_without_stopping_client(self) -> None: + listener_errors: list[tuple[str, str | None, str]] = [] + client = self.make_client() + client.on_notification( + lambda notification: (_ for _ in ()).throw(RuntimeError("boom")) + if notification.type == "turn_start" + else None + ) + client.on_listener_error( + lambda event: listener_errors.append((event.listener_kind, event.source_type, str(event.error))) + ) + + try: + client.start() + turn = client.prompt_and_wait("say hello", timeout=2.0) + finally: + client.stop() + + self.assertEqual(turn.require_assistant_text(), "pong") + self.assertEqual(listener_errors, [("notification", "turn_start", "boom")]) + self.assertEqual(len(client.listener_errors), 1) + self.assertEqual(client.listener_errors[0].listener_kind, "notification") + + def test_stderr_history_is_bounded(self) -> None: + client = self.make_client(server=STDERR_SERVER, max_stderr_chunks=1) + + try: + client.start() + finally: + client.stop() + + self.assertEqual(client.stderr, "second\n") + + 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: + client.prompt_and_wait("say hello", timeout=2.0) + + self.assertIn("max_event_history", str(ctx.exception)) + if __name__ == "__main__": unittest.main() diff --git a/python/omp-rpc/tests/test_protocol.py b/python/omp-rpc/tests/test_protocol.py index 3eb56234b..f96813c54 100644 --- a/python/omp-rpc/tests/test_protocol.py +++ b/python/omp-rpc/tests/test_protocol.py @@ -8,6 +8,7 @@ from omp_rpc import ( SessionState, TodoReminderEvent, assistant_text, + assistant_text_with_thinking, parse_notification, parse_session_state, ) @@ -157,6 +158,18 @@ class ProtocolParsingTests(unittest.TestCase): self.assertEqual(notification.todos[0].content, "Map tools") self.assertEqual(notification.todos[0].status, "pending") + def test_assistant_text_excludes_thinking_by_default(self) -> None: + message = { + "role": "assistant", + "content": [ + {"type": "thinking", "thinking": "internal"}, + {"type": "text", "text": "visible"}, + ], + } + + self.assertEqual(assistant_text(message), "visible") + self.assertEqual(assistant_text_with_thinking(message), "internalvisible") + if __name__ == "__main__": unittest.main()