Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
20 changes: 20 additions & 0 deletions packages/types/src/__tests__/provider-settings.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand Down
1 change: 1 addition & 0 deletions packages/types/src/provider-settings.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
2 changes: 1 addition & 1 deletion packages/types/src/provider-settings/index.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down
5 changes: 5 additions & 0 deletions packages/types/src/provider-settings/openai-codex.ts
Original file line number Diff line number Diff line change
@@ -1,3 +1,5 @@
import { z } from "zod"

import { providerIdentifiers } from "../provider-identifiers.js"
import { openAiCodexServiceTierSchema } from "../model.js"
import {
Expand All @@ -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({
Expand All @@ -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(),
},
Expand Down
107 changes: 107 additions & 0 deletions src/api/providers/__tests__/openai-codex.spec.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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"

Expand All @@ -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",
Expand Down
39 changes: 32 additions & 7 deletions src/api/providers/openai-codex.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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"
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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()
}
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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<any>

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.",
)
Expand Down Expand Up @@ -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)
Expand Down
35 changes: 35 additions & 0 deletions src/core/config/__tests__/ProviderSettingsManager.spec.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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) => {
Expand Down
14 changes: 14 additions & 0 deletions src/core/webview/__tests__/ClineProvider.spec.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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) => {
Expand Down
2 changes: 1 addition & 1 deletion src/eslint-suppressions.json
Original file line number Diff line number Diff line change
Expand Up @@ -391,7 +391,7 @@
},
"api/providers/openai-codex.ts": {
"@typescript-eslint/no-explicit-any": {
"count": 31
"count": 28
}
},
"api/providers/openai-native.ts": {
Expand Down
Loading