diff --git a/packages/agent/src/agent.ts b/packages/agent/src/agent.ts index 809847bcb..27b649b28 100644 --- a/packages/agent/src/agent.ts +++ b/packages/agent/src/agent.ts @@ -10,9 +10,9 @@ import { type ImageContent, type Message, type Model, + type ProviderSessionState, streamSimple, type TextContent, - type ProviderSessionState, type ThinkingBudgets, type ToolChoice, type ToolResultMessage, diff --git a/packages/ai/CHANGELOG.md b/packages/ai/CHANGELOG.md index f7438a2c5..fd585d2aa 100644 --- a/packages/ai/CHANGELOG.md +++ b/packages/ai/CHANGELOG.md @@ -1,8 +1,10 @@ # Changelog ## [Unreleased] + ### Added +- Added automatic retry logic for WebSocket stream closures before response completion, with configurable retry budget to improve reliability on flaky connections - Added `providerSessionState` option to enable provider-scoped mutable state persistence across agent turns - Added WebSocket retry logic with configurable retry budget and delay via `PI_CODEX_WEBSOCKET_RETRY_BUDGET` and `PI_CODEX_WEBSOCKET_RETRY_DELAY_MS` environment variables - Added WebSocket idle timeout detection via `PI_CODEX_WEBSOCKET_IDLE_TIMEOUT_MS` environment variable to fail stalled connections @@ -33,10 +35,12 @@ ### Fixed +- Fixed WebSocket stream retry logic to properly handle mid-stream connection closures and retry before falling back to SSE transport +- Fixed `preferWebsockets` option handling to correctly respect explicit `false` values when determining transport preference - Fixed WebSocket append state not being reset after aborted requests, preventing stale state from affecting subsequent turns - Fixed WebSocket append state not being reset after stream errors, preventing failed append attempts from blocking future requests - - Fixed Codex model context window metadata to use 272000 input tokens (instead of 400000 total budget) for non-Spark Codex variants + ## [12.0.0] - 2026-02-12 ### Added diff --git a/packages/ai/src/providers/openai-codex-responses.ts b/packages/ai/src/providers/openai-codex-responses.ts index 166dda82f..f3006c5ab 100644 --- a/packages/ai/src/providers/openai-codex-responses.ts +++ b/packages/ai/src/providers/openai-codex-responses.ts @@ -177,6 +177,17 @@ function createCodexWebSocketTransportError(message: string): Error { return new Error(`${CODEX_WEBSOCKET_TRANSPORT_ERROR_PREFIX}: ${message}`); } +function isCodexWebSocketTransportError(error: unknown): boolean { + if (!(error instanceof Error)) return false; + return error.message.startsWith(CODEX_WEBSOCKET_TRANSPORT_ERROR_PREFIX); +} + +function isCodexWebSocketRetryableStreamError(error: unknown): boolean { + if (!(error instanceof Error) || !isCodexWebSocketTransportError(error)) return false; + const message = error.message.toLowerCase(); + return message.includes("websocket closed (") || message.includes("websocket closed before response completion"); +} + function toCodexHeaderRecord(value: unknown): Record | null { if (!value || typeof value !== "object") return null; const headers: Record = {}; @@ -455,217 +466,294 @@ export const streamOpenAICodexResponses: StreamFunction<"openai-codex-responses" let currentBlock: ThinkingContent | TextContent | (ToolCall & { partialJson: string }) | null = null; const blocks = output.content; const blockIndex = () => blocks.length - 1; - for await (const rawEvent of eventStream) { - const eventType = typeof rawEvent.type === "string" ? rawEvent.type : ""; - if (!eventType) continue; + let websocketStreamRetries = 0; + while (true) { + try { + for await (const rawEvent of eventStream) { + const eventType = typeof rawEvent.type === "string" ? rawEvent.type : ""; + if (!eventType) continue; - if (eventType === "response.output_item.added") { - if (!firstTokenTime) firstTokenTime = Date.now(); - const item = rawEvent.item as ResponseReasoningItem | ResponseOutputMessage | ResponseFunctionToolCall; - if (item.type === "reasoning") { - currentItem = item; - currentBlock = { type: "thinking", thinking: "" }; - output.content.push(currentBlock); - stream.push({ type: "thinking_start", contentIndex: blockIndex(), partial: output }); - } else if (item.type === "message") { - currentItem = item; - currentBlock = { type: "text", text: "" }; - output.content.push(currentBlock); - stream.push({ type: "text_start", contentIndex: blockIndex(), partial: output }); - } else if (item.type === "function_call") { - currentItem = item; - currentBlock = { - type: "toolCall", - id: `${item.call_id}|${item.id}`, - name: item.name, - arguments: {}, - partialJson: item.arguments || "", - }; - output.content.push(currentBlock); - stream.push({ type: "toolcall_start", contentIndex: blockIndex(), partial: output }); - } - } else if (eventType === "response.reasoning_summary_part.added") { - if (currentItem && currentItem.type === "reasoning") { - currentItem.summary = currentItem.summary || []; - currentItem.summary.push((rawEvent as { part: ResponseReasoningItem["summary"][number] }).part); - } - } else if (eventType === "response.reasoning_summary_text.delta") { - if (currentItem && currentItem.type === "reasoning" && currentBlock?.type === "thinking") { - currentItem.summary = currentItem.summary || []; - const lastPart = currentItem.summary[currentItem.summary.length - 1]; - if (lastPart) { - const delta = (rawEvent as { delta?: string }).delta || ""; - currentBlock.thinking += delta; - lastPart.text += delta; - stream.push({ - type: "thinking_delta", - contentIndex: blockIndex(), - delta, - partial: output, - }); - } - } - } else if (eventType === "response.reasoning_summary_part.done") { - if (currentItem && currentItem.type === "reasoning" && currentBlock?.type === "thinking") { - currentItem.summary = currentItem.summary || []; - const lastPart = currentItem.summary[currentItem.summary.length - 1]; - if (lastPart) { - currentBlock.thinking += "\n\n"; - lastPart.text += "\n\n"; - stream.push({ - type: "thinking_delta", - contentIndex: blockIndex(), - delta: "\n\n", - partial: output, - }); - } - } - } else if (eventType === "response.content_part.added") { - if (currentItem && currentItem.type === "message") { - currentItem.content = currentItem.content || []; - const part = (rawEvent as { part?: ResponseOutputMessage["content"][number] }).part; - if (part && (part.type === "output_text" || part.type === "refusal")) { - currentItem.content.push(part); - } - } - } else if (eventType === "response.output_text.delta") { - if (currentItem && currentItem.type === "message" && currentBlock?.type === "text") { - if (!currentItem.content || currentItem.content.length === 0) { - continue; - } - const lastPart = currentItem.content[currentItem.content.length - 1]; - if (lastPart && lastPart.type === "output_text") { - const delta = (rawEvent as { delta?: string }).delta || ""; - currentBlock.text += delta; - lastPart.text += delta; - stream.push({ - type: "text_delta", - contentIndex: blockIndex(), - delta, - partial: output, - }); - } - } - } else if (eventType === "response.refusal.delta") { - if (currentItem && currentItem.type === "message" && currentBlock?.type === "text") { - if (!currentItem.content || currentItem.content.length === 0) { - continue; - } - const lastPart = currentItem.content[currentItem.content.length - 1]; - if (lastPart && lastPart.type === "refusal") { - const delta = (rawEvent as { delta?: string }).delta || ""; - currentBlock.text += delta; - lastPart.refusal += delta; - stream.push({ - type: "text_delta", - contentIndex: blockIndex(), - delta, - partial: output, - }); - } - } - } else if (eventType === "response.function_call_arguments.delta") { - if (currentItem && currentItem.type === "function_call" && currentBlock?.type === "toolCall") { - const delta = (rawEvent as { delta?: string }).delta || ""; - currentBlock.partialJson += delta; - currentBlock.arguments = parseStreamingJson(currentBlock.partialJson); - stream.push({ - type: "toolcall_delta", - contentIndex: blockIndex(), - delta, - partial: output, - }); - } - } else if (eventType === "response.function_call_arguments.done") { - if (currentItem?.type === "function_call" && currentBlock?.type === "toolCall") { - const args = (rawEvent as { arguments?: string }).arguments; - if (typeof args === "string") { - currentBlock.partialJson = args; - currentBlock.arguments = parseStreamingJson(currentBlock.partialJson); - } - } - } else if (eventType === "response.output_item.done") { - const item = rawEvent.item as ResponseReasoningItem | ResponseOutputMessage | ResponseFunctionToolCall; - if (item.type === "reasoning" && currentBlock?.type === "thinking") { - currentBlock.thinking = item.summary?.map(s => s.text).join("\n\n") || ""; - currentBlock.thinkingSignature = JSON.stringify(item); - stream.push({ - type: "thinking_end", - contentIndex: blockIndex(), - content: currentBlock.thinking, - partial: output, - }); - currentBlock = null; - } else if (item.type === "message" && currentBlock?.type === "text") { - currentBlock.text = item.content.map(c => (c.type === "output_text" ? c.text : c.refusal)).join(""); - currentBlock.textSignature = item.id; - stream.push({ - type: "text_end", - contentIndex: blockIndex(), - content: currentBlock.text, - partial: output, - }); - currentBlock = null; - } else if (item.type === "function_call") { - const toolCall: ToolCall = { - type: "toolCall", - id: `${item.call_id}|${item.id}`, - name: item.name, - arguments: JSON.parse(item.arguments), - }; - stream.push({ type: "toolcall_end", contentIndex: blockIndex(), toolCall, partial: output }); - } - } else if (eventType === "response.created") { - if (usingWebsocket && websocketState) { - const createdResponse = (rawEvent as { response?: { id?: string } }).response; - if (typeof createdResponse?.id === "string" && createdResponse.id.length > 0) { - websocketState.lastResponseId = createdResponse.id; - } - } - } else if (eventType === "response.completed" || eventType === "response.done") { - const response = ( - rawEvent as { - response?: { - id?: string; - usage?: { - input_tokens?: number; - output_tokens?: number; - total_tokens?: number; - input_tokens_details?: { cached_tokens?: number }; + if (eventType === "response.output_item.added") { + if (!firstTokenTime) firstTokenTime = Date.now(); + const item = rawEvent.item as + | ResponseReasoningItem + | ResponseOutputMessage + | ResponseFunctionToolCall; + if (item.type === "reasoning") { + currentItem = item; + currentBlock = { type: "thinking", thinking: "" }; + output.content.push(currentBlock); + stream.push({ type: "thinking_start", contentIndex: blockIndex(), partial: output }); + } else if (item.type === "message") { + currentItem = item; + currentBlock = { type: "text", text: "" }; + output.content.push(currentBlock); + stream.push({ type: "text_start", contentIndex: blockIndex(), partial: output }); + } else if (item.type === "function_call") { + currentItem = item; + currentBlock = { + type: "toolCall", + id: `${item.call_id}|${item.id}`, + name: item.name, + arguments: {}, + partialJson: item.arguments || "", }; - status?: string; - }; + output.content.push(currentBlock); + stream.push({ type: "toolcall_start", contentIndex: blockIndex(), partial: output }); + } + } else if (eventType === "response.reasoning_summary_part.added") { + if (currentItem && currentItem.type === "reasoning") { + currentItem.summary = currentItem.summary || []; + currentItem.summary.push((rawEvent as { part: ResponseReasoningItem["summary"][number] }).part); + } + } else if (eventType === "response.reasoning_summary_text.delta") { + if (currentItem && currentItem.type === "reasoning" && currentBlock?.type === "thinking") { + currentItem.summary = currentItem.summary || []; + const lastPart = currentItem.summary[currentItem.summary.length - 1]; + if (lastPart) { + const delta = (rawEvent as { delta?: string }).delta || ""; + currentBlock.thinking += delta; + lastPart.text += delta; + stream.push({ + type: "thinking_delta", + contentIndex: blockIndex(), + delta, + partial: output, + }); + } + } + } else if (eventType === "response.reasoning_summary_part.done") { + if (currentItem && currentItem.type === "reasoning" && currentBlock?.type === "thinking") { + currentItem.summary = currentItem.summary || []; + const lastPart = currentItem.summary[currentItem.summary.length - 1]; + if (lastPart) { + currentBlock.thinking += "\n\n"; + lastPart.text += "\n\n"; + stream.push({ + type: "thinking_delta", + contentIndex: blockIndex(), + delta: "\n\n", + partial: output, + }); + } + } + } else if (eventType === "response.content_part.added") { + if (currentItem && currentItem.type === "message") { + currentItem.content = currentItem.content || []; + const part = (rawEvent as { part?: ResponseOutputMessage["content"][number] }).part; + if (part && (part.type === "output_text" || part.type === "refusal")) { + currentItem.content.push(part); + } + } + } else if (eventType === "response.output_text.delta") { + if (currentItem && currentItem.type === "message" && currentBlock?.type === "text") { + if (!currentItem.content || currentItem.content.length === 0) { + continue; + } + const lastPart = currentItem.content[currentItem.content.length - 1]; + if (lastPart && lastPart.type === "output_text") { + const delta = (rawEvent as { delta?: string }).delta || ""; + currentBlock.text += delta; + lastPart.text += delta; + stream.push({ + type: "text_delta", + contentIndex: blockIndex(), + delta, + partial: output, + }); + } + } + } else if (eventType === "response.refusal.delta") { + if (currentItem && currentItem.type === "message" && currentBlock?.type === "text") { + if (!currentItem.content || currentItem.content.length === 0) { + continue; + } + const lastPart = currentItem.content[currentItem.content.length - 1]; + if (lastPart && lastPart.type === "refusal") { + const delta = (rawEvent as { delta?: string }).delta || ""; + currentBlock.text += delta; + lastPart.refusal += delta; + stream.push({ + type: "text_delta", + contentIndex: blockIndex(), + delta, + partial: output, + }); + } + } + } else if (eventType === "response.function_call_arguments.delta") { + if (currentItem && currentItem.type === "function_call" && currentBlock?.type === "toolCall") { + const delta = (rawEvent as { delta?: string }).delta || ""; + currentBlock.partialJson += delta; + currentBlock.arguments = parseStreamingJson(currentBlock.partialJson); + stream.push({ + type: "toolcall_delta", + contentIndex: blockIndex(), + delta, + partial: output, + }); + } + } else if (eventType === "response.function_call_arguments.done") { + if (currentItem?.type === "function_call" && currentBlock?.type === "toolCall") { + const args = (rawEvent as { arguments?: string }).arguments; + if (typeof args === "string") { + currentBlock.partialJson = args; + currentBlock.arguments = parseStreamingJson(currentBlock.partialJson); + } + } + } else if (eventType === "response.output_item.done") { + const item = rawEvent.item as + | ResponseReasoningItem + | ResponseOutputMessage + | ResponseFunctionToolCall; + if (item.type === "reasoning" && currentBlock?.type === "thinking") { + currentBlock.thinking = item.summary?.map(s => s.text).join("\n\n") || ""; + currentBlock.thinkingSignature = JSON.stringify(item); + stream.push({ + type: "thinking_end", + contentIndex: blockIndex(), + content: currentBlock.thinking, + partial: output, + }); + currentBlock = null; + } else if (item.type === "message" && currentBlock?.type === "text") { + currentBlock.text = item.content + .map(c => (c.type === "output_text" ? c.text : c.refusal)) + .join(""); + currentBlock.textSignature = item.id; + stream.push({ + type: "text_end", + contentIndex: blockIndex(), + content: currentBlock.text, + partial: output, + }); + currentBlock = null; + } else if (item.type === "function_call") { + const toolCall: ToolCall = { + type: "toolCall", + id: `${item.call_id}|${item.id}`, + name: item.name, + arguments: JSON.parse(item.arguments), + }; + stream.push({ type: "toolcall_end", contentIndex: blockIndex(), toolCall, partial: output }); + } + } else if (eventType === "response.created") { + if (usingWebsocket && websocketState) { + const createdResponse = (rawEvent as { response?: { id?: string } }).response; + if (typeof createdResponse?.id === "string" && createdResponse.id.length > 0) { + websocketState.lastResponseId = createdResponse.id; + } + } + } else if (eventType === "response.completed" || eventType === "response.done") { + const response = ( + rawEvent as { + response?: { + id?: string; + usage?: { + input_tokens?: number; + output_tokens?: number; + total_tokens?: number; + input_tokens_details?: { cached_tokens?: number }; + }; + status?: string; + }; + } + ).response; + if (response?.usage) { + const cachedTokens = response.usage.input_tokens_details?.cached_tokens || 0; + output.usage = { + input: (response.usage.input_tokens || 0) - cachedTokens, + output: response.usage.output_tokens || 0, + cacheRead: cachedTokens, + cacheWrite: 0, + totalTokens: response.usage.total_tokens || 0, + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 }, + }; + } + if (usingWebsocket && websocketState) { + websocketState.lastRequest = cloneRequestBody(requestBodyForState); + if (typeof response?.id === "string" && response.id.length > 0) { + websocketState.lastResponseId = response.id; + } + websocketState.canAppend = eventType === "response.done"; + } + calculateCost(model, output.usage); + output.stopReason = mapStopReason(response?.status); + if (output.content.some(b => b.type === "toolCall") && output.stopReason === "stop") { + output.stopReason = "toolUse"; + } + } else if (eventType === "error") { + const code = (rawEvent as { code?: string }).code || ""; + const message = (rawEvent as { message?: string }).message || ""; + throw new Error(formatCodexErrorEvent(rawEvent, code, message)); + } else if (eventType === "response.failed") { + throw new Error(formatCodexFailure(rawEvent) ?? "Codex response failed"); } - ).response; - if (response?.usage) { - const cachedTokens = response.usage.input_tokens_details?.cached_tokens || 0; - output.usage = { - input: (response.usage.input_tokens || 0) - cachedTokens, - output: response.usage.output_tokens || 0, - cacheRead: cachedTokens, - cacheWrite: 0, - totalTokens: response.usage.total_tokens || 0, - cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 }, - }; } - if (usingWebsocket && websocketState) { - websocketState.lastRequest = cloneRequestBody(requestBodyForState); - if (typeof response?.id === "string" && response.id.length > 0) { - websocketState.lastResponseId = response.id; + + break; + } catch (error) { + if ( + usingWebsocket && + websocketState && + isCodexWebSocketRetryableStreamError(error) && + output.content.length === 0 && + !options?.signal?.aborted + ) { + const activateFallback = websocketStreamRetries >= getCodexWebSocketRetryBudget(); + recordCodexWebSocketFailure(websocketState, activateFallback); + logCodexDebug("codex websocket stream fallback", { + error: error instanceof Error ? error.message : String(error), + retry: websocketStreamRetries, + retryBudget: getCodexWebSocketRetryBudget(), + activated: activateFallback, + }); + if (!activateFallback) { + websocketStreamRetries += 1; + await abortableSleep(getCodexWebSocketRetryDelayMs(websocketStreamRetries), options?.signal); + const websocketV2Enabled = isCodexWebSocketV2Enabled(); + const websocketHeaders = createCodexHeaders( + requestHeaders, + accountId, + apiKey, + options?.sessionId, + "websocket", + websocketState, + websocketV2Enabled, + ); + const websocketRequest = buildCodexWebSocketRequest( + transformedBody, + websocketState, + websocketV2Enabled, + ); + requestBodyForState = cloneRequestBody(transformedBody); + eventStream = await openCodexWebSocketEventStream( + toWebSocketUrl(url), + websocketHeaders, + websocketRequest, + websocketState, + options?.signal, + ); + usingWebsocket = true; + websocketState.lastTransport = "websocket"; + continue; } - websocketState.canAppend = eventType === "response.done"; + eventStream = await openCodexSseEventStream( + url, + requestHeaders, + accountId, + apiKey, + options?.sessionId, + transformedBody, + websocketState, + options?.signal, + ); + usingWebsocket = false; + websocketState.lastTransport = "sse"; + requestBodyForState = cloneRequestBody(transformedBody); + continue; } - calculateCost(model, output.usage); - output.stopReason = mapStopReason(response?.status); - if (output.content.some(b => b.type === "toolCall") && output.stopReason === "stop") { - output.stopReason = "toolUse"; - } - } else if (eventType === "error") { - const code = (rawEvent as { code?: string }).code || ""; - const message = (rawEvent as { message?: string }).message || ""; - throw new Error(formatCodexErrorEvent(rawEvent, code, message)); - } else if (eventType === "response.failed") { - throw new Error(formatCodexFailure(rawEvent) ?? "Codex response failed"); + throw error; } } @@ -795,6 +883,7 @@ function shouldUseCodexWebSocket( preferWebsockets?: boolean, ): boolean { if (!state || state.disableWebsocket) return false; + if (preferWebsockets === false) return false; return isCodexWebSocketEnvEnabled() || preferWebsockets === true || model.preferWebsockets === true; } @@ -821,7 +910,9 @@ export function getOpenAICodexTransportDetails( ): OpenAICodexTransportDetails { const baseUrl = options?.baseUrl || model.baseUrl || CODEX_BASE_URL; const websocketPreferred = - isCodexWebSocketEnvEnabled() || options?.preferWebsockets === true || model.preferWebsockets === true; + options?.preferWebsockets === false + ? false + : isCodexWebSocketEnvEnabled() || options?.preferWebsockets === true || model.preferWebsockets === true; const providerSessionState = getCodexProviderSessionState(options?.providerSessionState); const publicSessionKey = getCodexPublicSessionKey(options?.sessionId, model, baseUrl); const privateSessionKey = publicSessionKey diff --git a/packages/ai/test/openai-codex-stream.test.ts b/packages/ai/test/openai-codex-stream.test.ts index 402af3f7d..377ed5cb7 100644 --- a/packages/ai/test/openai-codex-stream.test.ts +++ b/packages/ai/test/openai-codex-stream.test.ts @@ -1060,6 +1060,155 @@ describe("openai-codex streaming", () => { expect(fetchMock).not.toHaveBeenCalled(); }); + it("retries websocket stream closes before surfacing transport errors", async () => { + const tempDir = TempDir.createSync("@pi-codex-stream-"); + setAgentDir(tempDir.path()); + Bun.env.PI_CODEX_WEBSOCKET_RETRY_BUDGET = "1"; + Bun.env.PI_CODEX_WEBSOCKET_RETRY_DELAY_MS = "1"; + + const payload = Buffer.from( + JSON.stringify({ "https://api.openai.com/auth": { chatgpt_account_id: "acc_test" } }), + "utf8", + ).toBase64(); + const token = `aaa.${payload}.bbb`; + const fetchMock = vi.fn(async () => { + throw new Error("SSE fallback should not be called when websocket retry succeeds"); + }); + global.fetch = fetchMock as unknown as typeof fetch; + + type WsListener = (event: Event) => void; + let constructorCount = 0; + const requestTypes: string[] = []; + + class FlakyCloseWebSocket { + static readonly CONNECTING = 0; + static readonly OPEN = 1; + static readonly CLOSING = 2; + static readonly CLOSED = 3; + readyState = FlakyCloseWebSocket.CONNECTING; + #listeners = new Map>(); + + constructor(_url: string, _options?: { headers?: Record }) { + constructorCount += 1; + setTimeout(() => { + this.readyState = FlakyCloseWebSocket.OPEN; + this.#emit("open", new Event("open")); + }, 0); + } + + addEventListener(type: string, listener: unknown): void { + if (typeof listener !== "function") return; + const listeners = this.#listeners.get(type) ?? new Set(); + listeners.add(listener as WsListener); + this.#listeners.set(type, listeners); + } + + removeEventListener(type: string, listener: unknown): void { + if (typeof listener !== "function") return; + const listeners = this.#listeners.get(type); + listeners?.delete(listener as WsListener); + } + + send(data: string): void { + const request = JSON.parse(data) as { type?: string }; + requestTypes.push(typeof request.type === "string" ? request.type : ""); + if (requestTypes.length === 1) { + this.readyState = FlakyCloseWebSocket.CLOSED; + this.#emit("close", { code: 1012 } as unknown as Event); + return; + } + this.#emit("message", { + data: JSON.stringify({ + type: "response.output_item.added", + item: { + type: "message", + id: "msg_retry_close", + role: "assistant", + status: "in_progress", + content: [], + }, + }), + } as unknown as Event); + this.#emit("message", { + data: JSON.stringify({ type: "response.content_part.added", part: { type: "output_text", text: "" } }), + } as unknown as Event); + this.#emit("message", { + data: JSON.stringify({ type: "response.output_text.delta", delta: "Hello retry close" }), + } as unknown as Event); + this.#emit("message", { + data: JSON.stringify({ + type: "response.output_item.done", + item: { + type: "message", + id: "msg_retry_close", + role: "assistant", + status: "completed", + content: [{ type: "output_text", text: "Hello retry close" }], + }, + }), + } as unknown as Event); + this.#emit("message", { + data: JSON.stringify({ + type: "response.done", + response: { + id: "resp_retry_close", + status: "completed", + usage: { + input_tokens: 5, + output_tokens: 3, + total_tokens: 8, + input_tokens_details: { cached_tokens: 0 }, + }, + }, + }), + } as unknown as Event); + } + + close(): void { + this.readyState = FlakyCloseWebSocket.CLOSED; + } + + #emit(type: string, event: Event): void { + const listeners = this.#listeners.get(type); + if (!listeners) return; + for (const listener of listeners) { + listener(event); + } + } + } + + global.WebSocket = FlakyCloseWebSocket as unknown as typeof WebSocket; + + const model: Model<"openai-codex-responses"> = { + id: "gpt-5.3-codex-spark", + name: "GPT-5.3 Codex Spark", + api: "openai-codex-responses", + provider: "openai-codex", + baseUrl: "https://chatgpt.com/backend-api", + reasoning: true, + preferWebsockets: true, + input: ["text"], + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 }, + contextWindow: 128000, + maxTokens: 128000, + }; + const context: Context = { + systemPrompt: "You are a helpful assistant.", + messages: [{ role: "user", content: "Say hello", timestamp: Date.now() }], + }; + const providerSessionState = new Map(); + const result = await streamOpenAICodexResponses(model, context, { + apiKey: token, + sessionId: "ws-retry-close-session", + providerSessionState, + }).result(); + + expect(result.role).toBe("assistant"); + expect(constructorCount).toBe(2); + expect(requestTypes).toEqual(["response.create", "response.create"]); + expect(fetchMock).not.toHaveBeenCalled(); + }); + it("resets websocket append state after an aborted request closes the connection", async () => { const tempDir = TempDir.createSync("@pi-codex-stream-"); setAgentDir(tempDir.path()); diff --git a/packages/coding-agent/CHANGELOG.md b/packages/coding-agent/CHANGELOG.md index 9bfe527b1..67841c472 100644 --- a/packages/coding-agent/CHANGELOG.md +++ b/packages/coding-agent/CHANGELOG.md @@ -1,6 +1,7 @@ # Changelog ## [Unreleased] + ### Added - Added `providerSessionState` property to AgentSession for managing provider-scoped transport and session caches @@ -13,6 +14,7 @@ ### Changed +- Changed `providers.openaiWebsockets` setting from boolean to enum with values "auto", "off", "on" for more granular websocket policy control (auto uses model defaults, on forces websocket, off disables it) - Enhanced provider details display to include live provider session state information - Enhanced session info output to display active provider configuration and authentication details - Replaced `process.cwd()` with `getProjectDir()` throughout codebase for improved project directory detection and handling diff --git a/packages/coding-agent/src/config/settings-schema.ts b/packages/coding-agent/src/config/settings-schema.ts index cd5033b0b..5f3ea285b 100644 --- a/packages/coding-agent/src/config/settings-schema.ts +++ b/packages/coding-agent/src/config/settings-schema.ts @@ -593,12 +593,14 @@ export const SETTINGS_SCHEMA = { }, }, "providers.openaiWebsockets": { - type: "boolean", - default: false, + type: "enum", + values: ["auto", "off", "on"] as const, + default: "auto", ui: { tab: "services", label: "OpenAI websockets", - description: "Prefer websocket transport for OpenAI Codex models", + description: "Websocket policy for OpenAI Codex models (auto uses model defaults, on forces, off disables)", + submenu: true, }, }, diff --git a/packages/coding-agent/src/modes/controllers/command-controller.ts b/packages/coding-agent/src/modes/controllers/command-controller.ts index fb3be923e..e8b56a653 100644 --- a/packages/coding-agent/src/modes/controllers/command-controller.ts +++ b/packages/coding-agent/src/modes/controllers/command-controller.ts @@ -23,8 +23,8 @@ import { DynamicBorder } from "../../modes/components/dynamic-border"; import { PythonExecutionComponent } from "../../modes/components/python-execution"; import { getMarkdownTheme, getSymbolTheme, theme } from "../../modes/theme/theme"; import type { InteractiveModeContext } from "../../modes/types"; -import { createCompactionSummaryMessage } from "../../session/messages"; import type { AuthStorage } from "../../session/auth-storage"; +import { createCompactionSummaryMessage } from "../../session/messages"; import { outputMeta } from "../../tools/output-meta"; import { resolveToCwd } from "../../tools/path-utils"; import { getChangelogPath, parseChangelog } from "../../utils/changelog"; @@ -220,11 +220,14 @@ export class CommandController { info += `${theme.fg("dim", "No model selected")}\n`; } else { const authMode = resolveProviderAuthMode(this.ctx.session.modelRegistry.authStorage, model.provider); + const openaiWebsocketSetting = this.ctx.settings.get("providers.openaiWebsockets") ?? "auto"; + const preferOpenAICodexWebsockets = + openaiWebsocketSetting === "on" ? true : openaiWebsocketSetting === "off" ? false : undefined; const providerDetails = getProviderDetails({ model, sessionId: stats.sessionId, authMode, - preferWebsockets: this.ctx.settings.get("providers.openaiWebsockets") ?? false, + preferWebsockets: preferOpenAICodexWebsockets, providerSessionState: this.ctx.session.providerSessionState, }); info += renderProviderSection(providerDetails, theme); diff --git a/packages/coding-agent/src/sdk.ts b/packages/coding-agent/src/sdk.ts index b68d66dfa..34a013efe 100644 --- a/packages/coding-agent/src/sdk.ts +++ b/packages/coding-agent/src/sdk.ts @@ -1017,6 +1017,10 @@ export async function createAgentSession(options: CreateAgentSessionOptions = {} .map(name => toolRegistry.get(name)) .filter((tool): tool is AgentTool => tool !== undefined); + const openaiWebsocketSetting = settings.get("providers.openaiWebsockets") ?? "auto"; + const preferOpenAICodexWebsockets = + openaiWebsocketSetting === "on" ? true : openaiWebsocketSetting === "off" ? false : undefined; + agent = new Agent({ initialState: { systemPrompt, @@ -1037,7 +1041,7 @@ export async function createAgentSession(options: CreateAgentSessionOptions = {} thinkingBudgets: settings.getGroup("thinkingBudgets"), temperature: settings.get("temperature") >= 0 ? settings.get("temperature") : undefined, kimiApiFormat: settings.get("providers.kimiApiFormat") ?? "anthropic", - preferWebsockets: settings.get("providers.openaiWebsockets") ?? false, + preferWebsockets: preferOpenAICodexWebsockets, getToolContext: tc => toolContextStore.getContext(tc), getApiKey: async provider => { // Use the provider argument from the in-flight request; @@ -1095,7 +1099,7 @@ export async function createAgentSession(options: CreateAgentSessionOptions = {} await prewarmOpenAICodexResponses(model, { apiKey: await modelRegistry.getApiKey(model, sessionId), sessionId, - preferWebsockets: settings.get("providers.openaiWebsockets") ?? false, + preferWebsockets: preferOpenAICodexWebsockets, providerSessionState: session.providerSessionState, }); debugStartup("sdk:prewarmCodexWebsocket:done"); diff --git a/packages/coding-agent/src/session/agent-session.ts b/packages/coding-agent/src/session/agent-session.ts index 919b65a9f..0759d406b 100644 --- a/packages/coding-agent/src/session/agent-session.ts +++ b/packages/coding-agent/src/session/agent-session.ts @@ -21,12 +21,12 @@ import type { ImageContent, Message, Model, + ProviderSessionState, TextContent, ToolCall, ToolChoice, Usage, UsageReport, - ProviderSessionState, } from "@oh-my-pi/pi-ai"; import { isContextOverflow, modelsAreEqual, supportsXhigh } from "@oh-my-pi/pi-ai"; import { abortableSleep, isEnoent, logger } from "@oh-my-pi/pi-utils"; diff --git a/packages/coding-agent/test/session-provider-section.test.ts b/packages/coding-agent/test/session-provider-section.test.ts index 4ad95f5e9..59d96a8ae 100644 --- a/packages/coding-agent/test/session-provider-section.test.ts +++ b/packages/coding-agent/test/session-provider-section.test.ts @@ -1,6 +1,7 @@ import { describe, expect, it } from "bun:test"; import { getProviderDetails, type Model } from "@oh-my-pi/pi-ai"; import { renderProviderSection } from "@oh-my-pi/pi-coding-agent/modes/controllers/command-controller"; + describe("session provider section", () => { it("renders codex provider details with transport fields", () => { const model: Model<"openai-codex-responses"> = { diff --git a/packages/natives/test/native.test.ts b/packages/natives/test/native.test.ts index 9ec91c810..5e42ea7d3 100644 --- a/packages/natives/test/native.test.ts +++ b/packages/natives/test/native.test.ts @@ -2,7 +2,6 @@ import { beforeAll, describe, expect, it } from "bun:test"; import * as fs from "node:fs/promises"; import * as os from "node:os"; import * as path from "node:path"; -import { getProjectDir } from "@oh-my-pi/pi-utils/dirs"; import { FileType, fuzzyFind, type GlobMatch, glob, grep, htmlToMarkdown, invalidateFsScanCache } from "../src/index"; let testDir: string; @@ -136,7 +135,7 @@ describe("pi-natives", () => { const newFile = path.join(testDir, "newly-added.ts"); await fs.writeFile(newFile, "export const newer = true;\n"); - const relativePath = path.relative(getProjectDir(), newFile); + const relativePath = path.relative(process.cwd(), newFile); invalidateFsScanCache(relativePath); const result = await glob({ pattern: "newly-added.ts", path: testDir, cache: true });