diff --git a/packages/coding-agent/src/modes/rpc/rpc-frame.ts b/packages/coding-agent/src/modes/rpc/rpc-frame.ts index a86b1edfc..f4a86a507 100644 --- a/packages/coding-agent/src/modes/rpc/rpc-frame.ts +++ b/packages/coding-agent/src/modes/rpc/rpc-frame.ts @@ -1,3 +1,4 @@ +import { isDeepStrictEqual } from "node:util"; import { isRecord } from "@oh-my-pi/pi-utils"; /** Maximum UTF-8 size of one newline-delimited RPC frame, including the newline. */ @@ -55,11 +56,37 @@ function shrinkValue(value: unknown, pass: ShrinkPass): unknown { return value; } -function compactTerminalFrame(frame: object, streamedMessageCount: number): object { +function jsonSnapshot(value: unknown): unknown { + const json = JSON.stringify(value); + return json === undefined ? undefined : JSON.parse(json); +} + +function encodedMessageSnapshot(encoded: string): { message: unknown } | undefined { + const frame = JSON.parse(encoded); + return isRecord(frame) && frame.type === "message_end" && Object.hasOwn(frame, "message") + ? { message: frame.message } + : undefined; +} + +function compactTerminalFrame( + frame: object, + streamedMessageCount: number, + streamedMessages?: readonly unknown[], +): object { if (!isRecord(frame) || frame.type !== "agent_end" || !Array.isArray(frame.messages)) return frame; - const streamed = Number.isSafeInteger(streamedMessageCount) + let streamed = Number.isSafeInteger(streamedMessageCount) ? Math.min(Math.max(0, streamedMessageCount), frame.messages.length) : 0; + if (streamedMessages) { + streamed = 0; + const limit = Math.min(streamedMessages.length, frame.messages.length); + while ( + streamed < limit && + isDeepStrictEqual(streamedMessages[streamed], jsonSnapshot(frame.messages[streamed])) + ) { + streamed++; + } + } return { ...frame, messages: frame.messages.slice(streamed), @@ -93,8 +120,8 @@ function overflowFrame(frame: object): object { } /** Serialize a complete JSONL frame while enforcing the transport byte ceiling. */ -export function encodeRpcFrame(frame: object, streamedMessageCount = 0): string { - const compacted = compactTerminalFrame(frame, streamedMessageCount); +export function encodeRpcFrame(frame: object, streamedMessageCount = 0, streamedMessages?: readonly unknown[]): string { + const compacted = compactTerminalFrame(frame, streamedMessageCount, streamedMessages); let json = JSON.stringify(compacted); if (serializedFrameBytes(json) <= MAX_RPC_FRAME_BYTES) return `${json}\n`; if (isRecord(compacted) && compacted.type === "response") { @@ -111,14 +138,16 @@ export function encodeRpcFrame(frame: object, streamedMessageCount = 0): string /** Stateful encoder that tracks which messages a client has already received. */ export class RpcFrameEncoder { - #streamedMessageCount = 0; + #streamedMessages: unknown[] = []; encode(frame: object): string { - if (isRecord(frame) && frame.type === "agent_start") this.#streamedMessageCount = 0; - const encoded = encodeRpcFrame(frame, this.#streamedMessageCount); + if (isRecord(frame) && frame.type === "agent_start") this.#streamedMessages = []; + const encoded = encodeRpcFrame(frame, this.#streamedMessages.length, this.#streamedMessages); if (!isRecord(frame)) return encoded; - if (frame.type === "message_end") this.#streamedMessageCount++; - else if (frame.type === "agent_end") this.#streamedMessageCount = 0; + if (frame.type === "message_end") { + const snapshot = encodedMessageSnapshot(encoded); + if (snapshot) this.#streamedMessages.push(snapshot.message); + } else if (frame.type === "agent_end") this.#streamedMessages = []; return encoded; } } diff --git a/packages/coding-agent/test/rpc-frame.test.ts b/packages/coding-agent/test/rpc-frame.test.ts index ecd791d3c..e08117416 100644 --- a/packages/coding-agent/test/rpc-frame.test.ts +++ b/packages/coding-agent/test/rpc-frame.test.ts @@ -23,7 +23,7 @@ describe("RPC frame encoding", () => { expect(decoded).toEqual({ type: "agent_end", messages: [], messageCount: 10_000, telemetry: { stepCount: 42 } }); }); - it("retains terminal messages that were not emitted as message events", () => { + it("retains a terminal error emitted only by agent_end after earlier message events", () => { const streamed = { role: "assistant", content: [{ type: "text", text: "done" }] }; const aborted = { role: "assistant", @@ -34,12 +34,72 @@ describe("RPC frame encoding", () => { const encoder = new RpcFrameEncoder(); encoder.encode({ type: "agent_start" }); encoder.encode({ type: "message_end", message: streamed }); - const decoded = decode(encoder.encode({ type: "agent_end", messages: [streamed, aborted] })); + const decoded = decode(encoder.encode({ type: "agent_end", messages: [aborted] })); expect(decoded).toEqual({ type: "agent_end", messages: [aborted], - messageCount: 2, + messageCount: 1, + }); + }); + + it("compacts stateful agent_end messages that match earlier message events", () => { + const streamed = { role: "assistant", content: [{ type: "text", text: "done" }] }; + const encoder = new RpcFrameEncoder(); + encoder.encode({ type: "agent_start" }); + encoder.encode({ type: "message_end", message: streamed }); + const decoded = decode( + encoder.encode({ + type: "agent_end", + messages: [{ role: "assistant", content: [{ type: "text", text: "done" }] }], + }), + ); + + expect(decoded).toEqual({ + type: "agent_end", + messages: [], + messageCount: 1, + }); + }); + + it("matches terminal messages in the JSON shape sent by message_end", () => { + const encoder = new RpcFrameEncoder(); + encoder.encode({ type: "agent_start" }); + encoder.encode({ + type: "message_end", + message: { + role: "assistant", + content: [{ type: "text", text: "done" }], + disabledFeatures: undefined, + toolCallAbortMessages: undefined, + }, + }); + const decoded = decode( + encoder.encode({ + type: "agent_end", + messages: [{ role: "assistant", content: [{ type: "text", text: "done" }] }], + }), + ); + + expect(decoded).toEqual({ + type: "agent_end", + messages: [], + messageCount: 1, + }); + }); + + it("does not let later mutation rewrite the message_end snapshot", () => { + const streamed = { role: "assistant", content: [{ type: "text", text: "before" }] }; + const encoder = new RpcFrameEncoder(); + encoder.encode({ type: "agent_start" }); + encoder.encode({ type: "message_end", message: streamed }); + streamed.content[0].text = "after"; + const decoded = decode(encoder.encode({ type: "agent_end", messages: [streamed] })); + + expect(decoded).toEqual({ + type: "agent_end", + messages: [streamed], + messageCount: 1, }); });