From 9bbbeb048966604fdd3ab896a204c08c42e22e90 Mon Sep 17 00:00:00 2001 From: Peter Steinberger Date: Sun, 9 Aug 2026 09:18:16 -0700 Subject: [PATCH] fix(openai): preserve Chat response hook lifecycle (#117056) --- .../ai/src/providers/openai-completions.ts | 31 ++- ...ompletions-transport.response-hook.test.ts | 261 ++++++++++++++++++ .../openai-completions-transport.ts | 33 ++- .../src/transports/openai-transport-shared.ts | 12 + .../src/transports/transport-stream-shared.ts | 41 +++ .../src/utils/stream-first-event-timeout.ts | 2 +- 6 files changed, 359 insertions(+), 21 deletions(-) create mode 100644 packages/ai/src/transports/openai-completions-transport.response-hook.test.ts diff --git a/packages/ai/src/providers/openai-completions.ts b/packages/ai/src/providers/openai-completions.ts index 47b300d47521..e0c493aa0155 100644 --- a/packages/ai/src/providers/openai-completions.ts +++ b/packages/ai/src/providers/openai-completions.ts @@ -19,11 +19,16 @@ import { } from "../transports/openai-completions-compat.js"; import { resolveOpenAIReasoningEffortMap } from "../transports/openai-reasoning-compat.js"; import { + createOpenAIResponseHook, isOpenAICompletionsThinkingEnabled, parseOpenAICompletionsUsage, readOpenAICompletionsContentDeltas, } from "../transports/openai-transport-shared.js"; -import { transportAbortError } from "../transports/transport-stream-shared.js"; +import { + assignTransportErrorDetails, + transportAbortError, + withProviderResponseHook, +} from "../transports/transport-stream-shared.js"; import type { AssistantMessage, CacheRetention, @@ -44,7 +49,6 @@ import { type PendingCommentaryTags, } from "../utils/assistant-text-phase.js"; import { AssistantMessageEventStream } from "../utils/event-stream.js"; -import { headersToRecord } from "../utils/headers.js"; import { parseStreamingJson } from "../utils/json-parse.js"; import { notifyLlmRequestActivity } from "../utils/llm-request-activity.js"; import { formatProviderError } from "../utils/provider-error.js"; @@ -187,11 +191,13 @@ export const streamOpenAICompletions: StreamFunction< requestOptions, ) .withResponse(); - await options?.onResponse?.( - { status: response.status, headers: headersToRecord(response.headers) }, - model, - ); - stream.push({ type: "start", partial: output }); + const hookedOpenAIStream = withProviderResponseHook({ + stream: openaiStream, + signal: firstEventAbort.signal, + abort: firstEventAbort.abort, + hook: createOpenAIResponseHook(options?.onResponse, response, model), + onReady: () => stream.push({ type: "start", partial: output }), + }); interface StreamingToolCallBlock extends ToolCall { partialArgs?: string; @@ -380,7 +386,7 @@ export const streamOpenAICompletions: StreamFunction< } }; - const guardedOpenaiStream = withFirstStreamEventTimeout(openaiStream, { + const guardedOpenaiStream = withFirstStreamEventTimeout(hookedOpenAIStream, { provider: model.provider, api: model.api, model: model.id, @@ -605,7 +611,12 @@ export const streamOpenAICompletions: StreamFunction< stream.push({ type: "done", reason: output.stopReason, message: output }); stream.end(); } catch (error) { - output.stopReason = options?.signal?.aborted ? "aborted" : "error"; + const errorReason = options?.signal?.aborted ? "aborted" : "error"; + if (options?.signal?.aborted) { + assignTransportErrorDetails(output, error, options.signal); + } else { + output.stopReason = errorReason; + } finalizeOpenAICompletionsToolCalls(output, { allowSilentToolCallPromotion: false }); for (const block of output.content) { delete (block as { index?: number }).index; @@ -620,7 +631,7 @@ export const streamOpenAICompletions: StreamFunction< if (rawMetadata && !output.errorMessage.includes(rawMetadata)) { output.errorMessage += `\n${rawMetadata}`; } - stream.push({ type: "error", reason: output.stopReason, error: output }); + stream.push({ type: "error", reason: errorReason, error: output }); stream.end(); } finally { firstEventAbort?.dispose(); diff --git a/packages/ai/src/transports/openai-completions-transport.response-hook.test.ts b/packages/ai/src/transports/openai-completions-transport.response-hook.test.ts new file mode 100644 index 000000000000..a5a4c30e9189 --- /dev/null +++ b/packages/ai/src/transports/openai-completions-transport.response-hook.test.ts @@ -0,0 +1,261 @@ +import { afterEach, beforeEach, describe, expect, it, type MockInstance, vi } from "vitest"; +import { configureAiTransportHost, getAiTransportHost } from "../host.js"; +import { streamOpenAICompletions } from "../providers/openai-completions.js"; +import type { AssistantMessageEventStreamLike, Context, Model, StreamOptions } from "../types.js"; +import { createOpenAICompletionsTransportStreamFn } from "./openai-completions-transport.js"; + +const model = { + id: "gpt-5.5", + name: "Response hook lifecycle", + api: "openai-completions", + provider: "openai", + baseUrl: "https://api.openai.com/v1", + reasoning: false, + input: ["text"], + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 }, + contextWindow: 128_000, + maxTokens: 4_096, +} satisfies Model<"openai-completions">; +const context = { + messages: [{ role: "user", content: "hello", timestamp: 1 }], +} satisfies Context; + +type RequestLifecycle = { + requestAborted: ReturnType; + assertListenersRemoved(): void; +}; + +function installResponse(): RequestLifecycle { + let addAbortListener: MockInstance | undefined; + let removeAbortListener: MockInstance | undefined; + const requestAborted = vi.fn(); + const chunk = { + id: "chatcmpl-response-hook", + object: "chat.completion.chunk", + created: 1, + model: model.id, + choices: [{ index: 0, delta: { content: "hello" }, finish_reason: "stop" }], + }; + const body = `data: ${JSON.stringify(chunk)}\n\ndata: [DONE]\n\n`; + configureAiTransportHost({ + buildModelFetch: () => async (_input, init) => { + const signal = init?.signal; + if (signal) { + signal.addEventListener("abort", requestAborted, { once: true }); + addAbortListener = vi.spyOn(signal, "addEventListener"); + removeAbortListener = vi.spyOn(signal, "removeEventListener"); + } + return new Response(body, { + status: 202, + headers: { + "content-type": "text/event-stream", + "x-ratelimit-remaining-requests": "42", + "x-request-id": "req_observable", + }, + }); + }, + }); + return { + requestAborted, + assertListenersRemoved() { + for (const [event, listener] of addAbortListener?.mock.calls ?? []) { + if (event === "abort") { + expect(removeAbortListener).toHaveBeenCalledWith("abort", listener); + } + } + }, + }; +} + +const createManagedStream = createOpenAICompletionsTransportStreamFn(); + +function createManagedFixtureStream( + fixtureModel: Model<"openai-completions">, + fixtureContext: Context, + fixtureOptions?: StreamOptions, +): AssistantMessageEventStreamLike { + const stream = createManagedStream(fixtureModel, fixtureContext, fixtureOptions); + if (stream instanceof Promise) { + throw new Error("OpenAI Chat transport must return its event stream synchronously"); + } + return stream; +} + +async function settleWithin(promise: Promise, timeoutMs = 250): Promise { + let timer: ReturnType | undefined; + try { + return await Promise.race([ + promise, + new Promise((_resolve, reject) => { + timer = setTimeout(() => reject(new Error("response hook lifecycle timed out")), timeoutMs); + }), + ]); + } finally { + clearTimeout(timer); + } +} + +let previousHost: ReturnType; + +beforeEach(() => { + previousHost = getAiTransportHost(); +}); + +afterEach(() => { + configureAiTransportHost(previousHost); +}); + +describe.each([ + { name: "package", createStream: streamOpenAICompletions }, + { name: "managed", createStream: createManagedFixtureStream }, +])("$name OpenAI Chat response hook", ({ createStream }) => { + it("awaits response metadata before exposing the first stream event", async () => { + installResponse(); + const order: string[] = []; + let continueHook!: () => void; + const hookCompleted = new Promise((resolve) => { + continueHook = resolve; + }); + const onResponse = vi.fn>(async () => { + order.push("hook:start"); + await hookCompleted; + order.push("hook:end"); + }); + const stream = createStream(model, context, { apiKey: "fixture-token", onResponse }); + const consume = (async () => { + for await (const event of stream) { + order.push(event.type); + } + })(); + + await vi.waitFor(() => expect(onResponse).toHaveBeenCalledOnce()); + expect(order).toEqual(["hook:start"]); + expect(onResponse).toHaveBeenCalledWith( + { + status: 202, + headers: { + "content-type": "text/event-stream", + "x-ratelimit-remaining-requests": "42", + "x-request-id": "req_observable", + }, + }, + model, + ); + + continueHook(); + await consume; + expect((await stream.result()).stopReason).toBe("stop"); + expect(order.slice(0, 3)).toEqual(["hook:start", "hook:end", "start"]); + }); + + it.each(["throw", "reject"] as const)( + "preserves a hook %s and closes the unread request", + async (failure) => { + const lifecycle = installResponse(); + const hookError = new Error("after_provider_response hook failed"); + const onResponse = vi.fn>(() => { + if (failure === "throw") { + throw hookError; + } + return Promise.reject(hookError); + }); + const stream = createStream(model, context, { apiKey: "fixture-token", onResponse }); + const eventTypes: string[] = []; + for await (const event of stream) { + eventTypes.push(event.type); + } + const result = await stream.result(); + + expect(onResponse).toHaveBeenCalledOnce(); + expect(result).toMatchObject({ + stopReason: "error", + errorMessage: "after_provider_response hook failed", + }); + expect(eventTypes).toEqual(["error"]); + expect(lifecycle.requestAborted).toHaveBeenCalledOnce(); + lifecycle.assertListenersRemoved(); + }, + ); + + it("applies the first-event timeout while the hook is pending", async () => { + const lifecycle = installResponse(); + const onFirstEventTimeout = vi.fn(); + const onResponse = vi.fn(() => new Promise(() => {})); + const stream = createStream(model, context, { + apiKey: "fixture-token", + firstEventTimeoutMs: 20, + onFirstEventTimeout, + onResponse, + }); + const eventTypes: string[] = []; + const consume = (async () => { + for await (const event of stream) { + eventTypes.push(event.type); + } + })(); + const result = await settleWithin(stream.result()); + await consume; + + expect(onResponse).toHaveBeenCalledOnce(); + expect(result.stopReason).toBe("error"); + expect(result.errorMessage).toMatch( + /completions HTTP stream opened but did not deliver a first SSE event within 20ms/, + ); + expect(onFirstEventTimeout).toHaveBeenCalledWith(expect.any(Error)); + expect(eventTypes).toEqual(["error"]); + expect(lifecycle.requestAborted).toHaveBeenCalledOnce(); + lifecycle.assertListenersRemoved(); + }); + + it.each(["resolve", "reject"] as const)( + "keeps caller cancellation terminal after a late hook %s", + async (settlement) => { + const lifecycle = installResponse(); + const controller = new AbortController(); + let settleHook!: () => void; + const pendingHook = new Promise((resolve, reject) => { + settleHook = () => { + if (settlement === "resolve") { + resolve(); + } else { + reject(new Error("late response hook rejection")); + } + }; + }); + const onResponse = vi.fn(() => pendingHook); + const stream = createStream(model, context, { + apiKey: "fixture-token", + signal: controller.signal, + onResponse, + }); + const eventTypes: string[] = []; + const consume = (async () => { + for await (const event of stream) { + eventTypes.push(event.type); + } + })(); + + await vi.waitFor(() => expect(onResponse).toHaveBeenCalledOnce()); + const abortReason = Object.assign(new Error("caller canceled the provider response"), { + code: "CALLER_ABORTED", + }); + controller.abort(abortReason); + const result = await settleWithin(stream.result()); + await consume; + expect(result).toMatchObject({ + stopReason: "aborted", + errorCode: "CALLER_ABORTED", + errorMessage: "caller canceled the provider response", + }); + expect(eventTypes).toEqual(["error"]); + expect(lifecycle.requestAborted).toHaveBeenCalledOnce(); + + settleHook(); + await new Promise((resolve) => { + setImmediate(resolve); + }); + expect(eventTypes).toEqual(["error"]); + lifecycle.assertListenersRemoved(); + }, + ); +}); diff --git a/packages/ai/src/transports/openai-completions-transport.ts b/packages/ai/src/transports/openai-completions-transport.ts index 33a8f3c90df0..b95a4831057b 100644 --- a/packages/ai/src/transports/openai-completions-transport.ts +++ b/packages/ai/src/transports/openai-completions-transport.ts @@ -73,6 +73,7 @@ import { import { GEMINI_THOUGHT_SIGNATURE_VALIDATOR_SKIP, createModelStreamCooperativeScheduler, + createOpenAIResponseHook, isOpenAICompletionsThinkingEnabled, log, parseOpenAICompletionsUsage, @@ -84,7 +85,11 @@ import { type OpenAICompletionsContentDelta as CompletionsReasoningDelta, type OpenAIModeModel, } from "./openai-transport-shared.js"; -import { failTransportStream, finalizeTransportStream } from "./transport-stream-shared.js"; +import { + failTransportStream, + finalizeTransportStream, + withProviderResponseHook, +} from "./transport-stream-shared.js"; import { CHARS_PER_TOKEN_ESTIMATE, estimateStringChars, @@ -342,15 +347,23 @@ export function createOpenAICompletionsTransportStreamFn(): StreamFn { options as OpenAICompletionsOptions | undefined, ); firstEventAbort = createFirstStreamEventAbortController(options?.signal); - const responseStream = (await client.chat.completions.create( - params as never, - buildOpenAISdkRequestOptions(model, firstEventAbort.signal, { - timeoutMs: options?.timeoutMs, - maxRetries: options?.maxRetries, - }), - )) as unknown as AsyncIterable; - stream.push({ type: "start", partial: output as never }); - await processOpenAICompletionsStream(responseStream, output, model, stream, { + const { data: responseStream, response } = await client.chat.completions + .create( + params as unknown as OpenAI.Chat.Completions.ChatCompletionCreateParamsStreaming, + buildOpenAISdkRequestOptions(model, firstEventAbort.signal, { + timeoutMs: options?.timeoutMs, + maxRetries: options?.maxRetries, + }), + ) + .withResponse(); + const hookedResponseStream = withProviderResponseHook({ + stream: responseStream, + signal: firstEventAbort.signal, + abort: firstEventAbort.abort, + hook: createOpenAIResponseHook(options?.onResponse, response, model), + onReady: () => stream.push({ type: "start", partial: output as never }), + }); + await processOpenAICompletionsStream(hookedResponseStream, output, model, stream, { signal: options?.signal, emitReasoning, firstEventTimeoutMs: getFirstStreamEventTimeoutMs(options), diff --git a/packages/ai/src/transports/openai-transport-shared.ts b/packages/ai/src/transports/openai-transport-shared.ts index 1a42dc93ac4e..e0dc4da00c33 100644 --- a/packages/ai/src/transports/openai-transport-shared.ts +++ b/packages/ai/src/transports/openai-transport-shared.ts @@ -5,6 +5,7 @@ import { applyProviderReportedUsageCost, calculateCost } from "../model-utils.js import type { BaseOpenAIStreamOptions } from "../provider-options.js"; /** Shared options, usage shape, cache identity, ordering, and stream scheduling for OpenAI APIs. */ import { clampOpenAIPromptCacheKey } from "../providers/openai-prompt-cache.js"; +import { headersToRecord } from "../utils/headers.js"; import { transportAbortError } from "./transport-stream-shared.js"; export { sortPromptCacheToolsByName as sortTransportToolsByName } from "../utils/prompt-cache-stability.js"; @@ -101,6 +102,17 @@ export function parseOpenAICompletionsUsage( return usage; } +export function createOpenAIResponseHook( + onResponse: BaseOpenAIStreamOptions["onResponse"], + response: Response, + model: Model, +): (() => void | Promise) | undefined { + return onResponse + ? () => + onResponse({ status: response.status, headers: headersToRecord(response.headers) }, model) + : undefined; +} + type ModelStreamCooperativeScheduler = { afterEvent: () => Promise; }; diff --git a/packages/ai/src/transports/transport-stream-shared.ts b/packages/ai/src/transports/transport-stream-shared.ts index 93c779eeb4bd..af45b80e62a8 100644 --- a/packages/ai/src/transports/transport-stream-shared.ts +++ b/packages/ai/src/transports/transport-stream-shared.ts @@ -136,6 +136,47 @@ export function transportAbortError(signal?: AbortSignal): Error { : new Error("Request was aborted"); } +/** Run a provider-response hook before start/body consumption inside the first-event deadline. */ +export function withProviderResponseHook(params: { + stream: AsyncIterable; + signal: AbortSignal; + abort: (reason: Error) => void; + hook?: () => void | Promise; + onReady: () => void; +}): AsyncIterable { + return { + async *[Symbol.asyncIterator]() { + let onAbort: (() => void) | undefined; + try { + if (params.signal.aborted) { + throw transportAbortError(params.signal); + } + if (params.hook) { + await Promise.race([ + Promise.resolve().then(params.hook), + new Promise((_resolve, reject) => { + onAbort = () => reject(transportAbortError(params.signal)); + params.signal.addEventListener("abort", onAbort, { once: true }); + }), + ]); + } + } catch (error) { + params.abort(error instanceof Error ? error : new Error(String(error))); + throw error; + } finally { + if (onAbort) { + params.signal.removeEventListener("abort", onAbort); + } + } + if (params.signal.aborted) { + throw transportAbortError(params.signal); + } + params.onReady(); + yield* params.stream; + }, + }; +} + export function finalizeTransportStream(params: { stream: WritableTransportStream; output: TransportOutputShape; diff --git a/packages/ai/src/utils/stream-first-event-timeout.ts b/packages/ai/src/utils/stream-first-event-timeout.ts index 27b703207a5f..aac8b5fec956 100644 --- a/packages/ai/src/utils/stream-first-event-timeout.ts +++ b/packages/ai/src/utils/stream-first-event-timeout.ts @@ -104,8 +104,8 @@ export function withFirstStreamEventTimeout( timer = setTimeout(() => { const timeoutError = createFirstStreamEventTimeoutError(timeoutContext); timeoutContext.onTimeout?.(timeoutError); - timeoutContext.abort?.(timeoutError); reject(timeoutError); + timeoutContext.abort?.(timeoutError); }, timeoutMs); timer.unref?.(); iterator.next().then(resolve, reject);