diff --git a/packages/ai/CHANGELOG.md b/packages/ai/CHANGELOG.md index febc1bd5ca..52dafa1674 100644 --- a/packages/ai/CHANGELOG.md +++ b/packages/ai/CHANGELOG.md @@ -3,6 +3,7 @@ ## [Unreleased] - Cursor native tool calls (shell/read/write/… oneof variants) now convert their protobuf payloads into plain JSON-safe data before attaching them as toolCall `arguments`: `$typeName` markers are stripped, safe-range bigints become numbers (decimal strings beyond `Number.MAX_SAFE_INTEGER`), byte arrays become base64 strings, and cycles/functions collapse to null. Raw protobuf-es payloads carry `bigint` fields (`fileSize`, `durationMs`, `fileOutputThresholdBytes`, …) that defeat `JSON.stringify`, which broke managed snapshot staging, JSONL transcript persistence, and provider replay — the issue #4578 local-snapshot producer defect class fixed at its producer boundary. - Generic OpenAI-compatible `/v1/models` discovery now reads served context-window and output-limit metadata instead of defaulting every dynamically listed model to the unknown-window sentinel. `max_model_len` (vLLM/SGLang/oMLX), `context_length`, `context_window`, `max_context_length` (LM Studio), and `max_position_embeddings` populate `contextWindow` in that precedence order, while `max_tokens`/`max_output_tokens` populate `maxTokens`; total-window fields never leak into the output-token ceiling. Malformed values (non-finite, zero, negative, non-numeric) are rejected per-field with fallback to the next candidate, so a `1e400`-style catalog entry can no longer poison compaction thresholds or compact-input budgets. +- Codex websocket requests now abort and close their transport when the downstream event-stream consumer returns early (including managed provisional-buffer rejection), so the next turn opens a clean connection instead of inheriting `websocket request already in progress` (#4534). - Refreshed the bundled ZAI catalog with GLM-5.3 and made it the provider's default model. - Added the typed `local_snapshot_failure` and `local_buffer_overflow` assistant error kinds so downstream retry policy can distinguish local event-snapshot and staging-buffer failures from provider failures. - Anthropic first-event timeouts now report safe elapsed time, serialized request bytes, canonical-vs-custom endpoint class, and the `PI_STREAM_FIRST_EVENT_TIMEOUT_MS` override without exposing URL credentials, query tokens, or body content. Large requests through custom endpoints receive one bounded two-minute observation grace so a slightly later proxy 529 can surface without extending explicit-zero, small-request, or canonical deadlines; full-window multi-megabyte requests are never automatically re-uploaded and small requests get at most one session replay. Credit: @probepark (#4464). diff --git a/packages/ai/src/providers/openai-codex-responses.ts b/packages/ai/src/providers/openai-codex-responses.ts index c651036f9f..b09c7d9a51 100644 --- a/packages/ai/src/providers/openai-codex-responses.ts +++ b/packages/ai/src/providers/openai-codex-responses.ts @@ -1855,24 +1855,29 @@ export const streamOpenAICodexResponses: StreamFunction<"openai-codex-responses" context: Context, options?: OpenAICodexResponsesOptions, ): AssistantMessageEventStream => { - const stream = new AssistantMessageEventStream(); + const consumerAbortController = new AbortController(); + const stream = new AssistantMessageEventStream(() => consumerAbortController.abort()); + const signal = options?.signal + ? AbortSignal.any([options.signal, consumerAbortController.signal]) + : consumerAbortController.signal; + const streamOptions = { ...options, signal }; (async () => { const startTime = Date.now(); const output = createAssistantOutput(model); - const requestSetup = createRequestSetup(options); + const requestSetup = createRequestSetup(streamOptions); let processingContext: CodexStreamProcessingContext | undefined; try { - const requestContext = await buildCodexRequestContext(model, context, options, output); + const requestContext = await buildCodexRequestContext(model, context, streamOptions, output); let initialTransport: CodexInitialTransport; try { - initialTransport = await openInitialCodexEventStream(model, options, requestSetup, requestContext); + initialTransport = await openInitialCodexEventStream(model, streamOptions, requestSetup, requestContext); } catch (error) { - if (options?.fallbackManaged) throw error; + if (streamOptions.fallbackManaged) throw error; initialTransport = await retryCodexInitialTransportWithoutToolChoice( model, - options, + streamOptions, requestSetup, requestContext, stream, @@ -1891,7 +1896,7 @@ export const streamOpenAICodexResponses: StreamFunction<"openai-codex-responses" model, output, stream, - options, + options: streamOptions, requestSetup, requestContext, startTime, @@ -1909,7 +1914,7 @@ export const streamOpenAICodexResponses: StreamFunction<"openai-codex-responses" model, output, stream, - options, + options: streamOptions, requestSetup, requestContext: { apiKey: "", diff --git a/packages/ai/src/providers/register-builtins.ts b/packages/ai/src/providers/register-builtins.ts index 9f61623773..ae8c74a999 100644 --- a/packages/ai/src/providers/register-builtins.ts +++ b/packages/ai/src/providers/register-builtins.ts @@ -339,12 +339,15 @@ function createLazyStream( limits?: LazyStreamLimits, ): (model: Model, context: Context, options: OptionsForApi) => EventStreamImpl { return (model, context, options) => { - const outer = new EventStreamImpl(); + let abortTracker: AbortSourceTracker | undefined; + const outer = new EventStreamImpl(() => + abortTracker?.abortLocally(new Error("Provider stream consumer stopped before completion")), + ); const streamOptions = (options ?? {}) as OptionsForApi; loadModule() .then(module => { - const abortTracker = createAbortSourceTracker(streamOptions.signal); + abortTracker = createAbortSourceTracker(streamOptions.signal); const providerOptions = { ...streamOptions, signal: abortTracker.requestSignal } as OptionsForApi; const inner = module.stream(model, context, providerOptions); forwardStream(outer, inner, model, streamOptions, abortTracker, limits); diff --git a/packages/ai/src/stream.ts b/packages/ai/src/stream.ts index 1a555b24fa..4eb5b9d01b 100644 --- a/packages/ai/src/stream.ts +++ b/packages/ai/src/stream.ts @@ -493,7 +493,11 @@ export function streamSimple( } const retryApiKey = options?.onAuthError ? (options.apiKey ?? getEnvApiKey(model.provider)) : undefined; if (retryApiKey) { - const outer = new AssistantMessageEventStream(); + const consumerAbortController = new AbortController(); + const outer = new AssistantMessageEventStream(() => consumerAbortController.abort()); + const requestSignal = options?.signal + ? AbortSignal.any([options.signal, consumerAbortController.signal]) + : consumerAbortController.signal; const onAuthError = options!.onAuthError!; const runAttempt = async (apiKey: string, captureAuthFailure: boolean): Promise => { const bufferedEvents: AssistantMessageEvent[] = []; @@ -504,7 +508,12 @@ export function streamSimple( }; try { - const inner = streamSimple(model, context, { ...options, apiKey, onAuthError: undefined }); + const inner = streamSimple(model, context, { + ...options, + apiKey, + onAuthError: undefined, + signal: requestSignal, + }); for await (const event of inner) { if (!emittedReplayUnsafeEvent && event.type === "start") { bufferedEvents.push(event); diff --git a/packages/ai/src/utils/event-stream.ts b/packages/ai/src/utils/event-stream.ts index f591d7c13f..f235b837d0 100644 --- a/packages/ai/src/utils/event-stream.ts +++ b/packages/ai/src/utils/event-stream.ts @@ -35,8 +35,9 @@ export class EventStream implements AsyncIterable { rejectFinalResult!: (err: unknown) => void; isComplete: (event: T) => boolean; extractResult: (event: T) => R; + #onConsumerClose?: () => void; - constructor(isComplete: (event: T) => boolean, extractResult: (event: T) => R) { + constructor(isComplete: (event: T) => boolean, extractResult: (event: T) => R, onConsumerClose?: () => void) { const { promise, resolve, reject } = Promise.withResolvers(); // Prevent an unhandled rejection when fail() is called but nobody awaits result(). // Callers who do await result() still receive the rejection normally. @@ -46,6 +47,7 @@ export class EventStream implements AsyncIterable { this.rejectFinalResult = reject; this.isComplete = isComplete; this.extractResult = extractResult; + this.#onConsumerClose = onConsumerClose; } #enqueue(node: QueueNode): void { @@ -240,6 +242,7 @@ export class EventStream implements AsyncIterable { } finally { this.#activeConsumerCount -= 1; this.#settleAllConsumerDrains("reject", new Error("Event stream consumer stopped before drain completed")); + if (!this.done) this.#onConsumerClose?.(); } } @@ -249,7 +252,7 @@ export class EventStream implements AsyncIterable { } export class AssistantMessageEventStream extends EventStream { - constructor() { + constructor(onConsumerClose?: () => void) { super( event => event.type === "done" || event.type === "error", event => { @@ -260,6 +263,7 @@ export class AssistantMessageEventStream extends EventStream { expect(transportDetails.fallbackCount).toBe(1); }); + it("releases an in-flight websocket request when the stream consumer returns early", async () => { + const tempDir = TempDir.createSync("@pi-codex-stream-"); + setAgentDir(tempDir.path()); + const token = createCodexTestToken(); + const providerSessionState = new Map(); + const sentTypesByConnection: string[][] = []; + let constructorCount = 0; + + class ConsumerReturnWebSocket extends MockWebSocket { + #connectionIndex: number; + + constructor(url: string, options?: { headers?: WsHeaders }) { + super(url, options); + this.#connectionIndex = constructorCount++; + sentTypesByConnection[this.#connectionIndex] = []; + this.scheduleOpen(); + } + + send(data: string): void { + const request = JSON.parse(data) as { type?: string }; + const requestType = typeof request.type === "string" ? request.type : ""; + sentTypesByConnection[this.#connectionIndex]?.push(requestType); + if (this.#connectionIndex === 0) { + this.sendJson({ + type: "response.output_item.added", + item: { type: "message", id: "msg_1", role: "assistant", status: "in_progress", content: [] }, + }); + this.sendJson({ type: "response.content_part.added", part: { type: "output_text", text: "" } }); + this.sendJson({ type: "response.output_text.delta", delta: "oversized provisional payload" }); + return; + } + this.emitCodexResponse({ messageId: "msg_2", responseId: "resp_2", text: "clean successor" }); + } + } + + global.WebSocket = ConsumerReturnWebSocket as unknown as typeof WebSocket; + global.fetch = vi.fn(async () => { + throw new Error("SSE fallback should not be called"); + }) as unknown as typeof fetch; + const model = createCodexTestModel("https://chatgpt.com/backend-api"); + const first = streamOpenAICodexResponses(model, createCodexTestContext(), { + apiKey: token, + sessionId: "ws-consumer-return-session", + providerSessionState, + }); + const iterator = first[Symbol.asyncIterator](); + for (let i = 0; i < 3; i++) { + const event = await iterator.next(); + expect(event.done).toBe(false); + } + await iterator.return?.(); + + const successor = await streamOpenAICodexResponses(model, createCodexTestContext(), { + apiKey: token, + sessionId: "ws-consumer-return-session", + providerSessionState, + }).result(); + + expect(successor.stopReason).toBe("stop"); + expect(successor.content).toEqual([expect.objectContaining({ type: "text", text: "clean successor" })]); + expect(constructorCount).toBe(2); + expect(sentTypesByConnection).toEqual([["response.create"], ["response.create"]]); + }); + it("resets websocket append state after an aborted request closes the connection", async () => { const tempDir = TempDir.createSync("@pi-codex-stream-"); setAgentDir(tempDir.path()); diff --git a/packages/ai/test/register-builtins.test.ts b/packages/ai/test/register-builtins.test.ts index 541e83102a..0319c26629 100644 --- a/packages/ai/test/register-builtins.test.ts +++ b/packages/ai/test/register-builtins.test.ts @@ -88,6 +88,43 @@ describe("register-builtins lazy streams", () => { expect(result).toEqual(finalMessage); }); + it("aborts the lazy provider request when the public stream consumer returns early", async () => { + const partialMessage = createAssistantMessage("stop"); + let providerSignal: AbortSignal | undefined; + let providerAborted = false; + const source = { + async *[Symbol.asyncIterator]() { + yield { type: "start", partial: partialMessage } as const; + const { promise, reject } = Promise.withResolvers(); + providerSignal?.addEventListener( + "abort", + () => { + providerAborted = true; + reject(new Error("Request was aborted")); + }, + { once: true }, + ); + await promise; + }, + } as unknown as AssistantMessageEventStream; + + setBedrockProviderModule({ + streamBedrock: (_model, _context, options) => { + providerSignal = options.signal; + return source; + }, + }); + + const stream = streamBedrock(createModel(), baseContext, {}); + const iterator = stream[Symbol.asyncIterator](); + expect((await iterator.next()).done).toBe(false); + await iterator.return?.(); + await Bun.sleep(0); + + expect(providerSignal?.aborted).toBe(true); + expect(providerAborted).toBe(true); + }); + it("turns iterator failures into terminal error results", async () => { const partialMessage = createAssistantMessage("stop"); const source = { diff --git a/packages/ai/test/stream-auth-retry.test.ts b/packages/ai/test/stream-auth-retry.test.ts index 5be64c28c0..4e00c7097b 100644 --- a/packages/ai/test/stream-auth-retry.test.ts +++ b/packages/ai/test/stream-auth-retry.test.ts @@ -64,6 +64,45 @@ describe("streamSimple auth retry", () => { unregisterCustomApis(SOURCE_ID); }); + it("aborts the active auth-retry request when the public consumer returns early", async () => { + let providerSignal: AbortSignal | undefined; + let providerAborted = false; + registerCustomApi( + API, + (_model: Model, _context: Context, options?: SimpleStreamOptions) => { + providerSignal = options?.signal; + const stream = new AssistantMessageEventStream(); + queueMicrotask(() => { + const message = assistant(["partial"]); + stream.push({ type: "start", partial: message }); + stream.push({ type: "text_delta", contentIndex: 0, delta: "partial", partial: message }); + }); + providerSignal?.addEventListener( + "abort", + () => { + providerAborted = true; + stream.fail(new Error("Request was aborted")); + }, + { once: true }, + ); + return stream; + }, + SOURCE_ID, + ); + + const stream = streamSimple(model(), context, { + apiKey: "old-key", + onAuthError: async () => "new-key", + }); + const iterator = stream[Symbol.asyncIterator](); + expect((await iterator.next()).done).toBe(false); + await iterator.return?.(); + await Bun.sleep(0); + + expect(providerSignal?.aborted).toBe(true); + expect(providerAborted).toBe(true); + }); + it("retries once with a fresh key when 401 happens before the first event", async () => { const keys: Array = []; let authCalls = 0;