mirror of
https://github.com/openclaw/openclaw.git
synced 2026-08-12 21:53:00 -06:00
fix(openai): preserve Chat response hook lifecycle (#117056)
This commit is contained in:
committed by
GitHub
parent
a39c24463e
commit
9bbbeb0489
@@ -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);
|
||||
|
||||
Reference in New Issue
Block a user