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
+6 -1
View File
@@ -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
+1 -1
View File
@@ -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
@@ -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<string, { controller: AbortController }>();
#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<RpcMessagesPage> {
const response = await this.#send({ type: "get_messages_page", ...options });
return this.#getData<RpcMessagesPage>(response);
}
/** Get all messages, draining stable pages when protocol v2 is available. */
async getMessages(): Promise<AgentMessage[]> {
if (this.#protocolVersion === 2) {
const messages: AgentMessage[] = [];
const seenCursors = new Set<string>();
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;
}
@@ -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,
};
}
@@ -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
// =================================================================
@@ -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
| {
+20
View File
@@ -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",
@@ -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 () => {
@@ -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<string, unknown> | undefined;
let stateResponse: Record<string, unknown> | undefined;
let pageResponse: Record<string, unknown> | undefined;
for await (const frame of readJsonl<unknown>(child.stdout as ReadableStream<Uint8Array>)) {
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);
});
@@ -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();
});
});
+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")