Skip to content
Merged
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
26 changes: 20 additions & 6 deletions packages/ai/src/api/openai-codex-responses.ts
Original file line number Diff line number Diff line change
Expand Up @@ -67,6 +67,7 @@ const REQUEST_COMPRESSION_ZSTD_LEVEL = 3;
const CODEX_TOOL_CALL_PROVIDERS = new Set(["openai", "openai-codex", "opencode"]);
const WEBSOCKET_MESSAGE_TOO_BIG_CLOSE_CODE = 1009;
const WEBSOCKET_CONNECTION_LIMIT_REACHED_CODE = "websocket_connection_limit_reached";
const PREVIOUS_RESPONSE_NOT_FOUND_CODE = "previous_response_not_found";

const CODEX_RESPONSE_STATUSES = new Set<CodexResponseStatus>([
"completed",
Expand Down Expand Up @@ -273,6 +274,7 @@ export const stream: StreamFunction<"openai-codex-responses", OpenAICodexRespons
const httpTimeoutMs = normalizeTimeoutMs(options?.timeoutMs);
const websocketConnectTimeoutMs = normalizeTimeoutMs(options?.websocketConnectTimeoutMs);
const transport = options?.transport || "auto";
let startEmitted = false;
const websocketDisabledForSession = transport !== "sse" && isWebSocketSseFallbackActive(options?.sessionId);
if (websocketDisabledForSession) {
recordWebSocketSseFallback(options?.sessionId);
Expand All @@ -281,6 +283,7 @@ export const stream: StreamFunction<"openai-codex-responses", OpenAICodexRespons
if (transport !== "sse" && !websocketDisabledForSession) {
let websocketStarted = false;
let retriedWebSocketConnectionLimit = false;
let retriedMissingWebSocketContinuation = false;
while (true) {
websocketStarted = false;
try {
Expand All @@ -293,6 +296,10 @@ export const stream: StreamFunction<"openai-codex-responses", OpenAICodexRespons
model,
() => {
websocketStarted = true;
if (!startEmitted) {
startEmitted = true;
stream.push({ type: "start", partial: output });
}
},
httpTimeoutMs,
websocketConnectTimeoutMs,
Expand All @@ -312,6 +319,11 @@ export const stream: StreamFunction<"openai-codex-responses", OpenAICodexRespons
} catch (error) {
const aborted = options?.signal?.aborted;
const connectionLimitBeforeStart = !websocketStarted && isWebSocketConnectionLimitReachedError(error);
const previousResponseNotFound = isPreviousResponseNotFoundError(error);
if (!aborted && previousResponseNotFound && !retriedMissingWebSocketContinuation) {
retriedMissingWebSocketContinuation = true;
continue;
}
if (!aborted && connectionLimitBeforeStart && !retriedWebSocketConnectionLimit) {
retriedWebSocketConnectionLimit = true;
continue;
Expand Down Expand Up @@ -432,7 +444,10 @@ export const stream: StreamFunction<"openai-codex-responses", OpenAICodexRespons
throw new Error("No response body");
}

stream.push({ type: "start", partial: output });
if (!startEmitted) {
startEmitted = true;
stream.push({ type: "start", partial: output });
}
await processStream(response, output, stream, model, options);

if (options?.signal?.aborted) {
Expand Down Expand Up @@ -636,6 +651,10 @@ function isWebSocketConnectionLimitReachedError(error: unknown): boolean {
return error instanceof CodexApiError && error.code === WEBSOCKET_CONNECTION_LIMIT_REACHED_CODE;
}

function isPreviousResponseNotFoundError(error: unknown): boolean {
return error instanceof CodexApiError && error.code === PREVIOUS_RESPONSE_NOT_FOUND_CODE;
}

function extractCodexEventError(event: Record<string, unknown>): { code?: string; message?: string } {
const nested = event.error && typeof event.error === "object" ? (event.error as Record<string, unknown>) : undefined;
return {
Expand Down Expand Up @@ -1358,16 +1377,13 @@ function buildCachedWebSocketRequestBody(entry: CachedWebSocketConnection, body:

async function* startWebSocketOutputOnFirstEvent(
events: AsyncIterable<ResponseStreamEvent>,
output: AssistantMessage,
stream: AssistantMessageEventStream,
onStart: () => void,
): AsyncGenerator<ResponseStreamEvent> {
let started = false;
for await (const event of events) {
if (!started) {
started = true;
onStart();
stream.push({ type: "start", partial: output });
}
yield event;
}
Expand Down Expand Up @@ -1422,8 +1438,6 @@ async function processWebSocketStream(
await processResponsesStream(
startWebSocketOutputOnFirstEvent(
mapCodexEvents(parseWebSocket(socket, options?.signal, idleTimeoutMs)),
output,
stream,
onStart,
),
output,
Expand Down
218 changes: 218 additions & 0 deletions packages/ai/test/openai-codex-stream.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -1802,6 +1802,224 @@ describe("openai-codex streaming", () => {
});
});

it.each(["websocket", "sse"] as const)(
"recovers a missing cached websocket continuation via %s",
async (recoveryTransport) => {
const token = mockToken();
const sessionId = `missing-continuation-${recoveryTransport}`;
const encoder = new TextEncoder();
const fetchMock = vi.fn(
async () =>
new Response(
new ReadableStream<Uint8Array>({
start(controller) {
controller.enqueue(encoder.encode(buildSSEPayload({ status: "completed" })));
controller.close();
},
}),
{ status: 200, headers: { "content-type": "text/event-stream" } },
),
);
vi.stubGlobal("fetch", fetchMock);
const sentBodies: Array<{
connectionId: number;
input: unknown[];
previous_response_id?: string;
}> = [];
let connections = 0;

class MockWebSocket {
static OPEN = 1;
static CLOSED = 3;
readyState = MockWebSocket.OPEN;
private readonly connectionId = ++connections;
private listeners = new Map<string, Set<(event: unknown) => void>>();

constructor(_url: string, _protocols?: string | string[] | { headers?: Record<string, string> }) {
queueMicrotask(() => this.dispatch("open", {}));
}

addEventListener(type: string, listener: (event: unknown) => void): void {
let listeners = this.listeners.get(type);
if (!listeners) {
listeners = new Set();
this.listeners.set(type, listeners);
}
listeners.add(listener);
}

removeEventListener(type: string, listener: (event: unknown) => void): void {
this.listeners.get(type)?.delete(listener);
}

send(data: string): void {
const body = JSON.parse(data) as { input: unknown[]; previous_response_id?: string };
sentBodies.push({ ...body, connectionId: this.connectionId });
if (sentBodies.length === 2) {
this.dispatchEvents([
{
type: "codex.rate_limits",
plan_type: "plus",
rate_limits: {
allowed: true,
limit_reached: false,
primary: {
used_percent: 7,
window_minutes: 10080,
reset_after_seconds: 556112,
reset_at: 1785269351,
},
secondary: null,
},
code_review_rate_limits: null,
additional_rate_limits: null,
credits: { has_credits: false, unlimited: false, balance: "0" },
promo: null,
},
{
type: "error",
status: 400,
error: {
code: "previous_response_not_found",
message: "Previous response with id 'resp_1' not found.",
param: "previous_response_id",
},
},
]);
return;
}
if (sentBodies.length === 3 && recoveryTransport === "sse") {
queueMicrotask(() => this.dispatch("error", { message: "retry websocket failed" }));
return;
}

const response =
sentBodies.length === 1
? { responseId: "resp_1", messageId: "msg_1", text: "Hello" }
: { responseId: "resp_2", messageId: "msg_2", text: "Recovered" };
this.dispatchEvents([
{ type: "response.created", response: { id: response.responseId } },
{
type: "response.output_item.added",
output_index: 0,
item: {
type: "message",
id: response.messageId,
role: "assistant",
status: "in_progress",
content: [],
},
},
{
type: "response.output_item.done",
output_index: 0,
item: {
type: "message",
id: response.messageId,
role: "assistant",
status: "completed",
content: [{ type: "output_text", text: response.text }],
},
},
{
type: "response.completed",
response: {
id: response.responseId,
status: "completed",
usage: { input_tokens: 5, output_tokens: 3, total_tokens: 8 },
},
},
]);
}

close(): void {
this.readyState = MockWebSocket.CLOSED;
}

private dispatchEvents(events: unknown[]): void {
queueMicrotask(() => {
for (const event of events) {
this.dispatch("message", { data: JSON.stringify(event) });
}
});
}

private dispatch(type: string, event: unknown): void {
for (const listener of this.listeners.get(type) ?? []) {
listener(event);
}
}
}

vi.stubGlobal("WebSocket", MockWebSocket);

const model: Model<"openai-codex-responses"> = {
id: "gpt-5.1-codex",
name: "GPT-5.1 Codex",
api: "openai-codex-responses",
provider: "openai-codex",
baseUrl: "https://chatgpt.com/backend-api",
reasoning: true,
input: ["text"],
cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 },
contextWindow: 400000,
maxTokens: 128000,
};
const firstContext: Context = {
systemPrompt: "You are a helpful assistant.",
messages: [{ role: "user", content: "Say hello", timestamp: 1 }],
};

const first = await streamOpenAICodexResponses(model, firstContext, {
apiKey: token,
sessionId,
transport: "websocket-cached",
}).result();
const secondContext: Context = {
systemPrompt: "You are a helpful assistant.",
messages: [...firstContext.messages, first, { role: "user", content: "Now finish", timestamp: 2 }],
};
const eventTypes: string[] = [];
const secondStream = streamOpenAICodexResponses(model, secondContext, {
apiKey: token,
sessionId,
transport: "websocket-cached",
});
for await (const event of secondStream) {
eventTypes.push(event.type);
}
const second = await secondStream.result();

expect(second.stopReason).toBe("stop");
expect(second.content.find((content) => content.type === "text")?.text).toBe(
recoveryTransport === "sse" ? "Hello" : "Recovered",
);
expect(eventTypes.filter((type) => type === "start")).toHaveLength(1);
expect(eventTypes).not.toContain("error");
expect(connections).toBe(2);
expect(sentBodies).toHaveLength(3);
expect(sentBodies.map((body) => body.connectionId)).toEqual([1, 1, 2]);
expect(sentBodies[1].previous_response_id).toBe("resp_1");
expect(sentBodies[1].input).toEqual([{ role: "user", content: [{ type: "input_text", text: "Now finish" }] }]);
expect(sentBodies[2].previous_response_id).toBeUndefined();
expect(sentBodies[2].input).toHaveLength(3);
expect(sentBodies[2].input.at(-1)).toEqual({
role: "user",
content: [{ type: "input_text", text: "Now finish" }],
});
expect(fetchMock).toHaveBeenCalledTimes(recoveryTransport === "sse" ? 1 : 0);
expect(getOpenAICodexWebSocketDebugStats(sessionId)).toMatchObject({
requests: 3,
connectionsCreated: 2,
connectionsReused: 1,
fullContextRequests: 2,
deltaRequests: 1,
websocketFailures: recoveryTransport === "sse" ? 1 : 0,
sseFallbacks: recoveryTransport === "sse" ? 1 : 0,
});
},
);

it.each([
["retry-after-ms", () => ({ "content-type": "application/json", "retry-after-ms": "1500" }), 1500],
["retry-after seconds", () => ({ "content-type": "application/json", "retry-after": "60" }), 60_000],
Expand Down
Loading