diff --git a/packages/types/src/__tests__/provider-settings.test.ts b/packages/types/src/__tests__/provider-settings.test.ts index c5a79002a4..cb14b21275 100644 --- a/packages/types/src/__tests__/provider-settings.test.ts +++ b/packages/types/src/__tests__/provider-settings.test.ts @@ -138,6 +138,26 @@ describe("OpenAI Codex provider settings", () => { }) }) +describe("OpenAI Codex WebSocket preference", () => { + it.each([true, false, undefined])("round-trips %s through both provider schemas", (enabled) => { + const settings = { + apiProvider: providerIdentifiers.openaiCodex, + ...(enabled !== undefined ? { openAiCodexUseWebSocket: enabled } : {}), + } + expect(providerSettingsSchema.parse(settings)).toEqual(settings) + expect(providerSettingsSchemaDiscriminated.parse(JSON.parse(JSON.stringify(settings)))).toEqual(settings) + expect(PROVIDER_SETTINGS_KEYS).toContain("openAiCodexUseWebSocket") + }) + it.each(["true", 1, null])("rejects non-boolean value %s", (enabled) => { + expect( + providerSettingsSchema.safeParse({ + apiProvider: providerIdentifiers.openaiCodex, + openAiCodexUseWebSocket: enabled, + }).success, + ).toBe(false) + }) +}) + describe("getApiProtocol", () => { it("preserves API protocol wire values", () => { expect(ANTHROPIC_API_PROTOCOL).toBe("anthropic") diff --git a/packages/types/src/provider-settings.ts b/packages/types/src/provider-settings.ts index b243971cc3..0aeb4888dd 100644 --- a/packages/types/src/provider-settings.ts +++ b/packages/types/src/provider-settings.ts @@ -5,6 +5,7 @@ import { API_PROVIDER_FIELD, SETTINGS_SHAPE_FIELD } from "./provider-settings/co export { DEFAULT_OPEN_AI_STRICT_TOOL_SCHEMAS, OPEN_AI_CODEX_SERVICE_TIER_KEY, + DEFAULT_OPEN_AI_CODEX_USE_WEBSOCKET, parseOpenAiExtraBody, kimiCodeAuthMethodSchema, type KimiCodeAuthMethod, diff --git a/packages/types/src/provider-settings/index.ts b/packages/types/src/provider-settings/index.ts index 5cdadb4b2c..814c367c45 100644 --- a/packages/types/src/provider-settings/index.ts +++ b/packages/types/src/provider-settings/index.ts @@ -36,7 +36,7 @@ import { basetenProviderDefinition } from "./baseten.js" import type { ProviderDefinition } from "./common.js" -export { OPEN_AI_CODEX_SERVICE_TIER_KEY } from "./openai-codex.js" +export { OPEN_AI_CODEX_SERVICE_TIER_KEY, DEFAULT_OPEN_AI_CODEX_USE_WEBSOCKET } from "./openai-codex.js" export { parseOpenAiExtraBody, DEFAULT_OPEN_AI_STRICT_TOOL_SCHEMAS } from "./openai.js" export { kimiCodeAuthMethodSchema, type KimiCodeAuthMethod } from "./kimi-code.js" export { zaiApiLineSchema, type ZaiApiLine } from "./zai.js" diff --git a/packages/types/src/provider-settings/openai-codex.ts b/packages/types/src/provider-settings/openai-codex.ts index 93117224ae..2c2edbfd5d 100644 --- a/packages/types/src/provider-settings/openai-codex.ts +++ b/packages/types/src/provider-settings/openai-codex.ts @@ -1,3 +1,5 @@ +import { z } from "zod" + import { providerIdentifiers } from "../provider-identifiers.js" import { openAiCodexServiceTierSchema } from "../model.js" import { @@ -7,6 +9,8 @@ import { createProviderDefinition, } from "./common.js" +export const DEFAULT_OPEN_AI_CODEX_USE_WEBSOCKET = false + export const OPEN_AI_CODEX_SERVICE_TIER_KEY = "openAiCodexServiceTier" export const openAiCodexProviderDefinition = createProviderDefinition({ @@ -15,6 +19,7 @@ export const openAiCodexProviderDefinition = createProviderDefinition({ getModelId: createModelIdAccessor(API_MODEL_ID_FIELD), schema: { ...apiModelIdProviderModelShape, + openAiCodexUseWebSocket: z.boolean().optional(), // Codex "Fast" mode maps to the Responses API priority service tier. [OPEN_AI_CODEX_SERVICE_TIER_KEY]: openAiCodexServiceTierSchema.optional(), }, diff --git a/src/api/providers/__tests__/openai-codex.spec.ts b/src/api/providers/__tests__/openai-codex.spec.ts index 35c1f5de60..5dd4f96ff8 100644 --- a/src/api/providers/__tests__/openai-codex.spec.ts +++ b/src/api/providers/__tests__/openai-codex.spec.ts @@ -11,6 +11,8 @@ vitest.mock("@roo-code/telemetry", () => ({ import { Anthropic } from "@anthropic-ai/sdk" import { OPEN_AI_CODEX_SERVICE_TIER_KEY, OpenAiCodexServiceTier, SERVICE_TIER_KEY } from "@roo-code/types" import { OpenAiCodexHandler, transformResponsesLiteBody } from "../openai-codex" +import { CodexWebSocketTransport, CodexWebSocketUnavailableError } from "../CodexWebSocketTransport" +import { CodexWebSocketTransportScope } from "../codex-websocket/scopes/CodexWebSocketTransportScope" import { openAiCodexOAuthManager } from "../../../integrations/openai-codex/oauth" import { asyncStreamFrom, collectStream } from "../../../test-utils/stream" @@ -28,6 +30,111 @@ function createCompletedStream() { ]) } +describe("OpenAiCodexHandler WebSocket transport", () => { + beforeEach(() => { + vitest.spyOn(openAiCodexOAuthManager, "getAccessToken").mockResolvedValue("test-token") + vitest.spyOn(openAiCodexOAuthManager, "getAccountId").mockResolvedValue("acct_test") + }) + afterEach(() => { + vitest.restoreAllMocks() + vitest.unstubAllEnvs() + vitest.unstubAllGlobals() + }) + + it("initializes the DI scope on the first request rather than in the provider constructor", async () => { + const init = vitest.spyOn(CodexWebSocketTransportScope.prototype, "init") + vitest.spyOn(CodexWebSocketTransport.prototype, "stream").mockImplementation(() => createCompletedStream()) + const handler = new OpenAiCodexHandler({ apiModelId: "gpt-5.6-sol", openAiCodexUseWebSocket: true }) + expect(init).not.toHaveBeenCalled() + await collectStream(handler.createMessage("System", [])) + await collectStream(handler.createMessage("System", [])) + expect(init).toHaveBeenCalledOnce() + }) + + it("routes Lite requests through the transport and returns existing usage chunks", async () => { + const stream = vitest + .spyOn(CodexWebSocketTransport.prototype, "stream") + .mockReturnValue(createCompletedStream()) + const reset = vitest.spyOn(CodexWebSocketTransport.prototype, "resetContinuation") + const handler = new OpenAiCodexHandler({ apiModelId: "gpt-6.1-sol", openAiCodexUseWebSocket: true }) + const chunks = await collectStream( + handler.createMessage("System", [], { + taskId: "task-1", + suppressPreviousResponseId: true, + }), + ) + expect(stream).toHaveBeenCalledWith( + expect.objectContaining({ + input: expect.arrayContaining([expect.objectContaining({ type: "additional_tools" })]), + }), + expect.objectContaining({ + headers: expect.objectContaining({ + Authorization: "Bearer test-token", + "ChatGPT-Account-Id": "acct_test", + session_id: "task-1", + }), + }), + ) + expect(reset).toHaveBeenCalledOnce() + expect(chunks).toContainEqual(expect.objectContaining({ type: "usage", inputTokens: 1, outputTokens: 1 })) + }) + + it.each([undefined, false])("keeps SDK streaming when WebSocket is %s", async (enabled) => { + vitest.stubEnv("ZOO_CODE_CODEX_WEBSOCKET", "1") + const stream = vitest.spyOn(CodexWebSocketTransport.prototype, "stream") + const handler = new OpenAiCodexHandler({ apiModelId: "gpt-5.6-sol", openAiCodexUseWebSocket: enabled }) + const create = vitest.fn().mockResolvedValue(createCompletedStream()) + Reflect.set(handler, "client", { responses: { create } }) + await collectStream(handler.createMessage("System", [])) + expect(create).toHaveBeenCalledOnce() + expect(stream).not.toHaveBeenCalled() + }) + + it("uses the HTTP fallback for a failed upgrade", async () => { + vitest.spyOn(CodexWebSocketTransport.prototype, "stream").mockImplementation(() => { + throw new CodexWebSocketUnavailableError("Upgrade rejected") + }) + const fetch = vitest.fn().mockResolvedValue({ + ok: true, + body: new ReadableStream({ + start(controller) { + controller.enqueue( + new TextEncoder().encode( + 'data: {"type":"response.completed","response":{"output":[],"usage":{"input_tokens":1,"output_tokens":1}}}\n\n', + ), + ) + controller.close() + }, + }), + }) + vitest.stubGlobal("fetch", fetch) + await collectStream( + new OpenAiCodexHandler({ apiModelId: "gpt-5.6-sol", openAiCodexUseWebSocket: true }).createMessage( + "System", + [], + ), + ) + expect(fetch).toHaveBeenCalledOnce() + }) + + it("does not replay an ambiguously accepted request over HTTP", async () => { + vitest.spyOn(CodexWebSocketTransport.prototype, "stream").mockImplementation(() => { + throw new Error("Connection closed after send") + }) + const fetch = vitest.fn() + vitest.stubGlobal("fetch", fetch) + await expect( + collectStream( + new OpenAiCodexHandler({ apiModelId: "gpt-5.6-sol", openAiCodexUseWebSocket: true }).createMessage( + "System", + [], + ), + ), + ).rejects.toThrow("Connection closed after send") + expect(fetch).not.toHaveBeenCalled() + }) +}) + describe("OpenAiCodexHandler.getModel", () => { it.each(["gpt-5.1", "gpt-5", "gpt-5.1-codex", "gpt-5-codex", "gpt-5-codex-mini", "gpt-5.3-codex-spark"])( "should return specified model when a valid model id is provided: %s", diff --git a/src/api/providers/openai-codex.ts b/src/api/providers/openai-codex.ts index ff8f4a24b0..b5328351ce 100644 --- a/src/api/providers/openai-codex.ts +++ b/src/api/providers/openai-codex.ts @@ -6,6 +6,7 @@ import OpenAI from "openai" import { type ModelInfo, OPEN_AI_CODEX_SERVICE_TIER_KEY, + DEFAULT_OPEN_AI_CODEX_USE_WEBSOCKET, OpenAiCodexServiceTier, openAiCodexDefaultModelId, OpenAiCodexModelId, @@ -24,6 +25,8 @@ import { ApiStream, ApiStreamUsageChunk } from "../transform/stream" import { getModelParams } from "../transform/model-params" import { BaseProvider } from "./base-provider" +import { type CodexWebSocketTransport, CodexWebSocketUnavailableError } from "./CodexWebSocketTransport" +import { CodexWebSocketTransportScope } from "./codex-websocket/scopes/CodexWebSocketTransportScope" import type { SingleCompletionHandler, ApiHandlerCreateMessageMetadata, CompletePromptOptions } from "../index" import { isMcpTool } from "../../utils/mcp-name" import { sanitizeOpenAiCallId } from "../../utils/tool-id" @@ -129,6 +132,8 @@ export class OpenAiCodexHandler extends BaseProvider implements SingleCompletion protected options: ApiHandlerOptions private readonly providerName = "OpenAI Codex" private client?: OpenAI + private readonly webSocketScope?: CodexWebSocketTransportScope + private webSocketTransport?: CodexWebSocketTransport // Complete response output array private lastResponseOutput: any[] | undefined // Last top-level response id @@ -181,6 +186,9 @@ export class OpenAiCodexHandler extends BaseProvider implements SingleCompletion constructor(options: ApiHandlerOptions) { super() this.options = options + if (options.openAiCodexUseWebSocket ?? DEFAULT_OPEN_AI_CODEX_USE_WEBSOCKET) { + this.webSocketScope = new CodexWebSocketTransportScope() + } // Generate a new session ID for standalone handler usage (fallback) this.sessionId = uuidv7() } @@ -248,6 +256,11 @@ export class OpenAiCodexHandler extends BaseProvider implements SingleCompletion this.sawTextDeltaInCurrentResponse = false this.sawSdkEventInCurrentResponse = false this.streamedToolCallIds.clear() + if (this.webSocketScope && !this.webSocketTransport) { + this.webSocketScope.init() + this.webSocketTransport = this.webSocketScope.transport + } + if (metadata?.suppressPreviousResponseId) this.webSocketTransport?.resetContinuation() // Get access token from OAuth manager let accessToken = await openAiCodexOAuthManager.getAccessToken() @@ -491,13 +504,20 @@ export class OpenAiCodexHandler extends BaseProvider implements SingleCompletion timeout: this.timeoutMs, }) - const stream = (await (client as any).responses.create(requestBody, { - signal: this.abortController.signal, - // If the SDK supports per-request overrides, ensure headers are present. - headers: codexHeaders, - })) as AsyncIterable - - if (typeof (stream as any)?.[Symbol.asyncIterator] !== "function") { + const sdkRequest: OpenAI.Responses.ResponseCreateParamsStreaming = requestBody + const stream = this.webSocketTransport + ? this.webSocketTransport.stream(requestBody, { + headers: { ...codexHeaders, Authorization: `Bearer ${accessToken}` }, + signal: this.abortController.signal, + timeoutMs: this.timeoutMs, + }) + : await client.responses.create(sdkRequest, { + signal: this.abortController.signal, + // If the SDK supports per-request overrides, ensure headers are present. + headers: codexHeaders, + }) + + if (typeof stream?.[Symbol.asyncIterator] !== "function") { throw new Error( "OpenAI SDK did not return an AsyncIterable for Responses API streaming. Falling back to SSE.", ) @@ -530,6 +550,11 @@ export class OpenAiCodexHandler extends BaseProvider implements SingleCompletion if (this.sawSdkEventInCurrentResponse || this.abortController?.signal.aborted) { throw sdkErr } + // A sent WebSocket request may be accepted even before its first event arrives. + // Only a failed HTTP upgrade is safe to replay through the existing fallback. + if (this.webSocketTransport && !(sdkErr instanceof CodexWebSocketUnavailableError)) { + throw sdkErr + } // Fallback to manual SSE via fetch (Codex backend). yield* this.makeCodexRequest(requestBody, model, accessToken, effectiveSessionId) diff --git a/src/core/config/__tests__/ProviderSettingsManager.spec.ts b/src/core/config/__tests__/ProviderSettingsManager.spec.ts index 13e1aeeb2d..297d9fd877 100644 --- a/src/core/config/__tests__/ProviderSettingsManager.spec.ts +++ b/src/core/config/__tests__/ProviderSettingsManager.spec.ts @@ -459,6 +459,41 @@ describe("ProviderSettingsManager", () => { expect(storedConfig).toEqual(expectedConfig) }) + it.each([true, false, undefined])( + "round-trips Codex WebSocket preference %s in a saved profile", + async (enabled) => { + mockSecrets.get.mockResolvedValue( + JSON.stringify({ + currentApiConfigName: "default", + apiConfigs: { default: {} }, + modeApiConfigs: {}, + }), + ) + const configuration = { + apiProvider: providerIdentifiers.openaiCodex, + apiModelId: "gpt-5.6-sol", + ...(enabled !== undefined ? { openAiCodexUseWebSocket: enabled } : {}), + } + await providerSettingsManager.saveConfig("codex", configuration) + const stored = mockSecrets.store.mock.calls.at(-1)?.[1] + if (typeof stored !== "string") throw new Error("Profile was not stored") + mockSecrets.get.mockResolvedValue(stored) + expect(await providerSettingsManager.getProfile({ name: "codex" })).toMatchObject(configuration) + expect((await providerSettingsManager.getProfile({ name: "codex" })).openAiCodexUseWebSocket).toBe( + enabled, + ) + const exported = await providerSettingsManager.export() + expect(exported.apiConfigs.codex.openAiCodexUseWebSocket).toBe(enabled) + await providerSettingsManager.import(exported) + const imported = mockSecrets.store.mock.calls.at(-1)?.[1] + if (typeof imported !== "string") throw new Error("Profile was not imported") + mockSecrets.get.mockResolvedValue(imported) + expect((await providerSettingsManager.getProfile({ name: "codex" })).openAiCodexUseWebSocket).toBe( + enabled, + ) + }, + ) + it.each([OpenAiCodexServiceTier.Default, OpenAiCodexServiceTier.Priority] as const)( "should persist the OpenAI Codex %s speed preference", async (openAiCodexServiceTier) => { diff --git a/src/core/webview/__tests__/ClineProvider.spec.ts b/src/core/webview/__tests__/ClineProvider.spec.ts index 0e4aa4aa18..f9f1ca6eff 100644 --- a/src/core/webview/__tests__/ClineProvider.spec.ts +++ b/src/core/webview/__tests__/ClineProvider.spec.ts @@ -1945,6 +1945,20 @@ describe("ClineProvider", () => { expect(postedState.apiConfiguration).toMatchObject(expectedConfiguration) }) + test.each([true, false, undefined])( + "returns saved Codex WebSocket preference %s through the webview round trip", + async (enabled) => { + const configuration: ProviderSettings = { + apiProvider: providerIdentifiers.openaiCodex, + ...(enabled !== undefined ? { openAiCodexUseWebSocket: enabled } : {}), + } + await provider.contextProxy.setProviderSettings(configuration) + expect(provider.contextProxy.getProviderSettings().openAiCodexUseWebSocket).toBe(enabled) + expect((await provider.getState()).apiConfiguration.openAiCodexUseWebSocket).toBe(enabled) + expect((await provider.getStateToPostToWebview()).apiConfiguration.openAiCodexUseWebSocket).toBe(enabled) + }, + ) + test.each([true, false, undefined])( "returns saved OpenAI-compatible reasoning settings to the webview when enabled is %s", async (enableReasoningEffort) => { diff --git a/src/eslint-suppressions.json b/src/eslint-suppressions.json index 908159f7ab..fb3369f638 100644 --- a/src/eslint-suppressions.json +++ b/src/eslint-suppressions.json @@ -391,7 +391,7 @@ }, "api/providers/openai-codex.ts": { "@typescript-eslint/no-explicit-any": { - "count": 31 + "count": 28 } }, "api/providers/openai-native.ts": {