fix(openai): preserve Chat response hook lifecycle (#117056)

This commit is contained in:
Peter Steinberger
2026-08-09 09:18:16 -07:00
committed by GitHub
parent a39c24463e
commit 9bbbeb0489
6 changed files with 359 additions and 21 deletions
+21 -10
View File
@@ -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();
@@ -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<typeof vi.fn>;
assertListenersRemoved(): void;
};
function installResponse(): RequestLifecycle {
let addAbortListener: MockInstance<AbortSignal["addEventListener"]> | undefined;
let removeAbortListener: MockInstance<AbortSignal["removeEventListener"]> | 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<T>(promise: Promise<T>, timeoutMs = 250): Promise<T> {
let timer: ReturnType<typeof setTimeout> | undefined;
try {
return await Promise.race([
promise,
new Promise<never>((_resolve, reject) => {
timer = setTimeout(() => reject(new Error("response hook lifecycle timed out")), timeoutMs);
}),
]);
} finally {
clearTimeout(timer);
}
}
let previousHost: ReturnType<typeof getAiTransportHost>;
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<void>((resolve) => {
continueHook = resolve;
});
const onResponse = vi.fn<NonNullable<StreamOptions["onResponse"]>>(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<NonNullable<StreamOptions["onResponse"]>>(() => {
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<void>(() => {}));
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<void>((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<void>((resolve) => {
setImmediate(resolve);
});
expect(eventTypes).toEqual(["error"]);
lifecycle.assertListenersRemoved();
},
);
});
@@ -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<ChatCompletionChunk>;
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),
@@ -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<void>) | undefined {
return onResponse
? () =>
onResponse({ status: response.status, headers: headersToRecord(response.headers) }, model)
: undefined;
}
type ModelStreamCooperativeScheduler = {
afterEvent: () => Promise<void>;
};
@@ -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<T>(params: {
stream: AsyncIterable<T>;
signal: AbortSignal;
abort: (reason: Error) => void;
hook?: () => void | Promise<void>;
onReady: () => void;
}): AsyncIterable<T> {
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<never>((_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;
@@ -104,8 +104,8 @@ export function withFirstStreamEventTimeout<T>(
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);