Merge PR #6330: feat(coding-agent): add lossless paged RPC transport (@wolfiesch)

This commit is contained in:
can1357
2026-07-23 12:27:38 +02:00
17 changed files with 1146 additions and 22 deletions
+45 -3
View File
@@ -25,15 +25,48 @@ Behavior notes:
- RPC mode disables automatic session title generation by default to avoid an extra model call. - RPC mode disables automatic session title generation by default to avoid an extra model call.
- RPC mode resets workflow-altering `todo.*`, `task.*`, `memory.backend`/`memories.enabled`, `advisor.*`, `async.*`, and `bash.autoBackground.*` settings to their built-in defaults instead of inheriting user overrides. - RPC mode resets workflow-altering `todo.*`, `task.*`, `memory.backend`/`memories.enabled`, `advisor.*`, `async.*`, and `bash.autoBackground.*` settings to their built-in defaults instead of inheriting user overrides.
- The process reads stdin as JSONL (`readJsonl(Bun.stdin.stream())`). - The process reads stdin as JSONL (`readJsonl(Bun.stdin.stream())`).
- At startup it writes `{ "type": "ready" }` before processing commands. - At startup it writes a `ready` frame before processing commands. The frame advertises supported protocol versions and transport limits.
- When stdin closes, pending host-tool calls and host-URI requests are rejected and the process exits with code `0`. - When stdin closes, pending host-tool calls and host-URI requests are rejected and the process exits with code `0`.
- Responses/events are written as one JSON object per line. - Responses/events are written as one JSON object per line.
## Transport and Framing ## Transport and Framing
Each frame is a single JSON object followed by `\n`. Protocol v1 frames are a single JSON object followed by `\n`. Every physical JSONL frame is limited to 1 MiB.
There is no envelope beyond the object shape itself. The initial ready frame uses protocol v1 and advertises the opt-in lossless transport:
```json
{
"type": "ready",
"protocolVersion": 1,
"supportedProtocolVersions": [1, 2],
"maxFrameBytes": 1048576,
"maxReassembledFrameBytes": 67108864
}
```
Clients that support protocol v2 SHOULD immediately send:
```json
{ "id": "protocol-1", "type": "negotiate_protocol", "protocolVersion": 2 }
```
After the success response, oversized stdout objects are emitted losslessly as an uninterrupted sequence of `rpc_chunk` frames. Each chunk carries a base64 segment of the original UTF-8 JSON object:
```json
{
"type": "rpc_chunk",
"chunkId": "rpc-1",
"index": 0,
"count": 7,
"byteLength": 1600042,
"data": "eyJ0eXBlIjoicmVzcG9uc2UiLC4uLn0="
}
```
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.
### Outbound frame categories (stdout) ### Outbound frame categories (stdout)
@@ -84,6 +117,10 @@ Important edge behavior from runtime:
- `{ id?, type: "abort_and_prompt", message: string, images?: ImageContent[] }` - `{ id?, type: "abort_and_prompt", message: string, images?: ImageContent[] }`
- `{ id?, type: "new_session", parentSession?: string }` - `{ id?, type: "new_session", parentSession?: string }`
### Protocol
- `{ id?, type: "negotiate_protocol", protocolVersion: 2 }`
### State ### State
- `{ id?, type: "get_state" }` - `{ id?, type: "get_state" }`
@@ -148,6 +185,11 @@ correlate it via `id`. Ordering across concurrent commands is not guaranteed
### Messages ### Messages
- `{ id?, type: "get_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. A v1 caller can page ordinary histories, but an individual message whose response exceeds that ceiling produces an overflow error; retrieving it losslessly 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, and fall back to its best-effort snapshot if the session begins streaming or compacting during a page walk. Direct `getMessagesPage()` and `get_messages_page()` calls remain strict so incremental hosts never mix snapshots silently.
### Login ### Login
+1
View File
@@ -27,6 +27,7 @@
- Reduced format-on-write latency by avoiding cold language-server startup when diagnostics are disabled. - Reduced format-on-write latency by avoiding cold language-server startup when diagnostics are disabled.
- Rewrote the `/guided-goal` interviewer rubric around loop-engineering: deterministic success criteria, verification commands, attempt caps, scope boundaries, and stop conditions. Ready objectives must use the five-section structured markdown form. - Rewrote the `/guided-goal` interviewer rubric around loop-engineering: deterministic success criteria, verification commands, attempt caps, scope boundaries, and stop conditions. Ready objectives must use the five-section structured markdown form.
- Added `task.isolation.apply` (default `true`) to choose whether successful isolated `task` runs automatically apply their changes to the parent checkout or retain patch/branch artifacts for later integration. - Added `task.isolation.apply` (default `true`) to choose whether successful isolated `task` runs automatically apply their changes to the parent checkout or retain patch/branch artifacts for later integration.
- 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 ## [17.0.8] - 2026-07-22
@@ -12,6 +12,8 @@ import { isRecord, ptree, readJsonl } from "@oh-my-pi/pi-utils";
import type { FileSink } from "bun"; import type { FileSink } from "bun";
import type { BashResult } from "../../exec/bash-executor"; import type { BashResult } from "../../exec/bash-executor";
import type { AgentSessionEvent, SessionStats } from "../../session/agent-session"; import type { AgentSessionEvent, SessionStats } from "../../session/agent-session";
import { MAX_RPC_FRAME_BYTES, MAX_RPC_REASSEMBLED_BYTES, RpcFrameDecoder, type RpcProtocolVersion } from "./rpc-frame";
import { RPC_MESSAGES_PAGE_BUSY_ERROR, type RpcMessagesPage, type RpcMessagesPageOptions } from "./rpc-messages";
import type { import type {
RpcAvailableCommandsUpdateFrame, RpcAvailableCommandsUpdateFrame,
RpcAvailableSlashCommand, RpcAvailableSlashCommand,
@@ -136,6 +138,16 @@ function isRpcResponse(value: unknown): value is RpcResponse {
return true; return true;
} }
function supportsRpcProtocolV2(value: Record<string, unknown>): boolean {
return (
value.type === "ready" &&
Array.isArray(value.supportedProtocolVersions) &&
value.supportedProtocolVersions.includes(2) &&
value.maxFrameBytes === MAX_RPC_FRAME_BYTES &&
value.maxReassembledFrameBytes === MAX_RPC_REASSEMBLED_BYTES
);
}
function isAgentEvent(value: unknown): value is AgentEvent { function isAgentEvent(value: unknown): value is AgentEvent {
if (!isRecord(value)) return false; if (!isRecord(value)) return false;
const type = value.type; const type = value.type;
@@ -218,6 +230,7 @@ export class RpcClient {
#customTools: RpcClientCustomTool[] = []; #customTools: RpcClientCustomTool[] = [];
#pendingHostToolCalls = new Map<string, { controller: AbortController }>(); #pendingHostToolCalls = new Map<string, { controller: AbortController }>();
#requestId = 0; #requestId = 0;
#protocolVersion: RpcProtocolVersion = 1;
#extensionUiListeners: Set<(req: RpcExtensionUIRequest) => void> = new Set(); #extensionUiListeners: Set<(req: RpcExtensionUIRequest) => void> = new Set();
#abortController = new AbortController(); #abortController = new AbortController();
@@ -242,6 +255,7 @@ export class RpcClient {
// Mint a fresh controller so a previous stop()'s abort does not // Mint a fresh controller so a previous stop()'s abort does not
// short-circuit the new stdout reader (issue #4079). // short-circuit the new stdout reader (issue #4079).
this.#abortController = new AbortController(); this.#abortController = new AbortController();
this.#protocolVersion = 1;
const cliPath = this.options.cliPath ?? "dist/cli.js"; const cliPath = this.options.cliPath ?? "dist/cli.js";
const args = ["--mode", "rpc"]; const args = ["--mode", "rpc"];
@@ -269,6 +283,9 @@ export class RpcClient {
// Wait for the "ready" signal or process exit // Wait for the "ready" signal or process exit
const { promise: readyPromise, resolve: readyResolve, reject: readyReject } = Promise.withResolvers<void>(); const { promise: readyPromise, resolve: readyResolve, reject: readyReject } = Promise.withResolvers<void>();
let readySettled = false; let readySettled = false;
let protocolV2Supported = false;
let protocolV2Enabled = false;
const frameDecoder = new RpcFrameDecoder();
const reapAfterOutputFailure = async (error: Error) => { const reapAfterOutputFailure = async (error: Error) => {
if (this.#process !== child) return; if (this.#process !== child) return;
@@ -294,11 +311,15 @@ export class RpcClient {
void (async () => { void (async () => {
for await (const line of lines) { for await (const line of lines) {
if (!readySettled && isRecord(line) && line.type === "ready") { if (!readySettled && isRecord(line) && line.type === "ready") {
protocolV2Supported = supportsRpcProtocolV2(line);
readySettled = true; readySettled = true;
readyResolve(); readyResolve();
continue; continue;
} }
this.#handleLine(line); if (isRecord(line) && line.type === "rpc_chunk" && !protocolV2Enabled)
throw new Error("RPC chunk received before protocol negotiation");
const decoded = frameDecoder.push(line);
if (decoded) this.#handleLine(decoded);
} }
// A closed stdout is terminal even if the child remains alive. Startup // A closed stdout is terminal even if the child remains alive. Startup
// failures are reaped by the readyPromise catch below; established // failures are reaped by the readyPromise catch below; established
@@ -359,6 +380,18 @@ export class RpcClient {
try { try {
await readyPromise; await readyPromise;
if (protocolV2Supported) {
protocolV2Enabled = true;
const response = await this.#send({ type: "negotiate_protocol", protocolVersion: 2 });
if (
!response.success ||
response.command !== "negotiate_protocol" ||
!isRecord(response.data) ||
response.data.protocolVersion !== 2
)
throw new Error("RPC protocol v2 negotiation failed");
this.#protocolVersion = 2;
}
if (this.#customTools.length > 0) { if (this.#customTools.length > 0) {
await this.setCustomTools(this.#customTools); await this.setCustomTools(this.#customTools);
} }
@@ -744,9 +777,42 @@ 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[]> { async getMessages(): Promise<AgentMessage[]> {
if (this.#protocolVersion === 2) {
try {
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;
} catch (error) {
if (!(error instanceof Error) || error.message !== RPC_MESSAGES_PAGE_BUSY_ERROR) throw error;
}
}
const response = await this.#send({ type: "get_messages" }); const response = await this.#send({ type: "get_messages" });
return this.#getData<{ messages: AgentMessage[] }>(response).messages; return this.#getData<{ messages: AgentMessage[] }>(response).messages;
} }
@@ -1,8 +1,24 @@
import { isDeepStrictEqual } from "node:util"; import { isDeepStrictEqual } from "node:util";
import { isRecord } from "@oh-my-pi/pi-utils"; import { isRecord } from "@oh-my-pi/pi-utils";
import type { RpcChunkFrame } from "./rpc-types";
/** Maximum UTF-8 size of one newline-delimited RPC frame, including the newline. */ /** Maximum UTF-8 size of one newline-delimited RPC frame, including the newline. */
export const MAX_RPC_FRAME_BYTES = 1024 * 1024; export const MAX_RPC_FRAME_BYTES = 1024 * 1024;
/** Maximum UTF-8 size of one logical frame reassembled by protocol v2. */
export const MAX_RPC_REASSEMBLED_BYTES = 64 * 1024 * 1024;
const RPC_CHUNK_PAYLOAD_BYTES = 256 * 1024;
export type RpcProtocolVersion = 1 | 2;
interface PendingRpcChunks {
chunkId: string;
count: number;
byteLength: number;
nextIndex: number;
chunks: Buffer[];
receivedBytes: number;
}
interface ShrinkPass { interface ShrinkPass {
stringCap: number; stringCap: number;
@@ -68,6 +84,103 @@ function encodedMessageSnapshot(encoded: string): { message: unknown } | undefin
: undefined; : undefined;
} }
function encodeChunkedRpcFrame(frame: object, chunkId: string): string {
const json = JSON.stringify(frame);
const bytes = Buffer.from(json, "utf8");
if (bytes.byteLength > MAX_RPC_REASSEMBLED_BYTES) return `${JSON.stringify(overflowFrame(frame))}\n`;
const count = Math.ceil(bytes.byteLength / RPC_CHUNK_PAYLOAD_BYTES);
let encoded = "";
for (let index = 0; index < count; index++) {
const chunk: RpcChunkFrame = {
type: "rpc_chunk",
chunkId,
index,
count,
byteLength: bytes.byteLength,
data: bytes
.subarray(index * RPC_CHUNK_PAYLOAD_BYTES, (index + 1) * RPC_CHUNK_PAYLOAD_BYTES)
.toString("base64"),
};
const line = `${JSON.stringify(chunk)}\n`;
if (serializedFrameBytes(line.slice(0, -1)) > MAX_RPC_FRAME_BYTES)
throw new Error("RPC chunk exceeded the transport limit");
encoded += line;
}
return encoded;
}
function isRpcChunkFrame(value: unknown): value is RpcChunkFrame {
return isRecord(value) && value.type === "rpc_chunk";
}
function decodeBase64(data: unknown): Buffer {
if (
typeof data !== "string" ||
data.length === 0 ||
!/^(?:[A-Za-z0-9+/]{4})*(?:[A-Za-z0-9+/]{2}==|[A-Za-z0-9+/]{3}=)?$/.test(data)
)
throw new Error("invalid rpc chunk data");
const bytes = Buffer.from(data, "base64");
if (bytes.toString("base64") !== data) throw new Error("invalid rpc chunk data");
return bytes;
}
/** Reassemble protocol v2 chunk frames after each JSONL line has been parsed. */
export class RpcFrameDecoder {
#pending?: PendingRpcChunks;
push(value: unknown): object | undefined {
if (!isRpcChunkFrame(value)) {
if (this.#pending) throw new Error("rpc chunk sequence interrupted");
if (!isRecord(value)) throw new Error("rpc frame must be an object");
return value;
}
const { chunkId, index, count, byteLength } = value;
if (
typeof chunkId !== "string" ||
chunkId.length === 0 ||
chunkId.length > 128 ||
!Number.isSafeInteger(index) ||
!Number.isSafeInteger(count) ||
!Number.isSafeInteger(byteLength) ||
index < 0 ||
count < 2 ||
count > Math.ceil(MAX_RPC_REASSEMBLED_BYTES / RPC_CHUNK_PAYLOAD_BYTES) ||
index >= count ||
byteLength < MAX_RPC_FRAME_BYTES ||
byteLength > MAX_RPC_REASSEMBLED_BYTES
)
throw new Error("invalid rpc chunk metadata");
const bytes = decodeBase64(value.data);
if (bytes.byteLength > RPC_CHUNK_PAYLOAD_BYTES) throw new Error("rpc chunk payload exceeds the transport limit");
if (!this.#pending) {
if (index !== 0) throw new Error("rpc chunk sequence must start at index 0");
this.#pending = { chunkId, count, byteLength, nextIndex: 0, chunks: [], receivedBytes: 0 };
}
const pending = this.#pending;
if (
pending.chunkId !== chunkId ||
pending.count !== count ||
pending.byteLength !== byteLength ||
pending.nextIndex !== index
)
throw new Error("rpc chunk sequence mismatch");
pending.chunks.push(bytes);
pending.receivedBytes += bytes.byteLength;
pending.nextIndex++;
if (pending.receivedBytes > pending.byteLength) throw new Error("rpc chunk sequence exceeds declared length");
if (pending.nextIndex < pending.count) return undefined;
if (pending.receivedBytes !== pending.byteLength) throw new Error("rpc chunk sequence length mismatch");
this.#pending = undefined;
const decoded = new TextDecoder("utf-8", { fatal: true }).decode(Buffer.concat(pending.chunks));
const frame: unknown = JSON.parse(decoded);
if (!isRecord(frame)) throw new Error("rpc frame must be an object");
return frame;
}
}
function compactTerminalFrame( function compactTerminalFrame(
frame: object, frame: object,
streamedMessageCount: number, streamedMessageCount: number,
@@ -142,13 +255,34 @@ export function encodeRpcFrame(frame: object, streamedMessageCount = 0, streamed
/** Stateful encoder that tracks which messages a client has already received. */ /** Stateful encoder that tracks which messages a client has already received. */
export class RpcFrameEncoder { export class RpcFrameEncoder {
#streamedMessages: unknown[] = []; #streamedMessages: unknown[] = [];
#protocolVersion: RpcProtocolVersion = 1;
#chunkCounter = 0;
setProtocolVersion(version: number): void {
if (version !== 1 && version !== 2) throw new Error(`Unsupported RPC protocol version: ${version}`);
this.#protocolVersion = version;
}
encode(frame: object): string { encode(frame: object): string {
if (isRecord(frame) && frame.type === "agent_start") this.#streamedMessages = []; if (isRecord(frame) && frame.type === "agent_start") this.#streamedMessages = [];
const encoded = encodeRpcFrame(frame, this.#streamedMessages.length, this.#streamedMessages); let json = JSON.stringify(frame);
let encoded: string;
if (this.#protocolVersion === 2 && serializedFrameBytes(json) > MAX_RPC_FRAME_BYTES) {
const compacted = compactTerminalFrame(frame, this.#streamedMessages.length, this.#streamedMessages);
json = JSON.stringify(compacted);
encoded =
serializedFrameBytes(json) > MAX_RPC_FRAME_BYTES
? encodeChunkedRpcFrame(compacted, `rpc-${++this.#chunkCounter}`)
: `${json}\n`;
} else {
encoded = encodeRpcFrame(frame, this.#streamedMessages.length, this.#streamedMessages);
}
if (!isRecord(frame)) return encoded; if (!isRecord(frame)) return encoded;
if (frame.type === "message_end") { if (frame.type === "message_end") {
const snapshot = encodedMessageSnapshot(encoded); const snapshot =
this.#protocolVersion === 2 && Object.hasOwn(frame, "message")
? { message: jsonSnapshot(frame.message) }
: encodedMessageSnapshot(encoded);
if (snapshot) this.#streamedMessages.push(snapshot.message); if (snapshot) this.#streamedMessages.push(snapshot.message);
} else if (frame.type === "agent_end" && frame.willContinue !== true) this.#streamedMessages = []; } else if (frame.type === "agent_end" && frame.willContinue !== true) this.#streamedMessages = [];
return encoded; return encoded;
@@ -0,0 +1,111 @@
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 const RPC_MESSAGES_PAGE_BUSY_ERROR = "Cannot page messages while the session is changing";
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,
};
}
@@ -34,8 +34,9 @@ import type { EventBus } from "../../utils/event-bus";
import { initializeExtensions } from "../runtime-init"; import { initializeExtensions } from "../runtime-init";
import { isRpcHostToolResult, isRpcHostToolUpdate, RpcHostToolBridge } from "./host-tools"; import { isRpcHostToolResult, isRpcHostToolUpdate, RpcHostToolBridge } from "./host-tools";
import { isRpcHostUriResult, RpcHostUriBridge } from "./host-uris"; import { isRpcHostUriResult, RpcHostUriBridge } from "./host-uris";
import { RpcFrameEncoder } from "./rpc-frame"; import { MAX_RPC_FRAME_BYTES, MAX_RPC_REASSEMBLED_BYTES, RpcFrameEncoder } from "./rpc-frame";
import { claimRpcInput } from "./rpc-input"; import { claimRpcInput } from "./rpc-input";
import { pageRpcMessages, RPC_MESSAGES_PAGE_BUSY_ERROR } from "./rpc-messages";
import { RpcSubagentRegistry, readRpcSubagentTranscript } from "./rpc-subagents"; import { RpcSubagentRegistry, readRpcSubagentTranscript } from "./rpc-subagents";
import type { import type {
RpcCommand, RpcCommand,
@@ -619,9 +620,19 @@ export async function runRpcMode(
process.env.PI_NOTIFICATIONS = "off"; process.env.PI_NOTIFICATIONS = "off";
const frameEncoder = new RpcFrameEncoder(); const frameEncoder = new RpcFrameEncoder();
process.stdout.write(frameEncoder.encode({ type: "ready" })); process.stdout.write(
frameEncoder.encode({
type: "ready",
protocolVersion: 1,
supportedProtocolVersions: [1, 2],
maxFrameBytes: MAX_RPC_FRAME_BYTES,
maxReassembledFrameBytes: MAX_RPC_REASSEMBLED_BYTES,
}),
);
const output = (obj: RpcResponse | RpcExtensionUIRequest | object) => { const output = (obj: RpcResponse | RpcExtensionUIRequest | object) => {
process.stdout.write(frameEncoder.encode(obj)); process.stdout.write(frameEncoder.encode(obj));
if (isRecord(obj) && obj.type === "response" && obj.command === "negotiate_protocol" && obj.success === true)
frameEncoder.setProtocolVersion(2);
}; };
const emitRpcTitles = shouldEmitRpcTitles(); const emitRpcTitles = shouldEmitRpcTitles();
@@ -936,6 +947,12 @@ export async function runRpcMode(
const id = command.id; const id = command.id;
switch (command.type) { switch (command.type) {
case "negotiate_protocol": {
if (command.protocolVersion !== 2)
return error(id, "negotiate_protocol", `Unsupported RPC protocol version: ${command.protocolVersion}`);
return success(id, "negotiate_protocol", { protocolVersion: 2 });
}
// ================================================================= // =================================================================
// Prompting // Prompting
// ================================================================= // =================================================================
@@ -1289,6 +1306,33 @@ export async function runRpcMode(
return success(id, "get_messages", { messages: session.messages }); return success(id, "get_messages", { messages: session.messages });
} }
case "get_messages_page": {
if (session.isStreaming || session.isCompacting)
return error(id, "get_messages_page", RPC_MESSAGES_PAGE_BUSY_ERROR);
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 // Login
// ================================================================= // =================================================================
@@ -19,12 +19,16 @@ import type {
SubagentProgressPayload, SubagentProgressPayload,
} from "../../task"; } from "../../task";
import type { TodoPhase } from "../../tools/todo"; import type { TodoPhase } from "../../tools/todo";
import type { RpcMessagesPage } from "./rpc-messages";
// ============================================================================ // ============================================================================
// RPC Commands (stdin) // RPC Commands (stdin)
// ============================================================================ // ============================================================================
export type RpcCommand = export type RpcCommand =
// Protocol
| { id?: string; type: "negotiate_protocol"; protocolVersion: number }
// Prompting // Prompting
| { id?: string; type: "prompt"; message: string; images?: ImageContent[]; streamingBehavior?: "steer" | "followUp" } | { id?: string; type: "prompt"; message: string; images?: ImageContent[]; streamingBehavior?: "steer" | "followUp" }
| { id?: string; type: "steer"; message: string; images?: ImageContent[] } | { id?: string; type: "steer"; message: string; images?: ImageContent[] }
@@ -81,6 +85,7 @@ export type RpcCommand =
// Messages // Messages
| { id?: string; type: "get_messages" } | { id?: string; type: "get_messages" }
| { id?: string; type: "get_messages_page"; cursor?: string; limit?: number }
// Login // Login
| { id?: string; type: "get_login_providers" } | { id?: string; type: "get_login_providers" }
@@ -132,6 +137,23 @@ export interface RpcPromptResultFrame {
agentInvoked: boolean; agentInvoked: boolean;
} }
export interface RpcReadyFrame {
type: "ready";
protocolVersion: 1;
supportedProtocolVersions: [1, 2];
maxFrameBytes: number;
maxReassembledFrameBytes: number;
}
export interface RpcChunkFrame {
type: "rpc_chunk";
chunkId: string;
index: number;
count: number;
byteLength: number;
data: string;
}
export interface RpcHandoffResult { export interface RpcHandoffResult {
savedPath?: string; savedPath?: string;
} }
@@ -168,6 +190,15 @@ export interface RpcSubagentMessagesResult {
// Success responses with data // Success responses with data
export type RpcResponse = export type RpcResponse =
// Protocol
| {
id?: string;
type: "response";
command: "negotiate_protocol";
success: true;
data: { protocolVersion: 2 };
}
// Prompting (async - events follow) // Prompting (async - events follow)
| { id?: string; type: "response"; command: "prompt"; success: true; data?: { agentInvoked: boolean } } | { id?: string; type: "response"; command: "prompt"; success: true; data?: { agentInvoked: boolean } }
| { id?: string; type: "response"; command: "steer"; success: true } | { id?: string; type: "response"; command: "steer"; success: true }
@@ -284,6 +315,7 @@ export type RpcResponse =
// Messages // Messages
| { id?: string; type: "response"; command: "get_messages"; success: true; data: { messages: AgentMessage[] } } | { 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 // Login
| { | {
+95 -6
View File
@@ -14,7 +14,43 @@ if (Bun.env.MOCK_RPC_IGNORE_SIGTERM === "1") {
process.on("SIGTERM", () => {}); process.on("SIGTERM", () => {});
} }
process.stdout.write(`${JSON.stringify({ type: "ready" })}\n`); const supportsProtocolV2 = Bun.env.MOCK_RPC_V2 === "1";
let protocolV2Enabled = false;
process.stdout.write(
`${JSON.stringify(
supportsProtocolV2
? {
type: "ready",
protocolVersion: 1,
supportedProtocolVersions: [1, 2],
maxFrameBytes: 1024 * 1024,
maxReassembledFrameBytes: 64 * 1024 * 1024,
}
: { type: "ready" },
)}\n`,
);
function writeFrame(frame: Record<string, unknown>): void {
const logical = Buffer.from(JSON.stringify(frame), "utf8");
if (!protocolV2Enabled || logical.byteLength <= 1024 * 1024) {
process.stdout.write(`${logical.toString("utf8")}\n`);
return;
}
const chunkBytes = 256 * 1024;
const count = Math.ceil(logical.byteLength / chunkBytes);
for (let index = 0; index < count; index++) {
process.stdout.write(
`${JSON.stringify({
type: "rpc_chunk",
chunkId: "mock-rpc-v2",
index,
count,
byteLength: logical.byteLength,
data: logical.subarray(index * chunkBytes, (index + 1) * chunkBytes).toString("base64"),
})}\n`,
);
}
}
// Bun's `console` is an AsyncIterable over stdin lines. // Bun's `console` is an AsyncIterable over stdin lines.
for await (const raw of console) { for await (const raw of console) {
@@ -32,15 +68,68 @@ for await (const raw of console) {
} }
if (Bun.env.MOCK_RPC_IGNORE_COMMANDS === "1") continue; if (Bun.env.MOCK_RPC_IGNORE_COMMANDS === "1") continue;
const id = typeof frame.id === "string" ? frame.id : undefined; const id = typeof frame.id === "string" ? frame.id : undefined;
process.stdout.write( if (frame.type === "negotiate_protocol" && frame.protocolVersion === 2) {
`${JSON.stringify({ writeFrame({
id, id,
type: "response", type: "response",
command: frame.type, command: frame.type,
success: true, success: true,
data: {}, data: { protocolVersion: 2 },
})}\n`, });
); protocolV2Enabled = true;
continue;
}
if (frame.type === "get_messages_page") {
if (Bun.env.MOCK_RPC_PAGE_BUSY === "1") {
writeFrame({
id,
type: "response",
command: frame.type,
success: false,
error: "Cannot page messages while the session is changing",
});
continue;
}
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;
}
if (frame.type === "get_messages" && Bun.env.MOCK_RPC_PAGE_BUSY === "1") {
writeFrame({
id,
type: "response",
command: frame.type,
success: true,
data: {
messages: [
{ role: "assistant", content: [{ type: "text", text: "streaming snapshot" }], timestamp: 3 },
],
},
});
continue;
}
writeFrame({
id,
type: "response",
command: frame.type,
success: true,
data: supportsProtocolV2 ? { payload: "😀".repeat(400_000) } : {},
});
} }
} catch { } catch {
// ignore parse errors — the test harness sends well-formed frames. // ignore parse errors — the test harness sends well-formed frames.
@@ -15,6 +15,34 @@ function isProcessAlive(pid: number): boolean {
} }
describe("RpcClient lifecycle (issue #4079 B)", () => { describe("RpcClient lifecycle (issue #4079 B)", () => {
test("auto-negotiates protocol v2 and reassembles an oversized response", async () => {
using client = new RpcClient({
cliPath: MOCK_AGENT,
env: { MOCK_RPC_V2: "1" },
});
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("preserves getMessages snapshot behavior while a v2 page walk is unavailable", async () => {
using client = new RpcClient({
cliPath: MOCK_AGENT,
env: { MOCK_RPC_V2: "1", MOCK_RPC_PAGE_BUSY: "1" },
});
await client.start();
await expect(client.getMessagesPage()).rejects.toThrow("Cannot page messages while the session is changing");
expect((await client.getMessages()) as unknown).toEqual([
{ role: "assistant", content: [{ type: "text", text: "streaming snapshot" }], timestamp: 3 },
]);
}, 20_000);
test("start() succeeds a second time after stop() on the same instance", async () => { test("start() succeeds a second time after stop() on the same instance", async () => {
using client = new RpcClient({ using client = new RpcClient({
cliPath: MOCK_AGENT, cliPath: MOCK_AGENT,
+109 -1
View File
@@ -1,5 +1,11 @@
import { describe, expect, it } from "bun:test"; import { describe, expect, it } from "bun:test";
import { encodeRpcFrame, MAX_RPC_FRAME_BYTES, RpcFrameEncoder } from "../src/modes/rpc/rpc-frame"; import {
encodeRpcFrame,
MAX_RPC_FRAME_BYTES,
MAX_RPC_REASSEMBLED_BYTES,
RpcFrameDecoder,
RpcFrameEncoder,
} from "../src/modes/rpc/rpc-frame";
function decode(frame: string): Record<string, unknown> { function decode(frame: string): Record<string, unknown> {
return JSON.parse(frame) as Record<string, unknown>; return JSON.parse(frame) as Record<string, unknown>;
@@ -178,4 +184,106 @@ describe("RPC frame encoding", () => {
expect(decoded.success).toBe(false); expect(decoded.success).toBe(false);
expect(decoded.id).toContain("chars elided for RPC frame"); expect(decoded.id).toContain("chars elided for RPC frame");
}); });
it("losslessly chunks oversized protocol v2 responses into bounded JSONL frames", () => {
const frame = {
id: "request-v2",
type: "response",
command: "get_messages",
success: true,
data: { messages: [{ role: "assistant", content: "😀".repeat(400_000) }] },
};
const encoder = new RpcFrameEncoder();
encoder.setProtocolVersion(2);
const encoded = encoder.encode(frame);
const lines = encoded.trimEnd().split("\n");
const decoder = new RpcFrameDecoder();
let decoded: object | undefined;
expect(lines.length).toBeGreaterThan(1);
for (const line of lines) {
expect(Buffer.byteLength(`${line}\n`, "utf8")).toBeLessThanOrEqual(MAX_RPC_FRAME_BYTES);
decoded = decoder.push(JSON.parse(line));
}
expect(decoded).toEqual(frame);
});
it("accepts a chunked logical frame at the exact physical-frame boundary", () => {
const frame = {
id: "request-boundary",
type: "response",
command: "get_state",
success: true,
data: { payload: "" },
};
const emptyBytes = Buffer.byteLength(JSON.stringify(frame), "utf8");
frame.data.payload = "x".repeat(MAX_RPC_FRAME_BYTES - emptyBytes);
expect(Buffer.byteLength(JSON.stringify(frame), "utf8")).toBe(MAX_RPC_FRAME_BYTES);
const encoder = new RpcFrameEncoder();
encoder.setProtocolVersion(2);
const decoder = new RpcFrameDecoder();
let decoded: object | undefined;
for (const line of encoder.encode(frame).trimEnd().split("\n")) decoded = decoder.push(JSON.parse(line));
expect(decoded).toEqual(frame);
});
it("preserves terminal message counts above the protocol v2 ceiling", () => {
const encoder = new RpcFrameEncoder();
encoder.setProtocolVersion(2);
const encoded = encoder.encode({
type: "agent_end",
messages: [{ role: "assistant", content: "x".repeat(MAX_RPC_REASSEMBLED_BYTES) }],
});
expect(decode(encoded)).toEqual({
type: "agent_end",
messages: [],
messageCount: 1,
});
});
it("rejects protocol v2 logical frames above the advertised reassembly ceiling", () => {
const encoder = new RpcFrameEncoder();
encoder.setProtocolVersion(2);
const encoded = encoder.encode({
id: "request-too-large",
type: "response",
command: "get_messages",
success: true,
data: { transcript: "x".repeat(MAX_RPC_REASSEMBLED_BYTES) },
});
expect(decode(encoded)).toEqual({
id: "request-too-large",
type: "response",
command: "get_messages",
success: false,
error: "RPC response exceeded the transport limit",
});
});
it("rejects interrupted protocol v2 chunk sequences", () => {
const decoder = new RpcFrameDecoder();
decoder.push({
type: "rpc_chunk",
chunkId: "chunk-1",
index: 0,
count: 2,
byteLength: MAX_RPC_FRAME_BYTES + 1,
data: "ew==",
});
expect(() =>
decoder.push({
type: "rpc_chunk",
chunkId: "chunk-2",
index: 1,
count: 2,
byteLength: MAX_RPC_FRAME_BYTES + 1,
data: "fQ==",
}),
).toThrow("rpc chunk sequence mismatch");
});
}); });
@@ -26,10 +26,12 @@ describe("RPC mode malformed stdin", () => {
// crashed the generator before the second was ever read. // crashed the generator before the second was ever read.
child.stdin.write("this is not json\n"); 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_state", id: "probe" })}\n`);
child.stdin.write(`${JSON.stringify({ type: "get_messages_page", id: "page-probe", limit: 1 })}\n`);
await child.stdin.flush(); await child.stdin.flush();
let parseError: Record<string, unknown> | undefined; let parseError: Record<string, unknown> | undefined;
let stateResponse: 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>)) { for await (const frame of readJsonl<unknown>(child.stdout as ReadableStream<Uint8Array>)) {
if (!isRecord(frame)) continue; if (!isRecord(frame)) continue;
@@ -38,8 +40,9 @@ describe("RPC mode malformed stdin", () => {
} }
if (frame.type === "response" && frame.id === "probe") { if (frame.type === "response" && frame.id === "probe") {
stateResponse = frame; stateResponse = frame;
break;
} }
if (frame.type === "response" && frame.id === "page-probe") pageResponse = frame;
if (stateResponse && pageResponse) break;
} }
child.stdin.end(); child.stdin.end();
@@ -50,5 +53,9 @@ describe("RPC mode malformed stdin", () => {
expect(String(parseError?.error)).toContain("Failed to parse command"); expect(String(parseError?.error)).toContain("Failed to parse command");
expect(stateResponse).toBeDefined(); expect(stateResponse).toBeDefined();
expect(stateResponse?.success).toBe(true); expect(stateResponse?.success).toBe(true);
expect(pageResponse).toMatchObject({
success: true,
data: { messages: [], totalMessages: 0 },
});
}, 30000); }, 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, - typed startup options for common `omp --mode rpc` flags such as thinking level,
tool selection, prompt appends, provider session IDs, and headless session toggles tool selection, prompt appends, provider session IDs, and headless session toggles
- typed protocol models for state, bash results, compaction, and session stats - 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 - a process-backed client that manages request correlation over stdio
- typed per-event listeners plus a typed catch-all notification hook - 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 - 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, HookMessage,
ImageContent, ImageContent,
MessageEndEvent, MessageEndEvent,
MessagesPage,
MessageStartEvent, MessageStartEvent,
MessageUpdateEvent, MessageUpdateEvent,
ModelCycleResult, ModelCycleResult,
@@ -136,6 +137,7 @@ __all__ = [
"ListenerErrorEvent", "ListenerErrorEvent",
"ListenerErrorListener", "ListenerErrorListener",
"MessageEndEvent", "MessageEndEvent",
"MessagesPage",
"MessageStartEvent", "MessageStartEvent",
"MessageUpdateEvent", "MessageUpdateEvent",
"ModelCost", "ModelCost",
+194 -4
View File
@@ -1,5 +1,7 @@
from __future__ import annotations from __future__ import annotations
import base64
import binascii
import json import json
import os import os
import queue import queue
@@ -34,6 +36,7 @@ from .protocol import (
JsonObject, JsonObject,
JsonValue, JsonValue,
MessageEndEvent, MessageEndEvent,
MessagesPage,
MessageStartEvent, MessageStartEvent,
MessageUpdateEvent, MessageUpdateEvent,
ModelCycleResult, ModelCycleResult,
@@ -111,6 +114,102 @@ THistoryItem = TypeVar("THistoryItem")
_ASYNC_COMMANDS = frozenset({"prompt", "abort_and_prompt"}) _ASYNC_COMMANDS = frozenset({"prompt", "abort_and_prompt"})
_DEFAULT_ERROR_HISTORY_LIMIT = 128 _DEFAULT_ERROR_HISTORY_LIMIT = 128
_TODO_STATUS_VALUES = frozenset({"pending", "in_progress", "completed", "abandoned"}) _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
_RPC_MESSAGES_PAGE_BUSY_ERROR = "Cannot page messages while the session is changing"
@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: def _process_group_id(process: subprocess.Popen[Any]) -> int | None:
@@ -417,6 +516,10 @@ class RpcClient:
self._closed_error: BaseException | None = None self._closed_error: BaseException | None = None
self._stopping = False self._stopping = False
self._ready_received = 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]( self._protocol_errors = _BoundedHistory[RpcProtocolError](
_DEFAULT_ERROR_HISTORY_LIMIT _DEFAULT_ERROR_HISTORY_LIMIT
) )
@@ -468,6 +571,10 @@ class RpcClient:
self._stopping = False self._stopping = False
self._closed_error = None self._closed_error = None
self._ready_received = False 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._events.clear()
self._async_errors.clear() self._async_errors.clear()
self._scheduled_agent_runs = 0 self._scheduled_agent_runs = 0
@@ -529,6 +636,24 @@ class RpcClient:
f"Timed out waiting for RPC ready signal. Stderr: {stderr}" 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: if self._custom_tools:
self.set_custom_tools(self._custom_tools) self.set_custom_tools(self._custom_tools)
if self._host_uris: if self._host_uris:
@@ -882,9 +1007,65 @@ class RpcClient:
return self.set_todos(()) return self.set_todos(())
def get_messages(self) -> tuple[AgentMessage, ...]: def get_messages(self) -> tuple[AgentMessage, ...]:
if self._protocol_version == 2:
try:
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)
except RpcCommandError as error:
if (
error.command != "get_messages_page"
or error.error != _RPC_MESSAGES_PAGE_BUSY_ERROR
):
raise
payload = self._request("get_messages") payload = self._request("get_messages")
return parse_agent_messages(cast(JsonValue | None, payload.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, ...]: def set_custom_tools(self, tools: Sequence[HostTool[Any, Any]]) -> tuple[str, ...]:
self._custom_tools = tuple(tools) self._custom_tools = tuple(tools)
if self._process is None: if self._process is None:
@@ -1100,9 +1281,8 @@ class RpcClient:
def _complete_agent_end_messages( def _complete_agent_end_messages(
events: tuple[RpcAgentEvent, ...], terminal: AgentEndEvent events: tuple[RpcAgentEvent, ...], terminal: AgentEndEvent
) -> tuple[AgentMessage, ...]: ) -> tuple[AgentMessage, ...]:
if ( if terminal.message_count is None or terminal.message_count <= len(
terminal.message_count is None terminal.messages
or terminal.message_count <= len(terminal.messages)
): ):
return terminal.messages return terminal.messages
@@ -1632,7 +1812,7 @@ class RpcClient:
continue continue
try: try:
payload = cast(JsonObject, json.loads(stripped)) raw_payload = json.loads(stripped)
except json.JSONDecodeError as exc: except json.JSONDecodeError as exc:
snippet = stripped snippet = stripped
if len(snippet) > 240: if len(snippet) > 240:
@@ -1640,6 +1820,15 @@ class RpcClient:
raise RpcError( raise RpcError(
f"Failed to decode RPC output on line {line_number}: {exc}. Frame: {snippet!r}" f"Failed to decode RPC output on line {line_number}: {exc}. Frame: {snippet!r}"
) from exc ) 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": if payload.get("type") == "response":
self._handle_response(payload) self._handle_response(payload)
continue continue
@@ -1666,6 +1855,7 @@ class RpcClient:
) )
if isinstance(notification, ReadyEvent): if isinstance(notification, ReadyEvent):
self._ready_event = notification
self._ready_received = True self._ready_received = True
self._ready.set() self._ready.set()
self._dispatch_listeners( self._dispatch_listeners(
+28 -1
View File
@@ -851,9 +851,20 @@ class SessionStats:
@dataclass(slots=True, frozen=True) @dataclass(slots=True, frozen=True)
class ReadyEvent: 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" 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) @dataclass(slots=True, frozen=True)
class ExtensionUiRequest: class ExtensionUiRequest:
id: str id: str
@@ -1490,7 +1501,23 @@ def parse_extension_error(payload: JsonObject) -> ExtensionError:
def parse_notification(payload: JsonObject) -> RpcNotification: def parse_notification(payload: JsonObject) -> RpcNotification:
event_type = payload.get("type") event_type = payload.get("type")
if event_type == "ready": 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": if event_type == "extension_ui_request":
return parse_extension_ui_request(payload) return parse_extension_ui_request(payload)
if event_type == "extension_error": if event_type == "extension_error":
+178
View File
@@ -1,5 +1,7 @@
from __future__ import annotations from __future__ import annotations
import base64
import json
import os import os
import shutil import shutil
import signal import signal
@@ -11,6 +13,7 @@ import time
import unittest import unittest
from omp_rpc import RpcClient, RpcCommandError, RpcConcurrencyError, RpcError, host_tool from omp_rpc import RpcClient, RpcCommandError, RpcConcurrencyError, RpcError, host_tool
from omp_rpc.client import _RpcFrameDecoder
FAKE_SERVER = textwrap.dedent( FAKE_SERVER = textwrap.dedent(
@@ -442,6 +445,129 @@ FAKE_SERVER = textwrap.dedent(
""" """
) )
V2_MESSAGES_SERVER = textwrap.dedent(
"""
import base64
import json
import os
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":
if os.environ.get("V2_MESSAGES_BUSY") == "1":
emit(
{
"id": request_id,
"type": "response",
"command": command_type,
"success": False,
"error": "Cannot page messages while the session is changing",
}
)
continue
emit(
{
"id": request_id,
"type": "response",
"command": command_type,
"success": True,
"data": {
"messages": [message],
"totalMessages": 1,
"nextCursor": None,
},
}
)
elif command_type == "get_messages":
emit(
{
"id": request_id,
"type": "response",
"command": command_type,
"success": True,
"data": {
"messages": [
{
"role": "assistant",
"content": [
{"type": "text", "text": "streaming snapshot"}
],
"timestamp": 3,
}
]
},
}
)
else:
emit(
{
"id": request_id,
"type": "response",
"command": command_type,
"success": False,
"error": f"unexpected command: {command_type}",
}
)
"""
)
IDLESS_ERROR_SERVER = textwrap.dedent( IDLESS_ERROR_SERVER = textwrap.dedent(
""" """
import json import json
@@ -612,6 +738,38 @@ class RpcClientTests(unittest.TestCase):
**kwargs, **kwargs,
) )
def test_protocol_v2_decoder_accepts_exact_logical_boundary(self) -> None:
frame = {
"id": "request-boundary",
"type": "response",
"command": "get_state",
"success": True,
"data": {"payload": ""},
}
encoded_empty = json.dumps(frame, separators=(",", ":")).encode("utf-8")
frame["data"]["payload"] = "x" * (1024 * 1024 - len(encoded_empty))
encoded = json.dumps(frame, separators=(",", ":")).encode("utf-8")
self.assertEqual(len(encoded), 1024 * 1024)
decoder = _RpcFrameDecoder()
chunk_size = 256 * 1024
count = (len(encoded) + chunk_size - 1) // chunk_size
decoded = None
for index in range(count):
chunk = encoded[index * chunk_size : (index + 1) * chunk_size]
decoded = decoder.push(
{
"type": "rpc_chunk",
"chunkId": "exact-boundary",
"index": index,
"count": count,
"byteLength": len(encoded),
"data": base64.b64encode(chunk).decode("ascii"),
}
)
self.assertEqual(decoded, frame)
def test_command_builder_supports_common_rpc_options(self) -> None: def test_command_builder_supports_common_rpc_options(self) -> None:
client = RpcClient( client = RpcClient(
executable="omp", executable="omp",
@@ -858,6 +1016,26 @@ class RpcClientTests(unittest.TestCase):
client.wait_for_idle(timeout=2.0) client.wait_for_idle(timeout=2.0)
self.assertEqual(client.get_last_assistant_text(), "pong") 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_protocol_v2_get_messages_falls_back_to_streaming_snapshot(self) -> None:
with self.make_client(
server=V2_MESSAGES_SERVER, env={"V2_MESSAGES_BUSY": "1"}
) as client:
with self.assertRaisesRegex(
RpcCommandError, "Cannot page messages while the session is changing"
):
client.get_messages_page()
messages = client.get_messages()
self.assertEqual(len(messages), 1)
self.assertEqual(messages[0]["content"][0]["text"], "streaming snapshot")
def test_collect_events_returns_turn_events(self) -> None: def test_collect_events_returns_turn_events(self) -> None:
with self.make_client() as client: with self.make_client() as client:
client.prompt("slow") client.prompt("slow")