feat(rpc): page stable message histories
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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":
|
||||
|
||||
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user