feat(omp-rpc): added error handling and history bounds to RPC client
- Added error handling callbacks (on_protocol_error, on_listener_error) and error history tracking (protocol_errors, listener_errors properties) for RPC failures. - Added max_event_history and max_stderr_chunks parameters to bound retained history in long-lived hosts, preventing unbounded memory growth. - Added assistant_text_with_thinking() and message_text_with_thinking() helper functions for backward compatibility with thinking block inclusion. - Improved error correlation for id-less parse/unknown-command failures to waiting requests and changed prompt scheduling failures to raise instead of timing out. - Refactored event dispatch to _dispatch_listeners() with exception handling and error recording for robust listener failure management.
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user