Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
78 changes: 47 additions & 31 deletions packages/agent/src/proxy.ts
Original file line number Diff line number Diff line change
Expand Up @@ -7,29 +7,17 @@
import {
type AssistantMessage,
type AssistantMessageEvent,
type AssistantMessageEventStream,
type Context,
EventStream,
createAssistantMessageEventStream,
createPendingToolCall,
type Model,
parseStreamingJson,
type PendingToolCall,
type SimpleStreamOptions,
type StopReason,
type ToolCall,
} from "@earendil-works/pi-ai";

// Create stream class matching ProxyMessageEventStream
class ProxyMessageEventStream extends EventStream<AssistantMessageEvent, AssistantMessage> {
constructor() {
super(
(event) => event.type === "done" || event.type === "error",
(event) => {
if (event.type === "done") return event.message;
if (event.type === "error") return event.error;
throw new Error("Unexpected event type");
},
);
}
}

/**
* Proxy event types - server sends these with partial field stripped to reduce bandwidth.
*/
Expand Down Expand Up @@ -117,8 +105,17 @@ function buildProxyRequestOptions(options: ProxyStreamOptions): ProxySerializabl
};
}

export function streamProxy(model: Model<any>, context: Context, options: ProxyStreamOptions): ProxyMessageEventStream {
const stream = new ProxyMessageEventStream();
export function streamProxy(
model: Model<any>,
context: Context,
options: ProxyStreamOptions,
): AssistantMessageEventStream {
const stream = createAssistantMessageEventStream();
const pendingCalls = new Map<number, PendingToolCall>();
const finishPendingCalls = () => {
for (const pending of pendingCalls.values()) pending.finish();
pendingCalls.clear();
};

(async () => {
// Initialize the partial message that we'll build up from events
Expand Down Expand Up @@ -190,9 +187,12 @@ export function streamProxy(model: Model<any>, context: Context, options: ProxyS
const data = line.slice(6).trim();
if (!data) return;
const proxyEvent = JSON.parse(data) as ProxyAssistantMessageEvent;
const event = processProxyEvent(proxyEvent, partial);
const event = processProxyEvent(proxyEvent, partial, pendingCalls);
if (event) {
if (event.type === "done" || event.type === "error") sawTerminalEvent = true;
if (event.type === "done" || event.type === "error") {
finishPendingCalls();
sawTerminalEvent = true;
}
stream.push(event);
}
};
Expand Down Expand Up @@ -231,6 +231,7 @@ export function streamProxy(model: Model<any>, context: Context, options: ProxyS
// consumers waiting on a result that never arrives.
partial.stopReason = "error";
partial.errorMessage = "Connection closed by proxy server before the response completed";
finishPendingCalls();
stream.push({
type: "error",
reason: "error",
Expand All @@ -244,6 +245,7 @@ export function streamProxy(model: Model<any>, context: Context, options: ProxyS
const reason = options.signal?.aborted ? "aborted" : "error";
partial.stopReason = reason;
partial.errorMessage = errorMessage;
finishPendingCalls();
stream.push({
type: "error",
reason,
Expand All @@ -266,6 +268,7 @@ export function streamProxy(model: Model<any>, context: Context, options: ProxyS
function processProxyEvent(
proxyEvent: ProxyAssistantMessageEvent,
partial: AssistantMessage,
pendingCalls: Map<number, PendingToolCall>,
): AssistantMessageEvent | undefined {
switch (proxyEvent.type) {
case "start":
Expand Down Expand Up @@ -335,22 +338,33 @@ function processProxyEvent(
throw new Error("Received thinking_end for non-thinking content");
}

case "toolcall_start":
partial.content[proxyEvent.contentIndex] = {
type: "toolCall",
id: proxyEvent.id,
name: proxyEvent.toolName,
arguments: {},
partialJson: "",
} satisfies ToolCall & { partialJson: string } as ToolCall;
case "toolcall_start": {
const pending = createPendingToolCall(
{
type: "toolCall",
id: proxyEvent.id,
name: proxyEvent.toolName,
arguments: {},
partialJson: "",
},
true,
);
partial.content[proxyEvent.contentIndex] = pending.toolCall;
pendingCalls.set(proxyEvent.contentIndex, pending);
return { type: "toolcall_start", contentIndex: proxyEvent.contentIndex, partial };
}

case "toolcall_delta": {
const content = partial.content[proxyEvent.contentIndex];
if (content?.type === "toolCall") {
(content as any).partialJson += proxyEvent.delta;
content.arguments = parseStreamingJson((content as any).partialJson) || {};
partial.content[proxyEvent.contentIndex] = { ...content }; // Trigger reactivity
const block = content as ToolCall & { partialJson: string };
block.partialJson += proxyEvent.delta;
const pending = pendingCalls.get(proxyEvent.contentIndex)!;
pending.setJson(block.partialJson);
// Trigger reactivity without reading arguments or losing extra metadata.
const copy = pending.copy();
pendingCalls.set(proxyEvent.contentIndex, copy);
partial.content[proxyEvent.contentIndex] = copy.toolCall;
return {
type: "toolcall_delta",
contentIndex: proxyEvent.contentIndex,
Expand All @@ -365,6 +379,8 @@ function processProxyEvent(
const content = partial.content[proxyEvent.contentIndex];
if (content?.type === "toolCall") {
Object.assign(content, proxyEvent.toolCall);
pendingCalls.get(proxyEvent.contentIndex)?.finish();
pendingCalls.delete(proxyEvent.contentIndex);
delete (content as any).partialJson;
return {
type: "toolcall_end",
Expand Down
143 changes: 142 additions & 1 deletion packages/agent/test/proxy.test.ts
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
import type { AssistantMessage, AssistantMessageEvent, Model } from "@earendil-works/pi-ai";
import type { AssistantMessage, AssistantMessageEvent, Model, ToolCall } from "@earendil-works/pi-ai";
import { afterEach, describe, expect, it, vi } from "vitest";
import { type ProxyAssistantMessageEvent, streamProxy } from "../src/proxy.ts";

Expand Down Expand Up @@ -26,9 +26,150 @@ const usage: AssistantMessage["usage"] = {

afterEach(() => {
vi.unstubAllGlobals();
vi.restoreAllMocks();
});

describe("streamProxy", () => {
// #9265: copying blocks for reactivity must not evaluate lazy arguments.
it.each([true, false])("defers argument parsing and settles unread arguments (complete: %s)", async (complete) => {
const content = "a".repeat(32 * 512);
const deltas = ['{"content":"', ...Array<string>(32).fill("a".repeat(512))];
if (complete) deltas.push('"}');
const proxyEvents: ProxyAssistantMessageEvent[] = [
{ type: "start" },
{ type: "toolcall_start", contentIndex: 0, id: "call_test", toolName: "write" },
...deltas.map((delta): ProxyAssistantMessageEvent => ({ type: "toolcall_delta", contentIndex: 0, delta })),
...(complete
? ([
{
type: "toolcall_end",
contentIndex: 0,
toolCall: { type: "toolCall", id: "call_test", name: "write", arguments: { content } },
},
{ type: "done", reason: "toolUse", usage },
] satisfies ProxyAssistantMessageEvent[])
: []),
];
const body = proxyEvents.map((event) => `data: ${JSON.stringify(event)}\n\n`).join("");
vi.stubGlobal(
"fetch",
vi.fn(async () => new Response(body)),
);
const parse = vi.spyOn(JSON, "parse");
const result = await streamProxy(
model,
{ messages: [] },
{
authToken: "test",
proxyUrl: "https://proxy.example.com",
},
).result();
expect(result.stopReason, result.errorMessage).toBe(complete ? "toolUse" : "error");
if (complete) {
expect(parse.mock.calls.filter(([text]) => text.startsWith('{"content":'))).toEqual([]);
}
expect(Object.getOwnPropertyDescriptor(result.content[0], "arguments")).toMatchObject({
value: { content },
writable: true,
});
});

// #9265: lazy parsing must preserve the proxy's existing `parsed || {}` behavior.
it.each(["null", "false", "0", '""'])("preserves the empty-object fallback for %s arguments", async (json) => {
const proxyEvents: ProxyAssistantMessageEvent[] = [
{ type: "start" },
{ type: "toolcall_start", contentIndex: 0, id: "call_test", toolName: "write" },
{ type: "toolcall_delta", contentIndex: 0, delta: json },
{ type: "error", reason: "error", usage, errorMessage: "Interrupted" },
];
const body = proxyEvents.map((event) => `data: ${JSON.stringify(event)}\n\n`).join("");
vi.stubGlobal(
"fetch",
vi.fn(async () => new Response(body)),
);
const result = await streamProxy(
model,
{ messages: [] },
{
authToken: "test",
proxyUrl: "https://proxy.example.com",
},
).result();
expect(result.stopReason).toBe("error");
const block = result.content[0];
if (block.type !== "toolCall") throw new Error("Expected tool call");
expect(block.arguments).toEqual({});
});

// #9265: preserve main's shallow-copy identity and extension metadata during live updates.
it("copies live blocks without losing metadata or changing argument identity", async () => {
let transport!: ReadableStreamDefaultController<Uint8Array>;
const body = new ReadableStream<Uint8Array>({
start: (controller) => {
transport = controller;
},
});
const send = (...events: ProxyAssistantMessageEvent[]) => {
transport.enqueue(
new TextEncoder().encode(events.map((event) => `data: ${JSON.stringify(event)}\n\n`).join("")),
);
};
send(
{ type: "start" },
{ type: "toolcall_start", contentIndex: 0, id: "call_test", toolName: "write" },
{ type: "toolcall_delta", contentIndex: 0, delta: '{"content":"hel' },
);
vi.stubGlobal(
"fetch",
vi.fn(async () => new Response(body)),
);
const stream = streamProxy(model, { messages: [] }, { authToken: "test", proxyUrl: "https://proxy.example.com" });
const symbol = Symbol("metadata");
const metadata = { source: "extension" };
let first: ToolCall | undefined;
let firstArguments: ToolCall["arguments"] | undefined;
let last: ToolCall | undefined;
for await (const event of stream) {
if (event.type !== "toolcall_delta") continue;
const block = event.partial.content[0];
if (block.type !== "toolCall") throw new Error("Expected tool call");
if (!first) {
first = block;
firstArguments = block.arguments;
Object.assign(block, { extra: metadata, [symbol]: metadata });
send({ type: "toolcall_delta", contentIndex: 0, delta: 'lo"}' });
} else {
last = block;
expect(block).not.toBe(first);
expect(block.arguments).toEqual({ content: "hello" });
expect(block.arguments).toBe(first.arguments);
expect(firstArguments).toEqual({ content: "hel" });
expect(Reflect.get(block, "extra")).toBe(metadata);
expect(Reflect.get(block, symbol)).toBe(metadata);
send(
{
type: "toolcall_end",
contentIndex: 0,
toolCall: {
type: "toolCall",
id: "call_test",
name: "write",
arguments: { content: "authoritative" },
},
},
{ type: "done", reason: "toolUse", usage },
);
transport.close();
}
}
const result = await stream.result();
expect(result.content[0]).toBe(last);
expect(Object.getOwnPropertyDescriptor(result.content[0], "arguments")).toMatchObject({
value: { content: "authoritative" },
writable: true,
});
});

it("preserves tool-call metadata received only on toolcall_end", async () => {
const proxyEvents: ProxyAssistantMessageEvent[] = [
{ type: "start" },
Expand Down
4 changes: 3 additions & 1 deletion packages/ai/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -574,7 +574,9 @@ context.messages.push({

### Streaming Tool Calls with Partial JSON

During streaming, tool call arguments are progressively parsed as they arrive. This enables real-time UI updates before the complete arguments are available:
During streaming, tool call arguments are parsed on demand when `arguments` is read, and cached until the next argument delta. This enables real-time UI updates before the complete arguments are available without reparsing growing JSON for delta-only consumers. Read partial arguments only when needed; reading them after every delta still reparses each growing prefix. Terminal messages materialize any unread arguments, including on errors and aborts.

For example:

```typescript
const s = models.stream(model, context);
Expand Down
16 changes: 13 additions & 3 deletions packages/ai/src/api/anthropic-messages.ts
Original file line number Diff line number Diff line change
Expand Up @@ -37,6 +37,7 @@ import { appendAssistantMessageDiagnostic } from "../utils/diagnostics.ts";
import { AssistantMessageEventStream } from "../utils/event-stream.ts";
import { headersToRecord } from "../utils/headers.ts";
import { parseJsonWithRepair, parseStreamingJson } from "../utils/json-parse.ts";
import { createPendingToolCall, type PendingToolCall } from "../utils/pending-tool-call.ts";
import { getPiUserAgent } from "../utils/pi-user-agent.ts";
import { getProviderEnvValue } from "../utils/provider-env.ts";
import { retryProviderRequest } from "../utils/provider-retry.ts";
Expand Down Expand Up @@ -530,6 +531,7 @@ export const stream: StreamFunction<"anthropic-messages", AnthropicOptions> = (
timestamp: Date.now(),
};

const pendingCalls = new Map<ToolCall, PendingToolCall>();
try {
let client: Anthropic;
let isOAuth: boolean;
Expand Down Expand Up @@ -649,7 +651,7 @@ export const stream: StreamFunction<"anthropic-messages", AnthropicOptions> = (
output.content.push(block);
stream.push({ type: "thinking_start", contentIndex: output.content.length - 1, partial: output });
} else if (event.content_block.type === "tool_use") {
const block: Block = {
const pending = createPendingToolCall({
type: "toolCall",
id: event.content_block.id,
name: isOAuth
Expand All @@ -658,7 +660,9 @@ export const stream: StreamFunction<"anthropic-messages", AnthropicOptions> = (
arguments: (event.content_block.input as Record<string, any>) ?? {},
partialJson: "",
index: event.index,
};
});
const block: Block = pending.toolCall;
pendingCalls.set(block, pending);
output.content.push(block);
stream.push({ type: "toolcall_start", contentIndex: output.content.length - 1, partial: output });
}
Expand Down Expand Up @@ -692,7 +696,7 @@ export const stream: StreamFunction<"anthropic-messages", AnthropicOptions> = (
const block = blocks[index];
if (block && block.type === "toolCall") {
block.partialJson += event.delta.partial_json;
block.arguments = parseStreamingJson(block.partialJson);
pendingCalls.get(block)!.setJson(block.partialJson);
stream.push({
type: "toolcall_delta",
contentIndex: index,
Expand Down Expand Up @@ -729,6 +733,8 @@ export const stream: StreamFunction<"anthropic-messages", AnthropicOptions> = (
});
} else if (block.type === "toolCall") {
block.arguments = parseStreamingJson(block.partialJson);
pendingCalls.get(block)!.finish();
pendingCalls.delete(block);
// Finalize in-place and strip the scratch buffer so replay only
// carries parsed arguments.
delete (block as { partialJson?: string }).partialJson;
Expand Down Expand Up @@ -803,9 +809,13 @@ export const stream: StreamFunction<"anthropic-messages", AnthropicOptions> = (
});
}

for (const pending of pendingCalls.values()) pending.finish();
pendingCalls.clear();
stream.push({ type: "done", reason: output.stopReason, message: output });
stream.end();
} catch (error) {
for (const pending of pendingCalls.values()) pending.finish();
pendingCalls.clear();
for (const block of output.content) {
delete (block as { index?: number }).index;
// partialJson is only a streaming scratch buffer; never persist it.
Expand Down
Loading
Loading