diff --git a/docs/rpc.md b/docs/rpc.md index 064a828c2..05bf375ed 100644 --- a/docs/rpc.md +++ b/docs/rpc.md @@ -64,7 +64,7 @@ After the success response, oversized stdout objects are emitted losslessly as a } ``` -Clients MUST validate `chunkId`, `index`, `count`, and `byteLength`, reject interleaved or interrupted sequences, enforce the advertised reassembly limit, concatenate decoded bytes in index order, decode them as strict UTF-8, and parse the result as one JSON object. The exported `RpcFrameDecoder` implements this validation. `RpcClient` negotiates v2 automatically when the ready frame advertises it. +Clients MUST validate `chunkId`, `index`, `count`, and `byteLength`, reject interleaved or interrupted sequences, enforce the advertised reassembly limit, concatenate decoded bytes in index order, decode them as strict UTF-8, and parse the result as one JSON object. The exported TypeScript `RpcFrameDecoder` implements this validation. The bundled TypeScript and Python `RpcClient` implementations negotiate v2 automatically when the ready frame advertises it. Legacy clients may ignore the added ready fields and remain on v1. V1 retains its bounded fallback behavior for oversized output. Frames above the v2 reassembly ceiling still fail explicitly; large history APIs should use pagination rather than depending on arbitrarily large logical frames. @@ -185,6 +185,11 @@ correlate it via `id`. Ordering across concurrent commands is not guaranteed ### Messages - `{ id?, type: "get_messages" }` +- `{ id?, type: "get_messages_page", cursor?: string, limit?: number }` + +`get_messages_page` returns a stable chronological page with `messages`, `totalMessages`, and an opaque `nextCursor` when more messages remain. Cursors are bound to the session ID, durable leaf, and message count. The server rejects stale cursors if the session changes between requests, and refuses to start a paging walk while the session is streaming or compacting. Pages contain at most 256 messages and normally stay below the v1 physical-frame ceiling; an individually oversized message requires negotiated v2 framing. + +The bundled TypeScript `RpcClient.getMessages()` and Python `RpcClient.get_messages()` drain this paged endpoint automatically after negotiating v2. They retain the legacy monolithic command when connected to a v1 server. Hosts that render incrementally can call `getMessagesPage()` or `get_messages_page()` directly. ### Login diff --git a/packages/coding-agent/CHANGELOG.md b/packages/coding-agent/CHANGELOG.md index 699d398d5..16603e5d1 100644 --- a/packages/coding-agent/CHANGELOG.md +++ b/packages/coding-agent/CHANGELOG.md @@ -4,7 +4,7 @@ ### Added -- Added opt-in RPC protocol v2 negotiation with bounded, lossless chunking for stdout objects up to 64 MiB. Legacy JSONL clients remain on protocol v1, while the TypeScript RPC client negotiates and reassembles v2 automatically. +- Added opt-in RPC protocol v2 negotiation with bounded, lossless chunking for stdout objects up to 64 MiB, plus stable cursor-based message pages for histories that should not travel as one response. Legacy JSONL clients remain on protocol v1, while the bundled TypeScript and Python RPC clients negotiate, reassemble, and drain message pages automatically. ## [17.0.8] - 2026-07-22 diff --git a/packages/coding-agent/src/modes/rpc/rpc-client.ts b/packages/coding-agent/src/modes/rpc/rpc-client.ts index e259016bd..3eb9ef3e6 100644 --- a/packages/coding-agent/src/modes/rpc/rpc-client.ts +++ b/packages/coding-agent/src/modes/rpc/rpc-client.ts @@ -12,7 +12,8 @@ import { isRecord, ptree, readJsonl } from "@oh-my-pi/pi-utils"; import type { FileSink } from "bun"; import type { BashResult } from "../../exec/bash-executor"; import type { AgentSessionEvent, SessionStats } from "../../session/agent-session"; -import { MAX_RPC_FRAME_BYTES, MAX_RPC_REASSEMBLED_BYTES, RpcFrameDecoder } from "./rpc-frame"; +import { MAX_RPC_FRAME_BYTES, MAX_RPC_REASSEMBLED_BYTES, RpcFrameDecoder, type RpcProtocolVersion } from "./rpc-frame"; +import type { RpcMessagesPage, RpcMessagesPageOptions } from "./rpc-messages"; import type { RpcAvailableCommandsUpdateFrame, RpcAvailableSlashCommand, @@ -229,6 +230,7 @@ export class RpcClient { #customTools: RpcClientCustomTool[] = []; #pendingHostToolCalls = new Map(); #requestId = 0; + #protocolVersion: RpcProtocolVersion = 1; #extensionUiListeners: Set<(req: RpcExtensionUIRequest) => void> = new Set(); #abortController = new AbortController(); @@ -253,6 +255,7 @@ export class RpcClient { // Mint a fresh controller so a previous stop()'s abort does not // short-circuit the new stdout reader (issue #4079). this.#abortController = new AbortController(); + this.#protocolVersion = 1; const cliPath = this.options.cliPath ?? "dist/cli.js"; const args = ["--mode", "rpc"]; @@ -387,6 +390,7 @@ export class RpcClient { response.data.protocolVersion !== 2 ) throw new Error("RPC protocol v2 negotiation failed"); + this.#protocolVersion = 2; } if (this.#customTools.length > 0) { await this.setCustomTools(this.#customTools); @@ -773,9 +777,38 @@ export class RpcClient { } /** - * Get all messages in the session. + * Get one stable, byte-bounded message page. */ + async getMessagesPage(options: RpcMessagesPageOptions = {}): Promise { + const response = await this.#send({ type: "get_messages_page", ...options }); + return this.#getData(response); + } + + /** Get all messages, draining stable pages when protocol v2 is available. */ async getMessages(): Promise { + if (this.#protocolVersion === 2) { + const messages: AgentMessage[] = []; + const seenCursors = new Set(); + let totalMessages: number | undefined; + let cursor: string | undefined; + do { + const page = await this.getMessagesPage({ cursor, limit: 256 }); + if ( + !Number.isSafeInteger(page.totalMessages) || + page.totalMessages < 0 || + (totalMessages !== undefined && page.totalMessages !== totalMessages) + ) + throw new Error("RPC message pagination returned an inconsistent total"); + totalMessages = page.totalMessages; + messages.push(...page.messages); + cursor = page.nextCursor; + if (cursor && seenCursors.has(cursor)) throw new Error("RPC message pagination repeated a cursor"); + if (cursor) seenCursors.add(cursor); + } while (cursor); + if (messages.length !== totalMessages) + throw new Error("RPC message pagination ended before the advertised total"); + return messages; + } const response = await this.#send({ type: "get_messages" }); return this.#getData<{ messages: AgentMessage[] }>(response).messages; } diff --git a/packages/coding-agent/src/modes/rpc/rpc-messages.ts b/packages/coding-agent/src/modes/rpc/rpc-messages.ts new file mode 100644 index 000000000..383733d7d --- /dev/null +++ b/packages/coding-agent/src/modes/rpc/rpc-messages.ts @@ -0,0 +1,109 @@ +import type { AgentMessage } from "@oh-my-pi/pi-agent-core"; +import { isRecord } from "@oh-my-pi/pi-utils"; + +const DEFAULT_RPC_MESSAGE_PAGE_LIMIT = 100; +const MAX_RPC_MESSAGE_PAGE_LIMIT = 256; +const MAX_RPC_MESSAGE_PAGE_BYTES = 768 * 1024; +const MAX_RPC_MESSAGE_CURSOR_CHARS = 2048; + +export interface RpcMessageSnapshot { + sessionId: string; + leafId: string | null; + messageCount: number; +} + +export interface RpcMessagesPage { + messages: AgentMessage[]; + nextCursor?: string; + totalMessages: number; +} + +interface RpcMessageCursorPayload extends RpcMessageSnapshot { + version: 1; + offset: number; +} + +export interface RpcMessagesPageOptions { + cursor?: string; + limit?: number; +} + +function encodeCursor(snapshot: RpcMessageSnapshot, offset: number): string { + const payload: RpcMessageCursorPayload = { version: 1, ...snapshot, offset }; + return Buffer.from(JSON.stringify(payload), "utf8").toString("base64url"); +} + +function decodeCursor(cursor: string): RpcMessageCursorPayload { + if (cursor.length === 0 || cursor.length > MAX_RPC_MESSAGE_CURSOR_CHARS || !/^[A-Za-z0-9_-]+$/.test(cursor)) + throw new Error("Invalid RPC message cursor"); + const bytes = Buffer.from(cursor, "base64url"); + if (bytes.toString("base64url") !== cursor) throw new Error("Invalid RPC message cursor"); + let value: unknown; + try { + value = JSON.parse(new TextDecoder("utf-8", { fatal: true }).decode(bytes)); + } catch { + throw new Error("Invalid RPC message cursor"); + } + if (!isRecord(value)) throw new Error("Invalid RPC message cursor"); + const { version, sessionId, leafId, messageCount, offset } = value; + if ( + version !== 1 || + typeof sessionId !== "string" || + sessionId.length === 0 || + sessionId.length > 256 || + !((typeof leafId === "string" && leafId.length > 0 && leafId.length <= 256) || leafId === null) || + typeof messageCount !== "number" || + !Number.isSafeInteger(messageCount) || + messageCount < 0 || + typeof offset !== "number" || + !Number.isSafeInteger(offset) || + offset < 0 || + offset > messageCount + ) + throw new Error("Invalid RPC message cursor"); + return { version, sessionId, leafId, messageCount, offset }; +} + +function sameSnapshot(cursor: RpcMessageCursorPayload, snapshot: RpcMessageSnapshot): boolean { + return ( + cursor.sessionId === snapshot.sessionId && + cursor.leafId === snapshot.leafId && + cursor.messageCount === snapshot.messageCount + ); +} + +/** Page one stable in-memory message snapshot without crossing the v1 frame budget. */ +export function pageRpcMessages( + messages: readonly AgentMessage[], + snapshot: RpcMessageSnapshot, + options: RpcMessagesPageOptions = {}, +): RpcMessagesPage { + if (snapshot.messageCount !== messages.length) + throw new Error("RPC message snapshot does not match current messages"); + const limit = options.limit ?? DEFAULT_RPC_MESSAGE_PAGE_LIMIT; + if (!Number.isSafeInteger(limit) || limit < 1 || limit > MAX_RPC_MESSAGE_PAGE_LIMIT) + throw new Error(`RPC message page limit must be between 1 and ${MAX_RPC_MESSAGE_PAGE_LIMIT}`); + let offset = 0; + if (options.cursor !== undefined) { + const cursor = decodeCursor(options.cursor); + if (!sameSnapshot(cursor, snapshot)) throw new Error("RPC message cursor is stale"); + offset = cursor.offset; + } + + const page: AgentMessage[] = []; + let pageBytes = 2; + while (offset + page.length < messages.length && page.length < limit) { + const message = messages[offset + page.length]; + const messageBytes = Buffer.byteLength(JSON.stringify(message), "utf8") + (page.length === 0 ? 0 : 1); + if (page.length > 0 && pageBytes + messageBytes > MAX_RPC_MESSAGE_PAGE_BYTES) break; + page.push(message); + pageBytes += messageBytes; + } + + const nextOffset = offset + page.length; + return { + messages: page, + ...(nextOffset < messages.length ? { nextCursor: encodeCursor(snapshot, nextOffset) } : {}), + totalMessages: messages.length, + }; +} diff --git a/packages/coding-agent/src/modes/rpc/rpc-mode.ts b/packages/coding-agent/src/modes/rpc/rpc-mode.ts index 6d458363c..051b252e6 100644 --- a/packages/coding-agent/src/modes/rpc/rpc-mode.ts +++ b/packages/coding-agent/src/modes/rpc/rpc-mode.ts @@ -36,6 +36,7 @@ import { isRpcHostToolResult, isRpcHostToolUpdate, RpcHostToolBridge } from "./h import { isRpcHostUriResult, RpcHostUriBridge } from "./host-uris"; import { MAX_RPC_FRAME_BYTES, MAX_RPC_REASSEMBLED_BYTES, RpcFrameEncoder } from "./rpc-frame"; import { claimRpcInput } from "./rpc-input"; +import { pageRpcMessages } from "./rpc-messages"; import { RpcSubagentRegistry, readRpcSubagentTranscript } from "./rpc-subagents"; import type { RpcCommand, @@ -1293,6 +1294,33 @@ export async function runRpcMode( return success(id, "get_messages", { messages: session.messages }); } + case "get_messages_page": { + if (session.isStreaming || session.isCompacting) + return error(id, "get_messages_page", "Cannot page messages while the session is changing"); + const messages = session.messages; + try { + return success( + id, + "get_messages_page", + pageRpcMessages( + messages, + { + sessionId: session.sessionId, + leafId: session.sessionManager.getLeafId(), + messageCount: messages.length, + }, + { cursor: command.cursor, limit: command.limit }, + ), + ); + } catch (pageError) { + return error( + id, + "get_messages_page", + pageError instanceof Error ? pageError.message : String(pageError), + ); + } + } + // ================================================================= // Login // ================================================================= diff --git a/packages/coding-agent/src/modes/rpc/rpc-types.ts b/packages/coding-agent/src/modes/rpc/rpc-types.ts index ee75ccabe..babb08ccd 100644 --- a/packages/coding-agent/src/modes/rpc/rpc-types.ts +++ b/packages/coding-agent/src/modes/rpc/rpc-types.ts @@ -19,6 +19,7 @@ import type { SubagentProgressPayload, } from "../../task"; import type { TodoPhase } from "../../tools/todo"; +import type { RpcMessagesPage } from "./rpc-messages"; // ============================================================================ // RPC Commands (stdin) @@ -84,6 +85,7 @@ export type RpcCommand = // Messages | { id?: string; type: "get_messages" } + | { id?: string; type: "get_messages_page"; cursor?: string; limit?: number } // Login | { id?: string; type: "get_login_providers" } @@ -313,6 +315,7 @@ export type RpcResponse = // Messages | { id?: string; type: "response"; command: "get_messages"; success: true; data: { messages: AgentMessage[] } } + | { id?: string; type: "response"; command: "get_messages_page"; success: true; data: RpcMessagesPage } // Login | { diff --git a/packages/coding-agent/test/fixtures/mock-rpc-agent.ts b/packages/coding-agent/test/fixtures/mock-rpc-agent.ts index 9be1d0ff0..0e93424dd 100755 --- a/packages/coding-agent/test/fixtures/mock-rpc-agent.ts +++ b/packages/coding-agent/test/fixtures/mock-rpc-agent.ts @@ -79,6 +79,26 @@ for await (const raw of console) { protocolV2Enabled = true; continue; } + if (frame.type === "get_messages_page") { + const first = frame.cursor === undefined; + writeFrame({ + id, + type: "response", + command: frame.type, + success: true, + data: first + ? { + messages: [{ role: "user", content: "first", timestamp: 1 }], + nextCursor: "second-page", + totalMessages: 2, + } + : { + messages: [{ role: "assistant", content: [{ type: "text", text: "second" }], timestamp: 2 }], + totalMessages: 2, + }, + }); + continue; + } writeFrame({ id, type: "response", diff --git a/packages/coding-agent/test/rpc-client.restart.test.ts b/packages/coding-agent/test/rpc-client.restart.test.ts index 90ee2ab65..16aedbba7 100644 --- a/packages/coding-agent/test/rpc-client.restart.test.ts +++ b/packages/coding-agent/test/rpc-client.restart.test.ts @@ -24,6 +24,10 @@ describe("RpcClient lifecycle (issue #4079 B)", () => { await client.start(); const state = (await client.getState()) as unknown as { payload: string }; expect(state.payload).toBe("😀".repeat(400_000)); + expect((await client.getMessages()) as unknown).toEqual([ + { role: "user", content: "first", timestamp: 1 }, + { role: "assistant", content: [{ type: "text", text: "second" }], timestamp: 2 }, + ]); }, 20_000); test("start() succeeds a second time after stop() on the same instance", async () => { diff --git a/packages/coding-agent/test/rpc-malformed-input.test.ts b/packages/coding-agent/test/rpc-malformed-input.test.ts index 31bd39507..c6ef029bd 100644 --- a/packages/coding-agent/test/rpc-malformed-input.test.ts +++ b/packages/coding-agent/test/rpc-malformed-input.test.ts @@ -26,10 +26,12 @@ describe("RPC mode malformed stdin", () => { // crashed the generator before the second was ever read. child.stdin.write("this is not json\n"); child.stdin.write(`${JSON.stringify({ type: "get_state", id: "probe" })}\n`); + child.stdin.write(`${JSON.stringify({ type: "get_messages_page", id: "page-probe", limit: 1 })}\n`); await child.stdin.flush(); let parseError: Record | undefined; let stateResponse: Record | undefined; + let pageResponse: Record | undefined; for await (const frame of readJsonl(child.stdout as ReadableStream)) { if (!isRecord(frame)) continue; @@ -38,8 +40,9 @@ describe("RPC mode malformed stdin", () => { } if (frame.type === "response" && frame.id === "probe") { stateResponse = frame; - break; } + if (frame.type === "response" && frame.id === "page-probe") pageResponse = frame; + if (stateResponse && pageResponse) break; } child.stdin.end(); @@ -50,5 +53,9 @@ describe("RPC mode malformed stdin", () => { expect(String(parseError?.error)).toContain("Failed to parse command"); expect(stateResponse).toBeDefined(); expect(stateResponse?.success).toBe(true); + expect(pageResponse).toMatchObject({ + success: true, + data: { messages: [], totalMessages: 0 }, + }); }, 30000); }); diff --git a/packages/coding-agent/test/rpc-messages.test.ts b/packages/coding-agent/test/rpc-messages.test.ts new file mode 100644 index 000000000..d78e0008f --- /dev/null +++ b/packages/coding-agent/test/rpc-messages.test.ts @@ -0,0 +1,64 @@ +import { describe, expect, it } from "bun:test"; +import type { AgentMessage } from "@oh-my-pi/pi-agent-core"; +import { encodeRpcFrame, MAX_RPC_FRAME_BYTES } from "../src/modes/rpc/rpc-frame"; +import { pageRpcMessages, type RpcMessageSnapshot } from "../src/modes/rpc/rpc-messages"; + +function message(index: number, bytes = 32 * 1024): AgentMessage { + return { role: "user", content: `${index}:${"x".repeat(bytes)}`, timestamp: index }; +} + +const snapshot: RpcMessageSnapshot = { + sessionId: "session-1", + leafId: "leaf-1", + messageCount: 60, +}; + +describe("RPC message pagination", () => { + it("reconstructs a large history from v1-safe pages without loss or overlap", () => { + const messages = Array.from({ length: snapshot.messageCount }, (_, index) => message(index)); + const reconstructed: AgentMessage[] = []; + let cursor: string | undefined; + let pageCount = 0; + + do { + const page = pageRpcMessages(messages, snapshot, { cursor, limit: 256 }); + const encoded = encodeRpcFrame({ + id: `page-${pageCount}`, + type: "response", + command: "get_messages_page", + success: true, + data: page, + }); + expect(Buffer.byteLength(encoded, "utf8")).toBeLessThanOrEqual(MAX_RPC_FRAME_BYTES); + expect(JSON.parse(encoded).success).toBe(true); + reconstructed.push(...page.messages); + cursor = page.nextCursor; + pageCount++; + } while (cursor); + + expect(pageCount).toBeGreaterThan(1); + expect(reconstructed).toEqual(messages); + }); + + it("rejects a cursor after the session snapshot changes", () => { + const messages = Array.from({ length: snapshot.messageCount }, (_, index) => message(index, 1024)); + const first = pageRpcMessages(messages, snapshot, { limit: 5 }); + expect(first.nextCursor).toBeDefined(); + + expect(() => + pageRpcMessages(messages, { ...snapshot, leafId: "leaf-2" }, { cursor: first.nextCursor, limit: 5 }), + ).toThrow("RPC message cursor is stale"); + }); + + it("returns one individually oversized message so negotiated v2 can carry it losslessly", () => { + const messages = [message(0, 2 * 1024 * 1024), message(1, 128)]; + const first = pageRpcMessages( + messages, + { sessionId: "session-2", leafId: "leaf-2", messageCount: messages.length }, + { limit: 10 }, + ); + + expect(first.messages).toEqual([messages[0]]); + expect(first.nextCursor).toBeDefined(); + }); +}); diff --git a/python/omp-rpc/README.md b/python/omp-rpc/README.md index 51914f771..903f8d7c3 100644 --- a/python/omp-rpc/README.md +++ b/python/omp-rpc/README.md @@ -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 diff --git a/python/omp-rpc/src/omp_rpc/__init__.py b/python/omp-rpc/src/omp_rpc/__init__.py index f596d2161..d1fe16ec3 100644 --- a/python/omp-rpc/src/omp_rpc/__init__.py +++ b/python/omp-rpc/src/omp_rpc/__init__.py @@ -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", diff --git a/python/omp-rpc/src/omp_rpc/client.py b/python/omp-rpc/src/omp_rpc/client.py index 1418f59f1..7069b09ed 100644 --- a/python/omp-rpc/src/omp_rpc/client.py +++ b/python/omp-rpc/src/omp_rpc/client.py @@ -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( diff --git a/python/omp-rpc/src/omp_rpc/protocol.py b/python/omp-rpc/src/omp_rpc/protocol.py index 0034f3dd0..f14aea8f0 100644 --- a/python/omp-rpc/src/omp_rpc/protocol.py +++ b/python/omp-rpc/src/omp_rpc/protocol.py @@ -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": diff --git a/python/omp-rpc/tests/test_client.py b/python/omp-rpc/tests/test_client.py index 68114923d..e9d3b0193 100644 --- a/python/omp-rpc/tests/test_client.py +++ b/python/omp-rpc/tests/test_client.py @@ -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")