diff --git a/bun.lock b/bun.lock index 138cc25..b8da404 100644 --- a/bun.lock +++ b/bun.lock @@ -9,11 +9,13 @@ "@libsql/client": "^0.15.15", "drizzle-orm": "^0.45.1", "hono": "^4.12.2", + "ws": "8.21.0", "zod": "^4.3.6", }, "devDependencies": { "@biomejs/biome": "2.4.0", "@types/bun": "^1.3.9", + "@types/ws": "8.18.1", "@typescript/native-preview": "^7.0.0-dev.20260223.1", "drizzle-kit": "^0.31.9", "lefthook": "^2.1.1", @@ -262,12 +264,14 @@ "web-streams-polyfill": ["web-streams-polyfill@3.3.3", "", {}, "sha512-d2JWLCivmZYTSIoge9MsgFCZrt571BikcWGYkjC1khllbTeDlGqZ2D8vD8E/lJa8WGWbb7Plm8/XJYV7IJHZZw=="], - "ws": ["ws@8.18.0", "", { "peerDependencies": { "bufferutil": "^4.0.1", "utf-8-validate": ">=5.0.2" }, "optionalPeers": ["bufferutil", "utf-8-validate"] }, "sha512-8VbfWfHLbbwu3+N6OKsOMpBdT4kXPDDB9cJk2bJ6mh9ucxdlnNvH1e+roYkKmN9Nxw2yjz7VzeO9oOz2zJ04Pw=="], + "ws": ["ws@8.21.0", "", { "peerDependencies": { "bufferutil": "^4.0.1", "utf-8-validate": ">=5.0.2" }, "optionalPeers": ["bufferutil", "utf-8-validate"] }, "sha512-Vsp28b7DRcimFQvrqu2Wek3z1iYxDCWqHYB8Qsnk/S4RfaCQzPGPyBNuVjJV3cd6UiKtUtp6sNM77gWvzcCH+g=="], "zod": ["zod@4.3.6", "", {}, "sha512-rftlrkhHZOcjDwkGlnUtZZkvaPHCsDATp4pGpuOOMDaTdDDXF91wuVDJoWoPsKX/3YPQ5fHuF3STjcYyKr+Qhg=="], "@esbuild-kit/core-utils/esbuild": ["esbuild@0.18.20", "", { "optionalDependencies": { "@esbuild/android-arm": "0.18.20", "@esbuild/android-arm64": "0.18.20", "@esbuild/android-x64": "0.18.20", "@esbuild/darwin-arm64": "0.18.20", "@esbuild/darwin-x64": "0.18.20", "@esbuild/freebsd-arm64": "0.18.20", "@esbuild/freebsd-x64": "0.18.20", "@esbuild/linux-arm": "0.18.20", "@esbuild/linux-arm64": "0.18.20", "@esbuild/linux-ia32": "0.18.20", "@esbuild/linux-loong64": "0.18.20", "@esbuild/linux-mips64el": "0.18.20", "@esbuild/linux-ppc64": "0.18.20", "@esbuild/linux-riscv64": "0.18.20", "@esbuild/linux-s390x": "0.18.20", "@esbuild/linux-x64": "0.18.20", "@esbuild/netbsd-x64": "0.18.20", "@esbuild/openbsd-x64": "0.18.20", "@esbuild/sunos-x64": "0.18.20", "@esbuild/win32-arm64": "0.18.20", "@esbuild/win32-ia32": "0.18.20", "@esbuild/win32-x64": "0.18.20" }, "bin": { "esbuild": "bin/esbuild" } }, "sha512-ceqxoedUrcayh7Y7ZX6NdbbDzGROiyVBgC4PriJThBKSVPWnnFHZAkfI1lJT8QFkOwH4qOS2SJkS4wvpGl8BpA=="], + "@libsql/isomorphic-ws/ws": ["ws@8.18.0", "", { "peerDependencies": { "bufferutil": "^4.0.1", "utf-8-validate": ">=5.0.2" }, "optionalPeers": ["bufferutil", "utf-8-validate"] }, "sha512-8VbfWfHLbbwu3+N6OKsOMpBdT4kXPDDB9cJk2bJ6mh9ucxdlnNvH1e+roYkKmN9Nxw2yjz7VzeO9oOz2zJ04Pw=="], + "@esbuild-kit/core-utils/esbuild/@esbuild/android-arm": ["@esbuild/android-arm@0.18.20", "", { "os": "android", "cpu": "arm" }, "sha512-fyi7TDI/ijKKNZTUJAQqiG5T7YjJXgnzkURqmGj13C6dCqckZBLdl4h7bkhHt/t0WP+zO9/zwroDvANaOqO5Sw=="], "@esbuild-kit/core-utils/esbuild/@esbuild/android-arm64": ["@esbuild/android-arm64@0.18.20", "", { "os": "android", "cpu": "arm64" }, "sha512-Nz4rJcchGDtENV0eMKUNa6L12zz2zBDXuhj/Vjh18zGqB44Bi7MBMSXjgunJgjRhCmKOjnPuZp4Mb6OKqtMHLQ=="], diff --git a/package.json b/package.json index cd6c0b5..78b97b1 100644 --- a/package.json +++ b/package.json @@ -20,11 +20,13 @@ "@libsql/client": "^0.15.15", "drizzle-orm": "^0.45.1", "hono": "^4.12.2", + "ws": "8.21.0", "zod": "^4.3.6" }, "devDependencies": { "@biomejs/biome": "2.4.0", "@types/bun": "^1.3.9", + "@types/ws": "8.18.1", "@typescript/native-preview": "^7.0.0-dev.20260223.1", "drizzle-kit": "^0.31.9", "lefthook": "^2.1.1", diff --git a/src/http/routes/proxy.ts b/src/http/routes/proxy.ts index 9a2fcf6..0be351e 100644 --- a/src/http/routes/proxy.ts +++ b/src/http/routes/proxy.ts @@ -350,6 +350,7 @@ const proxyRequest = async ( body: requestBody, // Provider streams can pause for minutes while a model is thinking. timeout: false, + signal: context.req.raw.signal, }; if (headerTimeout) { upstreamRequestInit.signal = AbortSignal.any([ @@ -368,11 +369,15 @@ const proxyRequest = async ( headerTimeout?.clear(); } } catch (error) { + if (context.req.raw.signal.aborted) { + // Client disconnects are expected cancellations, not proxy failures. + throw error; + } logWarn("proxy_upstream_request_failed", { provider: route.provider, endpoint: route.endpoint, elapsedMs: Date.now() - startedAt, - aborted: context.req.raw.signal.aborted, + aborted: false, ...errorLogFields(error), }); usageRecorder.recordImmediate(500); diff --git a/src/providers/proxies/claude-proxy.ts b/src/providers/proxies/claude-proxy.ts index 6b00f1c..80152c9 100644 --- a/src/providers/proxies/claude-proxy.ts +++ b/src/providers/proxies/claude-proxy.ts @@ -15,7 +15,7 @@ import { } from "../../usage/token-usage"; import { errorLogFields, logWarn } from "../../utils/log"; import { isObjectRecord, type JsonObject } from "../../utils/object"; -import { createSseKeepAlive } from "./sse-keepalive"; +import { createSseKeepAlive, createSseResponseHeaders } from "./sse-keepalive"; // Anthropic OAuth sessions reject the feedback repo path used in OpenCode's // prompt URL and the opening `` wrapper emitted by OpenCode's @@ -332,6 +332,7 @@ const maybeTransformClaudeStreamResponse = ( let bytes = 0; let chunks = 0; let lastChunkAt = startedAt; + let lastWriteAt = startedAt; let closed = false; let clearKeepAlive: (() => void) | null = null; @@ -345,6 +346,7 @@ const maybeTransformClaudeStreamResponse = ( transport: "sse_transform", elapsedMs: Date.now() - startedAt, idleMs: Date.now() - lastChunkAt, + downstreamIdleMs: Date.now() - lastWriteAt, bytes, chunks, ...fields, @@ -457,87 +459,84 @@ const maybeTransformClaudeStreamResponse = ( const stream = new ReadableStream({ start(controller): void { - const keepAlive = createSseKeepAlive(controller, { + clearKeepAlive = createSseKeepAlive(controller, { provider: "claude", transport: "sse_transform", getElapsedMs: () => Date.now() - startedAt, - }); - clearKeepAlive = keepAlive.clear; - - const pump = async (): Promise => { - try { - while (true) { - const { done, value } = await reader.read(); - if (done) { - if (closed) { - clearKeepAlive?.(); - return; - } - buffer += decoder.decode(); - if (buffer) { - controller.enqueue( - encoder.encode( - transformSseEventChunk( - buffer, - toolPrefix, - readStreamUsage, - readStreamAnomaly - ) - ) - ); - buffer = ""; - } - onTokenUsage?.(streamUsage); - closed = true; + onKeepAlive: () => { + lastWriteAt = Date.now(); + }, + }).clear; + }, + async pull(controller): Promise { + try { + while (true) { + const { done, value } = await reader.read(); + if (done) { + if (closed) { clearKeepAlive?.(); - controller.close(); return; } - - if (!value) { - continue; + buffer += decoder.decode(); + if (buffer) { + controller.enqueue( + encoder.encode( + transformSseEventChunk( + buffer, + toolPrefix, + readStreamUsage, + readStreamAnomaly + ) + ) + ); + lastWriteAt = Date.now(); + buffer = ""; } + onTokenUsage?.(streamUsage); + closed = true; + clearKeepAlive?.(); + controller.close(); + return; + } - bytes += value.byteLength; - chunks++; - lastChunkAt = Date.now(); - buffer += decoder.decode(value, { stream: true }); - - let boundary = findSseEventBoundary(buffer); - while (boundary) { - const chunk = buffer.slice(0, boundary.index + boundary.length); - buffer = buffer.slice(boundary.index + boundary.length); - try { - controller.enqueue( - encoder.encode( - transformSseEventChunk( - chunk, - toolPrefix, - readStreamUsage, - readStreamAnomaly - ) + if (!value) { + continue; + } + + bytes += value.byteLength; + chunks++; + lastChunkAt = Date.now(); + buffer += decoder.decode(value, { stream: true }); + + let enqueued = false; + let boundary = findSseEventBoundary(buffer); + while (boundary) { + const chunk = buffer.slice(0, boundary.index + boundary.length); + buffer = buffer.slice(boundary.index + boundary.length); + try { + controller.enqueue( + encoder.encode( + transformSseEventChunk( + chunk, + toolPrefix, + readStreamUsage, + readStreamAnomaly ) - ); - } catch (error) { - logStreamAnomaly("claude_sse_enqueue_failed", {}, error); - throw error; - } - boundary = findSseEventBoundary(buffer); + ) + ); + lastWriteAt = Date.now(); + enqueued = true; + } catch (error) { + logStreamAnomaly("claude_sse_enqueue_failed", {}, error); + throw error; } + boundary = findSseEventBoundary(buffer); } - } catch (error) { - if (closed) { - clearKeepAlive?.(); + if (enqueued) { return; } - closed = true; - clearKeepAlive?.(); - logStreamAnomaly("claude_sse_stream_failed", {}, error); - controller.error(error); } - }; - - pump().catch((error: unknown) => { + } catch (error) { if (closed) { clearKeepAlive?.(); return; @@ -546,12 +545,9 @@ const maybeTransformClaudeStreamResponse = ( clearKeepAlive?.(); logStreamAnomaly("claude_sse_stream_failed", {}, error); controller.error(error); - }); + } }, cancel(reason): Promise { - if (!closed) { - logStreamAnomaly("claude_sse_downstream_cancelled", {}, reason); - } closed = true; clearKeepAlive?.(); return reader.cancel(reason); @@ -561,7 +557,7 @@ const maybeTransformClaudeStreamResponse = ( return new Response(stream, { status: response.status, statusText: response.statusText, - headers: response.headers, + headers: createSseResponseHeaders(response.headers), }); }; diff --git a/src/providers/proxies/codex-websocket.ts b/src/providers/proxies/codex-websocket.ts index 572019e..be04646 100644 --- a/src/providers/proxies/codex-websocket.ts +++ b/src/providers/proxies/codex-websocket.ts @@ -1,3 +1,5 @@ +import WebSocket from "ws"; + import { CODEX_RESPONSE_ENDPOINT, CODEX_WEBSOCKET_BETA_HEADER, @@ -11,6 +13,7 @@ import { readOpenAiResponsesUsageFromSseEvent } from "../../usage/token-usage"; import type { TokenUsage } from "../../usage/token-usage"; import { errorLogFields, logWarn } from "../../utils/log"; import { isObjectRecord, readBooleanField } from "../../utils/object"; +import { createSseKeepAlive, createSseResponseHeaders } from "./sse-keepalive"; const SESSION_SOCKET_TTL_MS = 5 * 60 * 1000; const CONNECT_TIMEOUT_MS = 15_000; @@ -18,8 +21,10 @@ const RESPONSE_IDLE_TIMEOUT_MS = 5 * 60 * 1000; const MAX_SOCKET_AGE_MS = 55 * 60 * 1000; const CONNECTION_LIMIT_RETRIES = 5; const STREAM_FAILURE_RETRIES = 5; +const MAX_TRACKED_FAILURE_SESSIONS = 1000; const CONNECTION_LIMIT_REACHED_CODE = "websocket_connection_limit_reached"; const WEBSOCKET_MESSAGE_TOO_BIG_CLOSE_CODE = 1009; +const textEncoder = new TextEncoder(); type WebSocketEventType = "open" | "message" | "error" | "close"; type WebSocketListener = (event: unknown) => void; @@ -27,7 +32,8 @@ type WebSocketListener = (event: unknown) => void; type WebSocketLike = { readonly readyState?: number; close(code?: number, reason?: string): void; - send(data: string): void; + terminate?: () => void; + send(data: string, callback?: (error?: Error) => void): void; addEventListener(type: WebSocketEventType, listener: WebSocketListener): void; removeEventListener( type: WebSocketEventType, @@ -37,9 +43,17 @@ type WebSocketLike = { type WebSocketConstructor = new ( url: string, - protocols?: string | string[] | { headers?: Record } + options?: { headers?: Record } ) => WebSocketLike; +let webSocketConstructorOverride: WebSocketConstructor | null = null; + +export const setCodexWebSocketConstructorForTests = ( + webSocketConstructor: WebSocketConstructor | null +): void => { + webSocketConstructorOverride = webSocketConstructor; +}; + type WebSocketCloseError = Error & { webSocketCloseCode?: number; }; @@ -91,9 +105,7 @@ const headersToRecord = (headers: Headers): Record => { const createSseResponse = (body: ReadableStream): Response => new Response(body, { - headers: { - "content-type": "text/event-stream", - }, + headers: createSseResponseHeaders(), }); const readString = (value: unknown): string | null => { @@ -105,6 +117,31 @@ const readString = (value: unknown): string | null => { return trimmed || null; }; +const isCompactionRequest = ( + body: Record, + headers: Headers +): boolean => { + const clientMetadata = isObjectRecord(body.client_metadata) + ? body.client_metadata + : null; + const rawMetadata = + readString(clientMetadata?.["x-codex-turn-metadata"]) ?? + readString(headers.get("x-codex-turn-metadata")); + if (!rawMetadata) { + return false; + } + + try { + const metadata = JSON.parse(rawMetadata) as unknown; + return ( + isObjectRecord(metadata) && + readString(metadata.request_kind)?.toLowerCase() === "compaction" + ); + } catch { + return false; + } +}; + const isSocketOpen = (socket: WebSocketLike): boolean => socket.readyState === undefined || socket.readyState === 1; @@ -119,6 +156,18 @@ const closeSocket = (socket: WebSocketLike): void => { } }; +const terminateSocket = (socket: WebSocketLike): void => { + try { + if (socket.terminate) { + socket.terminate(); + return; + } + socket.close(1001, "invalid"); + } catch { + // Ignore termination failures from already-closed sockets. + } +}; + const extractWebSocketError = (event: unknown): Error => { if (isObjectRecord(event)) { const message = readString(event.message); @@ -174,8 +223,8 @@ const clearSessionFallback = (key: string): void => { const timer = fallbackSocketKeys.get(key); if (timer) { clearTimeout(timer); - fallbackSocketKeys.delete(key); } + fallbackSocketKeys.delete(key); }; const markSessionFallback = (key: string | null): void => { @@ -202,6 +251,15 @@ const recordSessionStreamFailure = (key: string | null): number => { return 0; } + if ( + !streamFailureCounts.has(key) && + streamFailureCounts.size >= MAX_TRACKED_FAILURE_SESSIONS + ) { + const oldestKey = streamFailureCounts.keys().next().value; + if (typeof oldestKey === "string") { + streamFailureCounts.delete(oldestKey); + } + } const failures = (streamFailureCounts.get(key) ?? 0) + 1; streamFailureCounts.set(key, failures); if (failures > STREAM_FAILURE_RETRIES) { @@ -222,16 +280,15 @@ const scheduleExpiry = (key: string, cached: CachedSocket): void => { closeSocket(cached.socket); socketCache.delete(key); - clearSessionFallback(key); clearSessionStreamFailures(key); }, SESSION_SOCKET_TTL_MS); }; const getWebSocketConstructor = (): WebSocketConstructor | null => { - const websocket = globalThis.WebSocket; - return typeof websocket === "function" - ? (websocket as unknown as WebSocketConstructor) - : null; + if (webSocketConstructorOverride) { + return webSocketConstructorOverride; + } + return WebSocket as unknown as WebSocketConstructor; }; const connectWebSocket = ( @@ -282,14 +339,14 @@ const connectWebSocket = ( const onClose = (event: unknown): void => fail(extractWebSocketCloseError(event)); const onAbort = (): void => { - closeSocket(socket); fail(new Error("Request was aborted")); + terminateSocket(socket); }; const onTimeout = (): void => { - closeSocket(socket); fail( new Error(`WebSocket connect timeout after ${CONNECT_TIMEOUT_MS}ms`) ); + terminateSocket(socket); }; try { @@ -323,7 +380,13 @@ const acquireSocket = async ( const acquired = { socket, cached: null, - release: (): void => closeSocket(acquired.socket), + release(keep: boolean): void { + if (keep) { + closeSocket(acquired.socket); + return; + } + terminateSocket(acquired.socket); + }, }; return acquired; } @@ -348,7 +411,7 @@ const acquireSocket = async ( cached: existing, release(keep: boolean): void { if (!(keep && isSocketOpen(existing.socket))) { - closeSocket(existing.socket); + terminateSocket(existing.socket); socketCache.delete(cacheKey); return; } @@ -391,7 +454,7 @@ const acquireSocket = async ( cached, release(keep: boolean): void { if (!(keep && isSocketOpen(cached.socket))) { - closeSocket(cached.socket); + terminateSocket(cached.socket); socketCache.delete(cacheKey); return; } @@ -632,10 +695,9 @@ const buildRequestBody = ( }; const encodeSse = (payload: unknown): Uint8Array => - new TextEncoder().encode(`data: ${JSON.stringify(payload)}\n\n`); + textEncoder.encode(`data: ${JSON.stringify(payload)}\n\n`); -const encodeDoneSse = (): Uint8Array => - new TextEncoder().encode("data: [DONE]\n\n"); +const encodeDoneSse = (): Uint8Array => textEncoder.encode("data: [DONE]\n\n"); const decodeMessageData = async (data: unknown): Promise => { if (typeof data === "string") { @@ -793,6 +855,11 @@ export const tryProxyCodexWebSocket = async ( const cacheKey = sessionId ? `${input.accountKey}:${sessionId}` : null; const headers = buildWebSocketHeaders(input.headers, requestId); + if (isCompactionRequest(body, input.headers)) { + markSessionFallback(cacheKey); + return null; + } + if (cacheKey && fallbackSocketKeys.has(cacheKey)) { return null; } @@ -802,9 +869,16 @@ export const tryProxyCodexWebSocket = async ( sessionId ? { ...body, prompt_cache_key: requestId } : body ); let requestBody: Record; + let requestPayloadText: string; + let requestBytes: number; try { acquired = await acquireSocket(headers, cacheKey, input.signal); requestBody = buildRequestBody(fullBody, acquired.cached); + requestPayloadText = JSON.stringify({ + ...requestBody, + type: "response.create", + }); + requestBytes = textEncoder.encode(requestPayloadText).byteLength; } catch (error) { acquired?.release(false); if (!input.signal?.aborted && !isSessionConcurrencyError(error)) { @@ -828,6 +902,11 @@ export const tryProxyCodexWebSocket = async ( let connectionLimitAttempts = 0; let retryingConnectionLimit = false; let responseIdleTimer: ReturnType | null = null; + let clearStreamKeepAlive: (() => void) | null = null; + let payloadCount = 0; + let lastUpstreamEventAt = startedAt; + let lastDownstreamWriteAt = startedAt; + let socketEpoch = 0; let messageChain = Promise.resolve(); let wake: (() => void) | null = null; const queue: Record[] = []; @@ -842,6 +921,10 @@ export const tryProxyCodexWebSocket = async ( terminal, queueLength: queue.length, connectionLimitAttempts, + requestBytes, + payloadCount, + upstreamIdleMs: Date.now() - lastUpstreamEventAt, + downstreamIdleMs: Date.now() - lastDownstreamWriteAt, ...fields, }); }; @@ -906,10 +989,20 @@ export const tryProxyCodexWebSocket = async ( return; } try { - active.socket.send( - JSON.stringify({ ...requestBody, type: "response.create" }) - ); - resetResponseIdleTimer("idle_timeout_waiting_for_websocket"); + resetResponseIdleTimer("idle_timeout_sending_websocket_request"); + const sendingSocket = active.socket; + sendingSocket.send(requestPayloadText, (error?: Error) => { + // Terminated sockets flush pending send callbacks with an error; + // ignore callbacks from sockets replaced by a connection retry. + if (settled || sendingSocket !== active.socket) { + return; + } + if (error) { + fail("send_failed", error); + return; + } + resetResponseIdleTimer("idle_timeout_waiting_for_websocket"); + }); } catch (error) { fail("send_failed", error); } @@ -929,14 +1022,23 @@ export const tryProxyCodexWebSocket = async ( connectionLimitAttempts++; retryingConnectionLimit = true; + // Invalidate events already dispatched (or still queued) from the socket + // being replaced; they must not fail the retried stream. + socketEpoch++; clearResponseIdleTimer(); active.socket.removeEventListener("message", onMessage); active.socket.removeEventListener("error", onError); active.socket.removeEventListener("close", onClose); - closeSocket(active.socket); + terminateSocket(active.socket); try { const nextSocket = await connectWebSocket(headers, input.signal); + if (settled) { + // A downstream cancel or abort raced the retry; do not leak the + // freshly opened socket. + terminateSocket(nextSocket); + return; + } active.socket = nextSocket; if (active.cached) { active.cached.socket = nextSocket; @@ -958,14 +1060,10 @@ export const tryProxyCodexWebSocket = async ( settled = true; keepSocket = false; failure = error; + clearStreamKeepAlive?.(); cleanup(); active.release(false); - if (isUserCancelledStage(stage)) { - logStreamAnomaly("codex_websocket_stream_cancelled", { - stage, - ...errorLogFields(error), - }); - } else { + if (!isUserCancelledStage(stage)) { const webSocketCloseCode = readWebSocketCloseCode(error); const immediateFallback = stage === "socket_closed_before_terminal" && @@ -991,6 +1089,7 @@ export const tryProxyCodexWebSocket = async ( return; } settled = true; + clearStreamKeepAlive?.(); cleanup(); const canStoreContinuation = Boolean( active.cached && @@ -1021,13 +1120,30 @@ export const tryProxyCodexWebSocket = async ( const onAbort = (): void => fail("request_aborted", new Error("Request was aborted")); - const onError = (event: unknown): void => - fail("socket_error", extractWebSocketError(event)); + const onError = (event: unknown): void => { + const eventEpoch = socketEpoch; + // Serialize behind in-flight message handling so an error event cannot + // clobber a terminal payload that is still being processed. + messageChain = messageChain + .then(() => { + if (terminal || eventEpoch !== socketEpoch) { + wakePull(); + return; + } + fail("socket_error", extractWebSocketError(event)); + }) + .catch((error: unknown) => { + fail("message_parse_failed", error); + }); + }; const onClose = (event: unknown): void => { + const eventEpoch = socketEpoch; messageChain = messageChain .then(() => { - if (terminal) { + if (terminal || eventEpoch !== socketEpoch) { + // A close from a socket replaced by a connection-limit retry must + // not fail the stream or mark session fallback. wakePull(); return; } @@ -1057,6 +1173,8 @@ export const tryProxyCodexWebSocket = async ( return; } resetResponseIdleTimer("idle_timeout_waiting_for_websocket"); + payloadCount++; + lastUpstreamEventAt = Date.now(); if (!emittedPayload && isConnectionLimitPayload(payload)) { retryConnectionLimit().catch((error: unknown) => { @@ -1130,6 +1248,7 @@ export const tryProxyCodexWebSocket = async ( } controller.enqueue(encodeSse(payload)); + lastDownstreamWriteAt = Date.now(); if (isTerminalPayload(payload)) { terminal = true; finalEventType = payload.type; @@ -1152,11 +1271,22 @@ export const tryProxyCodexWebSocket = async ( } } controller.enqueue(encodeDoneSse()); + lastDownstreamWriteAt = Date.now(); finish(); } }; const stream = new ReadableStream({ + start(controller): void { + clearStreamKeepAlive = createSseKeepAlive(controller, { + provider: "codex", + transport: "websocket_sse", + getElapsedMs: () => Date.now() - startedAt, + onKeepAlive: () => { + lastDownstreamWriteAt = Date.now(); + }, + }).clear; + }, async pull(controller): Promise { while (!queue.length && !settled) { await new Promise((resolve) => { @@ -1185,6 +1315,7 @@ export const tryProxyCodexWebSocket = async ( controller.close(); }, cancel(): void { + clearStreamKeepAlive?.(); if (settled) { return; } diff --git a/src/providers/proxies/openai-sse-passthrough.ts b/src/providers/proxies/openai-sse-passthrough.ts index 0b1f3db..443debf 100644 --- a/src/providers/proxies/openai-sse-passthrough.ts +++ b/src/providers/proxies/openai-sse-passthrough.ts @@ -1,7 +1,7 @@ import type { TokenUsage } from "../../usage/token-usage"; import { errorLogFields, logWarn } from "../../utils/log"; import { isObjectRecord } from "../../utils/object"; -import { createSseKeepAlive } from "./sse-keepalive"; +import { createSseKeepAlive, createSseResponseHeaders } from "./sse-keepalive"; type SseUsageExtractor = (payload: unknown) => TokenUsage | null; @@ -9,6 +9,7 @@ type OpenAiSsePassthroughInput = { response: Response; extractUsage: SseUsageExtractor; onTokenUsage?: ((usage: TokenUsage) => void) | null | undefined; + keepAliveIntervalMs?: number; }; const readSseTerminalAnomaly = (payload: unknown): string | null => { @@ -104,6 +105,12 @@ export const createOpenAiSseUsagePassthrough = ( const reader = input.response.body.getReader(); const decoder = new TextDecoder(); const startedAt = Date.now(); + const contentType = ( + input.response.headers.get("content-type") ?? "" + ).toLowerCase(); + // Codex omits the content-type header on SSE streams; anything explicitly + // non-SSE (e.g. JSON error bodies) must not receive keepalive comments. + const isSseBody = !contentType || contentType.includes("text/event-stream"); const usageState = { eventDataLines: [] as string[], latestUsage: null as TokenUsage | null, @@ -113,6 +120,7 @@ export const createOpenAiSseUsagePassthrough = ( let bytes = 0; let chunks = 0; let lastChunkAt = startedAt; + let lastWriteAt = startedAt; let closed = false; let clearKeepAlive: (() => void) | null = null; @@ -126,6 +134,7 @@ export const createOpenAiSseUsagePassthrough = ( transport: "sse", elapsedMs: Date.now() - startedAt, idleMs: Date.now() - lastChunkAt, + downstreamIdleMs: Date.now() - lastWriteAt, bytes, chunks, ...fields, @@ -135,75 +144,71 @@ export const createOpenAiSseUsagePassthrough = ( const stream = new ReadableStream({ start(controller): void { - const keepAlive = createSseKeepAlive(controller, { + if (!isSseBody) { + return; + } + clearKeepAlive = createSseKeepAlive(controller, { provider: "openai", transport: "sse", getElapsedMs: () => Date.now() - startedAt, - }); - clearKeepAlive = keepAlive.clear; + onKeepAlive: () => { + lastWriteAt = Date.now(); + }, + ...(input.keepAliveIntervalMs + ? { intervalMs: input.keepAliveIntervalMs } + : {}), + }).clear; + }, + async pull(controller): Promise { + try { + let result = await reader.read(); + while (!(result.done || result.value)) { + result = await reader.read(); + } - const pump = async (): Promise => { - try { - while (true) { - const { done, value } = await reader.read(); - if (done) { - if (closed) { - clearKeepAlive?.(); - return; - } - pendingText += decoder.decode(); - pendingText = readLatestUsageFromSse( - `${pendingText}\n\n`, - usageState, - input.extractUsage - ); - if (usageState.latestUsage) { - input.onTokenUsage?.(usageState.latestUsage); - } - if (usageState.terminalAnomaly) { - logStreamAnomaly("openai_sse_terminal_anomaly", { - terminalAnomaly: usageState.terminalAnomaly, - }); - } - closed = true; - clearKeepAlive?.(); - controller.close(); - return; - } - - if (!value) { - continue; - } - - bytes += value.byteLength; - chunks++; - lastChunkAt = Date.now(); - pendingText += decoder.decode(value, { stream: true }); - pendingText = readLatestUsageFromSse( - pendingText, - usageState, - input.extractUsage - ); - try { - controller.enqueue(value); - } catch (error) { - logStreamAnomaly("openai_sse_enqueue_failed", {}, error); - throw error; - } - } - } catch (error) { + if (result.done) { if (closed) { clearKeepAlive?.(); return; } + pendingText += decoder.decode(); + pendingText = readLatestUsageFromSse( + `${pendingText}\n\n`, + usageState, + input.extractUsage + ); + if (usageState.latestUsage) { + input.onTokenUsage?.(usageState.latestUsage); + } + if (usageState.terminalAnomaly) { + logStreamAnomaly("openai_sse_terminal_anomaly", { + terminalAnomaly: usageState.terminalAnomaly, + }); + } closed = true; clearKeepAlive?.(); - logStreamAnomaly("openai_sse_stream_failed", {}, error); - controller.error(error); + controller.close(); + return; } - }; - pump().catch((error: unknown) => { + const value = result.value; + bytes += value.byteLength; + chunks++; + lastChunkAt = Date.now(); + pendingText += decoder.decode(value, { stream: true }); + pendingText = readLatestUsageFromSse( + pendingText, + usageState, + input.extractUsage + ); + try { + controller.enqueue(value); + lastWriteAt = Date.now(); + } catch (error) { + logStreamAnomaly("openai_sse_enqueue_failed", {}, error); + throw error; + } + } catch (error) { if (closed) { clearKeepAlive?.(); return; @@ -212,12 +217,9 @@ export const createOpenAiSseUsagePassthrough = ( clearKeepAlive?.(); logStreamAnomaly("openai_sse_stream_failed", {}, error); controller.error(error); - }); + } }, cancel(reason): Promise { - if (!closed) { - logStreamAnomaly("openai_sse_downstream_cancelled", {}, reason); - } closed = true; clearKeepAlive?.(); return reader.cancel(reason); @@ -227,6 +229,6 @@ export const createOpenAiSseUsagePassthrough = ( return new Response(stream, { status: input.response.status, statusText: input.response.statusText, - headers: input.response.headers, + headers: createSseResponseHeaders(input.response.headers), }); }; diff --git a/src/providers/proxies/sse-keepalive.ts b/src/providers/proxies/sse-keepalive.ts index c77d6b1..6a6c40a 100644 --- a/src/providers/proxies/sse-keepalive.ts +++ b/src/providers/proxies/sse-keepalive.ts @@ -3,10 +3,24 @@ import { errorLogFields, logWarn } from "../../utils/log"; const SSE_KEEPALIVE_INTERVAL_MS = 25_000; const SSE_KEEPALIVE_BYTES = new TextEncoder().encode(": kleis-keepalive\n\n"); +export const createSseResponseHeaders = (source?: HeadersInit): Headers => { + const headers = new Headers(source); + headers.delete("content-encoding"); + headers.delete("content-length"); + headers.set("cache-control", "no-cache, no-transform"); + if (!headers.has("content-type")) { + headers.set("content-type", "text/event-stream"); + } + headers.set("x-accel-buffering", "no"); + return headers; +}; + type SseKeepAliveInput = { provider: string; transport: string; getElapsedMs: () => number; + onKeepAlive?: () => void; + intervalMs?: number; }; export const createSseKeepAlive = ( @@ -21,6 +35,7 @@ export const createSseKeepAlive = ( try { controller.enqueue(SSE_KEEPALIVE_BYTES); + input.onKeepAlive?.(); } catch (error) { active = false; clearInterval(timer); @@ -31,7 +46,7 @@ export const createSseKeepAlive = ( ...errorLogFields(error), }); } - }, SSE_KEEPALIVE_INTERVAL_MS); + }, input.intervalMs ?? SSE_KEEPALIVE_INTERVAL_MS); return { clear(): void { diff --git a/tests/providers/proxy-contract.test.ts b/tests/providers/proxy-contract.test.ts index bf7d73a..b717796 100644 --- a/tests/providers/proxy-contract.test.ts +++ b/tests/providers/proxy-contract.test.ts @@ -19,16 +19,16 @@ import { prepareClaudeProxyRequest } from "../../src/providers/proxies/claude-pr import { prepareCodexProxyRequest } from "../../src/providers/proxies/codex-proxy"; import { closeCodexWebSocketSessions, + setCodexWebSocketConstructorForTests, tryProxyCodexWebSocket, } from "../../src/providers/proxies/codex-websocket"; import { prepareCopilotProxyRequest } from "../../src/providers/proxies/copilot-proxy"; +import { createOpenAiSseUsagePassthrough } from "../../src/providers/proxies/openai-sse-passthrough"; import type { TokenUsage } from "../../src/usage/token-usage"; -const originalWebSocket = globalThis.WebSocket; - afterEach(() => { closeCodexWebSocketSessions(); - globalThis.WebSocket = originalWebSocket; + setCodexWebSocketConstructorForTests(null); }); const createUsageCapture = () => { @@ -80,6 +80,10 @@ type MockCodexWebSocketResponse = { terminalType?: "response.completed" | "response.done" | "response.incomplete"; }; +type MockWebSocketOptions = { + headers?: Record; +}; + const installCodexWebSocketMock = ( responses: MockCodexWebSocketResponse[], sentBodies: unknown[], @@ -93,17 +97,9 @@ const installCodexWebSocketMock = ( Set<(event: unknown) => void> >(); - constructor( - _url: string, - protocols?: string | string[] | { headers?: Record } - ) { - if ( - protocols && - typeof protocols === "object" && - !Array.isArray(protocols) && - protocols.headers - ) { - constructorHeaders.push(protocols.headers); + constructor(_url: string, options?: MockWebSocketOptions) { + if (options?.headers) { + constructorHeaders.push(options.headers); } queueMicrotask(() => this.dispatch("open", {})); } @@ -121,7 +117,7 @@ const installCodexWebSocketMock = ( this.listeners.get(type)?.delete(listener); } - send(data: string): void { + send(data: string, callback?: (error?: Error) => void): void { sentBodies.push(JSON.parse(data) as unknown); const response = responses.shift(); if (!response) { @@ -129,6 +125,7 @@ const installCodexWebSocketMock = ( } queueMicrotask(() => { + callback?.(); for (const event of [ { type: "response.created", response: { id: response.id } }, ...(response.items ?? []).map((item) => ({ @@ -160,6 +157,10 @@ const installCodexWebSocketMock = ( this.readyState = 3; } + terminate(): void { + this.readyState = 3; + } + private dispatch(type: string, event: unknown): void { for (const listener of this.listeners.get(type) ?? []) { listener(event); @@ -167,11 +168,12 @@ const installCodexWebSocketMock = ( } } - globalThis.WebSocket = MockWebSocket as unknown as typeof WebSocket; + setCodexWebSocketConstructorForTests(MockWebSocket); }; type ManualCodexWebSocket = { dispatch(type: string, event: unknown): void; + readonly terminated: boolean; }; const installManualCodexWebSocketMock = ( @@ -183,6 +185,7 @@ const installManualCodexWebSocketMock = ( class MockWebSocket { static OPEN = 1; readyState = MockWebSocket.OPEN; + terminated = false; private readonly listeners = new Map< string, Set<(event: unknown) => void> @@ -208,14 +211,20 @@ const installManualCodexWebSocketMock = ( this.listeners.get(type)?.delete(listener); } - send(data: string): void { + send(data: string, callback?: (error?: Error) => void): void { sentBodies.push(JSON.parse(data) as unknown); + callback?.(); } close(): void { this.readyState = 3; } + terminate(): void { + this.terminated = true; + this.readyState = 3; + } + dispatch(type: string, event: unknown): void { for (const listener of this.listeners.get(type) ?? []) { listener(event); @@ -223,7 +232,7 @@ const installManualCodexWebSocketMock = ( } } - globalThis.WebSocket = MockWebSocket as unknown as typeof WebSocket; + setCodexWebSocketConstructorForTests(MockWebSocket); return sockets; }; @@ -585,6 +594,163 @@ describe("proxy contract: codex", () => { }); }); + test("applies downstream backpressure to OpenAI SSE reads", async () => { + const encoder = new TextEncoder(); + let pulls = 0; + const source = new Response( + new ReadableStream( + { + pull(controller): void { + pulls++; + controller.enqueue( + encoder.encode( + `data: ${JSON.stringify({ type: "response.output_text.delta", delta: pulls })}\n\n` + ) + ); + }, + }, + { highWaterMark: 0 } + ), + { headers: { "content-type": "text/event-stream" } } + ); + + const transformed = createOpenAiSseUsagePassthrough({ + response: source, + extractUsage: () => null, + }); + await new Promise((resolve) => setTimeout(resolve, 0)); + + expect(pulls).toBe(1); + await transformed.body?.cancel(); + }); + + test("preserves explicit JSON content type for streaming errors", async () => { + const body = JSON.stringify({ + error: { message: "bad request", type: "invalid_request_error" }, + }); + const response = createOpenAiSseUsagePassthrough({ + response: new Response(body, { + status: 400, + headers: { + "content-encoding": "gzip", + "content-length": String(body.length), + "content-type": "application/json", + }, + }), + extractUsage: () => null, + }); + + expect(response.status).toBe(400); + expect(response.headers.get("content-type")).toBe("application/json"); + expect(response.headers.has("content-length")).toBe(false); + expect(response.headers.has("content-encoding")).toBe(false); + expect(await response.json()).toEqual({ + error: { message: "bad request", type: "invalid_request_error" }, + }); + }); + + test("does not inject keepalives into non-SSE bodies", async () => { + const encoder = new TextEncoder(); + const createSlowResponse = ( + payload: string, + contentType: string | null + ): Response => + new Response( + new ReadableStream({ + async start(controller): Promise { + controller.enqueue(encoder.encode(payload.slice(0, 4))); + await new Promise((resolve) => setTimeout(resolve, 30)); + controller.enqueue(encoder.encode(payload.slice(4))); + controller.close(); + }, + }), + contentType ? { headers: { "content-type": contentType } } : {} + ); + + const jsonBody = JSON.stringify({ error: { message: "slow error" } }); + const jsonResponse = createOpenAiSseUsagePassthrough({ + response: createSlowResponse(jsonBody, "application/json"), + extractUsage: () => null, + keepAliveIntervalMs: 5, + }); + expect(await jsonResponse.text()).toBe(jsonBody); + + const sseBody = 'data: {"type":"response.completed"}\n\n'; + const sseResponse = createOpenAiSseUsagePassthrough({ + response: createSlowResponse(sseBody, null), + extractUsage: () => null, + keepAliveIntervalMs: 5, + }); + const sseText = await sseResponse.text(); + expect(sseText).toContain("response.completed"); + expect(sseText).toContain(": kleis-keepalive"); + }); + + test("routes compaction turns over HTTP instead of WebSocket", async () => { + const sentBodies: unknown[] = []; + const sockets = installManualCodexWebSocketMock(sentBodies); + const headers = new Headers({ + authorization: "Bearer codex-access", + [CODEX_ACCOUNT_ID_HEADER]: "acct_1", + "x-session-affinity": "compaction-session", + }); + + const response = await tryProxyCodexWebSocket({ + headers, + bodyJson: { + model: "gpt-5.5", + stream: true, + client_metadata: { + "x-codex-turn-metadata": JSON.stringify({ + request_kind: "compaction", + }), + }, + input: [{ role: "user", content: "compact" }], + }, + accountKey: "key-1:account-1", + }); + + expect(response).toBeNull(); + expect(sockets).toHaveLength(0); + expect(sentBodies).toHaveLength(0); + }); + + test("terminates websocket when the downstream request aborts", async () => { + const sentBodies: unknown[] = []; + const sockets = installManualCodexWebSocketMock(sentBodies); + const abortController = new AbortController(); + const warnings: string[] = []; + const originalWarn = console.warn; + console.warn = (message?: unknown): void => { + warnings.push(String(message)); + }; + + try { + const responsePromise = tryProxyCodexWebSocket({ + headers: new Headers({ + authorization: "Bearer codex-access", + [CODEX_ACCOUNT_ID_HEADER]: "acct_1", + "x-session-affinity": "abort-session", + }), + bodyJson: { + model: "gpt-5.5", + stream: true, + input: [{ role: "user", content: "hello" }], + }, + accountKey: "key-1:account-1", + signal: abortController.signal, + }); + await waitFor(() => sentBodies.length === 1); + abortController.abort(new Error("stop")); + expect(await responsePromise).toBeNull(); + } finally { + console.warn = originalWarn; + } + + expect(sockets[0]?.terminated).toBe(true); + expect(warnings).toHaveLength(0); + }); + test("uses websocket cached delta transport for streaming requests", async () => { const firstAssistantItem = { type: "message", @@ -641,6 +807,9 @@ describe("proxy contract: codex", () => { onTokenUsage: capture.onTokenUsage, }); expect(first).not.toBeNull(); + expect(first?.headers.get("cache-control")).toBe("no-cache, no-transform"); + expect(first?.headers.get("x-accel-buffering")).toBe("no"); + expect(first?.headers.get("content-length")).toBeNull(); await first?.text(); const second = await tryProxyCodexWebSocket({ @@ -1361,6 +1530,68 @@ describe("proxy contract: codex", () => { expect(sentBodies).toHaveLength(2); }); + test("survives a socket close racing a connection limit retry", async () => { + const sentBodies: unknown[] = []; + const sockets = installManualCodexWebSocketMock(sentBodies); + const headers = new Headers({ + authorization: "Bearer codex-access", + [CODEX_ACCOUNT_ID_HEADER]: "acct_1", + "x-session-affinity": "retry-close-race", + }); + const bodyJson = { + model: "gpt-5-codex", + stream: true, + input: [ + { role: "user", content: [{ type: "input_text", text: "Race" }] }, + ], + }; + + const responsePromise = tryProxyCodexWebSocket({ + headers, + bodyJson, + accountKey: "key-1:account-1", + }); + await waitFor(() => sentBodies.length === 1); + + // The upstream sends the connection limit error and immediately closes + // the socket; both events are already queued before the retry can + // detach its listeners. + sockets[0]?.dispatch("message", { + data: JSON.stringify({ + type: "error", + error: { code: "websocket_connection_limit_reached" }, + }), + }); + sockets[0]?.dispatch("close", { code: 1006, reason: "Connection ended" }); + + await waitFor(() => sentBodies.length === 2 && sockets.length === 2); + sockets[1]?.dispatch("message", { + data: JSON.stringify({ + type: "response.completed", + response: { id: "resp_race", status: "completed" }, + }), + }); + + const response = await responsePromise; + expect(response).not.toBeNull(); + expect(await response?.text()).toContain("response.completed"); + + // The stale close must not have marked the session for HTTP fallback. + const followUpPromise = tryProxyCodexWebSocket({ + headers, + bodyJson, + accountKey: "key-1:account-1", + }); + await waitFor(() => sentBodies.length === 3); + sockets.at(-1)?.dispatch("message", { + data: JSON.stringify({ + type: "response.completed", + response: { id: "resp_follow", status: "completed" }, + }), + }); + expect(await followUpPromise).not.toBeNull(); + }); + test("preserves websocket rate limit status_code errors", async () => { const sentBodies: unknown[] = []; const sockets = installManualCodexWebSocketMock(sentBodies); @@ -1576,6 +1807,88 @@ describe("proxy contract: codex", () => { expect(sockets).toHaveLength(6); }); + test("isolates close-before-terminal fallback to one session", async () => { + const sentBodies: unknown[] = []; + const sockets = installManualCodexWebSocketMock(sentBodies); + const headers = new Headers({ + authorization: "Bearer codex-access", + [CODEX_ACCOUNT_ID_HEADER]: "acct_1", + "x-session-affinity": "failed-session", + }); + const bodyJson = { + model: "gpt-5.5", + stream: true, + input: [{ role: "user", content: "hello" }], + }; + + const failedResponsePromise = tryProxyCodexWebSocket({ + headers, + bodyJson, + accountKey: "key-1:account-1", + }); + await waitFor(() => sentBodies.length === 1); + sockets[0]?.dispatch("message", { + data: JSON.stringify({ + type: "response.created", + response: { id: "resp_failed" }, + }), + }); + const failedResponse = await failedResponsePromise; + const failedText = failedResponse?.text(); + const warnings: string[] = []; + const originalWarn = console.warn; + console.warn = (message?: unknown): void => { + warnings.push(String(message)); + }; + let streamFailed: boolean | undefined; + try { + sockets[0]?.dispatch("close", { + code: 1006, + reason: "Connection ended", + }); + streamFailed = await failedText?.then( + () => false, + () => true + ); + } finally { + console.warn = originalWarn; + } + expect(streamFailed).toBe(true); + expect(warnings).toHaveLength(1); + const warning = JSON.parse(warnings[0] ?? "{}") as Record; + expect(warning.event).toBe("codex_websocket_stream_failed"); + expect(warning.payloadCount).toBe(1); + expect(typeof warning.requestBytes).toBe("number"); + expect(typeof warning.upstreamIdleMs).toBe("number"); + expect(typeof warning.downstreamIdleMs).toBe("number"); + + const sameSession = await tryProxyCodexWebSocket({ + headers, + bodyJson, + accountKey: "key-1:account-1", + }); + expect(sameSession).toBeNull(); + expect(sockets).toHaveLength(1); + + headers.set("x-session-affinity", "healthy-session"); + const healthyResponsePromise = tryProxyCodexWebSocket({ + headers, + bodyJson, + accountKey: "key-1:account-1", + }); + await waitFor(() => sentBodies.length === 2); + sockets[1]?.dispatch("message", { + data: JSON.stringify({ + type: "response.completed", + response: { id: "resp_healthy", status: "completed" }, + }), + }); + const healthyResponse = await healthyResponsePromise; + + expect(await healthyResponse?.text()).toContain("response.completed"); + expect(sockets).toHaveLength(2); + }); + test("treats same-session websocket connect races as busy", async () => { const sentBodies: unknown[] = []; const sockets = installManualCodexWebSocketMock(sentBodies, { @@ -2283,10 +2596,43 @@ describe("proxy contract: claude", () => { ); const transformedResponse = await result.transformResponse(sourceResponse); + expect(transformedResponse.headers.get("cache-control")).toBe( + "no-cache, no-transform" + ); + expect(transformedResponse.headers.get("x-accel-buffering")).toBe("no"); + expect(transformedResponse.headers.get("content-length")).toBeNull(); const transformedText = await transformedResponse.text(); expect(transformedText).toContain('"name":"shell"'); }); + test("applies downstream backpressure to Claude SSE reads", async () => { + const result = prepareClaudeUsageRequest(); + const encoder = new TextEncoder(); + let pulls = 0; + const source = new Response( + new ReadableStream( + { + pull(controller): void { + pulls++; + controller.enqueue( + encoder.encode( + `data: ${JSON.stringify({ type: "content_block_delta", index: pulls })}\n\n` + ) + ); + }, + }, + { highWaterMark: 0 } + ), + { headers: { "content-type": "text/event-stream" } } + ); + + const transformed = await result.transformResponse(source); + await new Promise((resolve) => setTimeout(resolve, 0)); + + expect(pulls).toBe(1); + await transformed.body?.cancel(); + }); + test("logs claude streaming error event details", async () => { const result = prepareClaudeUsageRequest(); const warnings: string[] = [];