feat(rpc): page stable message histories

This commit is contained in:
Wolfgang Schoenberger
2026-07-22 17:48:53 -07:00
parent 2a6bcc7984
commit ce8c9dc75f
15 changed files with 590 additions and 10 deletions
+1
View File
@@ -9,6 +9,7 @@ provides:
- typed startup options for common `omp --mode rpc` flags such as thinking level,
tool selection, prompt appends, provider session IDs, and headless session toggles
- typed protocol models for state, bash results, compaction, and session stats
- automatic protocol v2 negotiation, lossless chunk reassembly, and stable message pagination
- a process-backed client that manages request correlation over stdio
- typed per-event listeners plus a typed catch-all notification hook
- helpers for collecting prompt runs and handling extension UI requests in manual or headless mode
+2
View File
@@ -57,6 +57,7 @@ from .protocol import (
HookMessage,
ImageContent,
MessageEndEvent,
MessagesPage,
MessageStartEvent,
MessageUpdateEvent,
ModelCycleResult,
@@ -136,6 +137,7 @@ __all__ = [
"ListenerErrorEvent",
"ListenerErrorListener",
"MessageEndEvent",
"MessagesPage",
"MessageStartEvent",
"MessageUpdateEvent",
"ModelCost",
+183 -4
View File
@@ -1,5 +1,7 @@
from __future__ import annotations
import base64
import binascii
import json
import os
import queue
@@ -34,6 +36,7 @@ from .protocol import (
JsonObject,
JsonValue,
MessageEndEvent,
MessagesPage,
MessageStartEvent,
MessageUpdateEvent,
ModelCycleResult,
@@ -111,6 +114,101 @@ THistoryItem = TypeVar("THistoryItem")
_ASYNC_COMMANDS = frozenset({"prompt", "abort_and_prompt"})
_DEFAULT_ERROR_HISTORY_LIMIT = 128
_TODO_STATUS_VALUES = frozenset({"pending", "in_progress", "completed", "abandoned"})
_MAX_RPC_FRAME_BYTES = 1024 * 1024
_MAX_RPC_REASSEMBLED_BYTES = 64 * 1024 * 1024
_RPC_CHUNK_PAYLOAD_BYTES = 256 * 1024
@dataclass(slots=True)
class _PendingRpcChunks:
chunk_id: str
count: int
byte_length: int
next_index: int = 0
chunks: list[bytes] = field(default_factory=list)
received_bytes: int = 0
class _RpcFrameDecoder:
def __init__(self) -> None:
self._pending: _PendingRpcChunks | None = None
def push(self, value: object) -> JsonObject | None:
if not isinstance(value, dict) or value.get("type") != "rpc_chunk":
if self._pending is not None:
raise RpcError("RPC chunk sequence was interrupted")
if not isinstance(value, dict):
raise RpcError("RPC frame must be a JSON object")
return cast(JsonObject, value)
chunk_id = value.get("chunkId")
index = value.get("index")
count = value.get("count")
byte_length = value.get("byteLength")
data = value.get("data")
max_chunk_count = (
_MAX_RPC_REASSEMBLED_BYTES + _RPC_CHUNK_PAYLOAD_BYTES - 1
) // _RPC_CHUNK_PAYLOAD_BYTES
if (
not isinstance(chunk_id, str)
or not chunk_id
or len(chunk_id) > 128
or not isinstance(index, int)
or isinstance(index, bool)
or not isinstance(count, int)
or isinstance(count, bool)
or not isinstance(byte_length, int)
or isinstance(byte_length, bool)
or index < 0
or count < 2
or count > max_chunk_count
or index >= count
or byte_length <= _MAX_RPC_FRAME_BYTES
or byte_length > _MAX_RPC_REASSEMBLED_BYTES
or not isinstance(data, str)
or not data
):
raise RpcError("Invalid RPC chunk metadata")
try:
chunk = base64.b64decode(data, validate=True)
except (binascii.Error, ValueError) as exc:
raise RpcError("Invalid RPC chunk data") from exc
if base64.b64encode(chunk).decode("ascii") != data:
raise RpcError("Invalid RPC chunk data")
if len(chunk) > _RPC_CHUNK_PAYLOAD_BYTES:
raise RpcError("RPC chunk payload exceeds the transport limit")
if self._pending is None:
if index != 0:
raise RpcError("RPC chunk sequence must start at index 0")
self._pending = _PendingRpcChunks(chunk_id, count, byte_length)
pending = self._pending
if (
pending.chunk_id != chunk_id
or pending.count != count
or pending.byte_length != byte_length
or pending.next_index != index
):
raise RpcError("RPC chunk sequence mismatch")
pending.chunks.append(chunk)
pending.received_bytes += len(chunk)
pending.next_index += 1
if pending.received_bytes > pending.byte_length:
raise RpcError("RPC chunk sequence exceeds its declared length")
if pending.next_index < pending.count:
return None
if pending.received_bytes != pending.byte_length:
raise RpcError("RPC chunk sequence length mismatch")
self._pending = None
try:
decoded = b"".join(pending.chunks).decode("utf-8")
frame = json.loads(decoded)
except (UnicodeDecodeError, json.JSONDecodeError) as exc:
raise RpcError("Failed to decode reassembled RPC frame") from exc
if not isinstance(frame, dict):
raise RpcError("RPC frame must be a JSON object")
return cast(JsonObject, frame)
def _process_group_id(process: subprocess.Popen[Any]) -> int | None:
@@ -417,6 +515,10 @@ class RpcClient:
self._closed_error: BaseException | None = None
self._stopping = False
self._ready_received = False
self._ready_event: ReadyEvent | None = None
self._protocol_version = 1
self._protocol_v2_enabled = False
self._frame_decoder = _RpcFrameDecoder()
self._protocol_errors = _BoundedHistory[RpcProtocolError](
_DEFAULT_ERROR_HISTORY_LIMIT
)
@@ -468,6 +570,10 @@ class RpcClient:
self._stopping = False
self._closed_error = None
self._ready_received = False
self._ready_event = None
self._protocol_version = 1
self._protocol_v2_enabled = False
self._frame_decoder = _RpcFrameDecoder()
self._events.clear()
self._async_errors.clear()
self._scheduled_agent_runs = 0
@@ -529,6 +635,24 @@ class RpcClient:
f"Timed out waiting for RPC ready signal. Stderr: {stderr}"
)
ready_event = self._ready_event
if (
ready_event is not None
and ready_event.supported_protocol_versions is not None
and 2 in ready_event.supported_protocol_versions
and ready_event.max_frame_bytes == _MAX_RPC_FRAME_BYTES
and ready_event.max_reassembled_frame_bytes == _MAX_RPC_REASSEMBLED_BYTES
):
try:
self._protocol_v2_enabled = True
negotiation = self._request("negotiate_protocol", protocolVersion=2)
if negotiation.get("protocolVersion") != 2:
raise RpcError("RPC protocol v2 negotiation failed")
self._protocol_version = 2
except BaseException:
self.stop()
raise
if self._custom_tools:
self.set_custom_tools(self._custom_tools)
if self._host_uris:
@@ -882,9 +1006,55 @@ class RpcClient:
return self.set_todos(())
def get_messages(self) -> tuple[AgentMessage, ...]:
if self._protocol_version == 2:
messages: list[AgentMessage] = []
seen_cursors: set[str] = set()
total_messages: int | None = None
cursor: str | None = None
while True:
page = self.get_messages_page(cursor=cursor, limit=256)
if total_messages is not None and page.total_messages != total_messages:
raise RpcError(
"RPC message pagination returned an inconsistent total"
)
total_messages = page.total_messages
messages.extend(page.messages)
cursor = page.next_cursor
if cursor is None:
break
if cursor in seen_cursors:
raise RpcError("RPC message pagination repeated a cursor")
seen_cursors.add(cursor)
if len(messages) != total_messages:
raise RpcError(
"RPC message pagination ended before the advertised total"
)
return tuple(messages)
payload = self._request("get_messages")
return parse_agent_messages(cast(JsonValue | None, payload.get("messages")))
def get_messages_page(
self, *, cursor: str | None = None, limit: int | None = None
) -> MessagesPage:
payload = self._request("get_messages_page", cursor=cursor, limit=limit)
raw_total = payload.get("totalMessages")
if (
not isinstance(raw_total, int)
or isinstance(raw_total, bool)
or raw_total < 0
):
raise RpcError("get_messages_page response has an invalid totalMessages")
raw_cursor = payload.get("nextCursor")
if raw_cursor is not None and not isinstance(raw_cursor, str):
raise RpcError("get_messages_page response has an invalid nextCursor")
return MessagesPage(
messages=parse_agent_messages(
cast(JsonValue | None, payload.get("messages"))
),
total_messages=raw_total,
next_cursor=raw_cursor,
)
def set_custom_tools(self, tools: Sequence[HostTool[Any, Any]]) -> tuple[str, ...]:
self._custom_tools = tuple(tools)
if self._process is None:
@@ -1100,9 +1270,8 @@ class RpcClient:
def _complete_agent_end_messages(
events: tuple[RpcAgentEvent, ...], terminal: AgentEndEvent
) -> tuple[AgentMessage, ...]:
if (
terminal.message_count is None
or terminal.message_count <= len(terminal.messages)
if terminal.message_count is None or terminal.message_count <= len(
terminal.messages
):
return terminal.messages
@@ -1632,7 +1801,7 @@ class RpcClient:
continue
try:
payload = cast(JsonObject, json.loads(stripped))
raw_payload = json.loads(stripped)
except json.JSONDecodeError as exc:
snippet = stripped
if len(snippet) > 240:
@@ -1640,6 +1809,15 @@ class RpcClient:
raise RpcError(
f"Failed to decode RPC output on line {line_number}: {exc}. Frame: {snippet!r}"
) from exc
if (
isinstance(raw_payload, dict)
and raw_payload.get("type") == "rpc_chunk"
and not self._protocol_v2_enabled
):
raise RpcError("RPC chunk received before protocol negotiation")
payload = self._frame_decoder.push(raw_payload)
if payload is None:
continue
if payload.get("type") == "response":
self._handle_response(payload)
continue
@@ -1666,6 +1844,7 @@ class RpcClient:
)
if isinstance(notification, ReadyEvent):
self._ready_event = notification
self._ready_received = True
self._ready.set()
self._dispatch_listeners(
+28 -1
View File
@@ -851,9 +851,20 @@ class SessionStats:
@dataclass(slots=True, frozen=True)
class ReadyEvent:
protocol_version: int | None = None
supported_protocol_versions: tuple[int, ...] | None = None
max_frame_bytes: int | None = None
max_reassembled_frame_bytes: int | None = None
type: Literal["ready"] = "ready"
@dataclass(slots=True, frozen=True)
class MessagesPage:
messages: tuple[AgentMessage, ...]
total_messages: int
next_cursor: str | None
@dataclass(slots=True, frozen=True)
class ExtensionUiRequest:
id: str
@@ -1490,7 +1501,23 @@ def parse_extension_error(payload: JsonObject) -> ExtensionError:
def parse_notification(payload: JsonObject) -> RpcNotification:
event_type = payload.get("type")
if event_type == "ready":
return ReadyEvent()
raw_versions = payload.get("supportedProtocolVersions")
supported_versions: tuple[int, ...] | None = None
if raw_versions is not None:
if not isinstance(raw_versions, list) or any(
not isinstance(version, int) or isinstance(version, bool)
for version in raw_versions
):
raise ValueError("ready.supportedProtocolVersions must be integers")
supported_versions = tuple(raw_versions)
return ReadyEvent(
protocol_version=_optional_int(payload, "protocolVersion"),
supported_protocol_versions=supported_versions,
max_frame_bytes=_optional_int(payload, "maxFrameBytes"),
max_reassembled_frame_bytes=_optional_int(
payload, "maxReassembledFrameBytes"
),
)
if event_type == "extension_ui_request":
return parse_extension_ui_request(payload)
if event_type == "extension_error":
+98
View File
@@ -442,6 +442,97 @@ FAKE_SERVER = textwrap.dedent(
"""
)
V2_MESSAGES_SERVER = textwrap.dedent(
"""
import base64
import json
import sys
message = {
"role": "user",
"content": [{"type": "text", "text": "x" * (1024 * 1024)}],
"timestamp": 1,
}
def emit(payload):
encoded = json.dumps(payload, separators=(",", ":")).encode("utf-8")
if len(encoded) <= 1024 * 1024:
print(encoded.decode("utf-8"), flush=True)
return
chunk_size = 256 * 1024
count = (len(encoded) + chunk_size - 1) // chunk_size
for index in range(count):
chunk = encoded[index * chunk_size : (index + 1) * chunk_size]
print(
json.dumps(
{
"type": "rpc_chunk",
"chunkId": "test-page",
"index": index,
"count": count,
"byteLength": len(encoded),
"data": base64.b64encode(chunk).decode("ascii"),
},
separators=(",", ":"),
),
flush=True,
)
print(
json.dumps(
{
"type": "ready",
"protocolVersion": 1,
"supportedProtocolVersions": [1, 2],
"maxFrameBytes": 1024 * 1024,
"maxReassembledFrameBytes": 64 * 1024 * 1024,
}
),
flush=True,
)
for raw_line in sys.stdin:
command = json.loads(raw_line)
request_id = command["id"]
command_type = command["type"]
if command_type == "negotiate_protocol":
emit(
{
"id": request_id,
"type": "response",
"command": command_type,
"success": True,
"data": {"protocolVersion": 2},
}
)
elif command_type == "get_messages_page":
emit(
{
"id": request_id,
"type": "response",
"command": command_type,
"success": True,
"data": {
"messages": [message],
"totalMessages": 1,
"nextCursor": None,
},
}
)
else:
emit(
{
"id": request_id,
"type": "response",
"command": command_type,
"success": False,
"error": f"unexpected command: {command_type}",
}
)
"""
)
IDLESS_ERROR_SERVER = textwrap.dedent(
"""
import json
@@ -858,6 +949,13 @@ class RpcClientTests(unittest.TestCase):
client.wait_for_idle(timeout=2.0)
self.assertEqual(client.get_last_assistant_text(), "pong")
def test_protocol_v2_reassembles_chunked_message_pages(self) -> None:
with self.make_client(server=V2_MESSAGES_SERVER) as client:
messages = client.get_messages()
self.assertEqual(len(messages), 1)
self.assertEqual(len(messages[0]["content"][0]["text"]), 1024 * 1024)
def test_collect_events_returns_turn_events(self) -> None:
with self.make_client() as client:
client.prompt("slow")