diff --git a/packages/types/src/__tests__/provider-settings.test.ts b/packages/types/src/__tests__/provider-settings.test.ts index b29a93ca3e..6a5294ab66 100644 --- a/packages/types/src/__tests__/provider-settings.test.ts +++ b/packages/types/src/__tests__/provider-settings.test.ts @@ -22,6 +22,39 @@ describe("provider settings discriminated union", () => { }) }) +describe("OpenAI-compatible extra body settings", () => { + it("accepts a JSON object with provider-specific nested fields", () => { + const settings = { + apiProvider: providerIdentifiers.openai, + openAiExtraBody: JSON.stringify({ metadata: { completion_window: "balanced" }, store: false }), + } + + expect(providerSettingsSchemaDiscriminated.parse(settings)).toEqual(settings) + expect(PROVIDER_SETTINGS_KEYS).toContain("openAiExtraBody") + }) + + it.each(["not json", "[]", "null", '"value"'])("rejects non-object JSON: %s", (openAiExtraBody) => { + expect( + providerSettingsSchemaDiscriminated.safeParse({ + apiProvider: providerIdentifiers.openai, + openAiExtraBody, + }).success, + ).toBe(false) + }) + + it.each(["model", "messages", "stream", "tools", "max_tokens", "__proto__"])( + "rejects the reserved request key %s", + (reservedKey) => { + expect( + providerSettingsSchemaDiscriminated.safeParse({ + apiProvider: providerIdentifiers.openai, + openAiExtraBody: JSON.stringify({ [reservedKey]: "override" }), + }).success, + ).toBe(false) + }, + ) +}) + describe("OpenAI Codex provider settings", () => { it("preserves the Fast preference in general and provider-specific schemas", () => { const settings = { diff --git a/packages/types/src/provider-settings.ts b/packages/types/src/provider-settings.ts index 4a432970c2..c09c9f8fe5 100644 --- a/packages/types/src/provider-settings.ts +++ b/packages/types/src/provider-settings.ts @@ -4,6 +4,8 @@ import { providerDefinitionList, type ProviderDefinition } from "./provider-sett import { API_PROVIDER_FIELD, SETTINGS_SHAPE_FIELD } from "./provider-settings/common.js" export { OPEN_AI_CODEX_SERVICE_TIER_KEY, + OPENAI_EXTRA_BODY_RESERVED_KEYS, + parseOpenAiExtraBody, kimiCodeAuthMethodSchema, type KimiCodeAuthMethod, nanoGptDefaultRoutingPreference, diff --git a/packages/types/src/provider-settings/index.ts b/packages/types/src/provider-settings/index.ts index 225b54fc29..17a3f36baf 100644 --- a/packages/types/src/provider-settings/index.ts +++ b/packages/types/src/provider-settings/index.ts @@ -37,6 +37,7 @@ import { basetenProviderDefinition } from "./baseten.js" import type { ProviderDefinition } from "./common.js" export { OPEN_AI_CODEX_SERVICE_TIER_KEY } from "./openai-codex.js" +export { OPENAI_EXTRA_BODY_RESERVED_KEYS, parseOpenAiExtraBody } from "./openai.js" export { kimiCodeAuthMethodSchema, type KimiCodeAuthMethod } from "./kimi-code.js" export { zaiApiLineSchema, type ZaiApiLine } from "./zai.js" export { diff --git a/packages/types/src/provider-settings/openai.ts b/packages/types/src/provider-settings/openai.ts index 651d90f60d..2176a981f6 100644 --- a/packages/types/src/provider-settings/openai.ts +++ b/packages/types/src/provider-settings/openai.ts @@ -6,6 +6,72 @@ import { baseProviderSettingsShape, createModelIdAccessor, createProviderDefinit export const OPEN_AI_MODEL_ID_FIELD = "openAiModelId" +export const OPENAI_EXTRA_BODY_RESERVED_KEYS = [ + "__proto__", + "constructor", + "max_completion_tokens", + "max_tokens", + "messages", + "model", + "parallel_tool_calls", + "prototype", + "reasoning", + "reasoning_effort", + "stream", + "stream_options", + "temperature", + "tool_choice", + "tools", +] as const + +type OpenAiExtraBodyParseResult = + | { success: true; data: Record } + | { + success: false + reason: "invalidJson" | "objectRequired" | "reservedKeys" + data: Record + reservedKeys?: string[] + } + +export function parseOpenAiExtraBody(value: string | undefined): OpenAiExtraBodyParseResult { + if (!value?.trim()) { + return { success: true, data: {} } + } + + let parsed: unknown + try { + parsed = JSON.parse(value) + } catch { + return { success: false, reason: "invalidJson", data: {} } + } + + if (parsed === null || typeof parsed !== "object" || Array.isArray(parsed)) { + return { success: false, reason: "objectRequired", data: {} } + } + + const entries = Object.entries(parsed) + const reservedKeys = entries + .map(([key]) => key) + .filter((key) => (OPENAI_EXTRA_BODY_RESERVED_KEYS as readonly string[]).includes(key)) + const data = Object.fromEntries(entries.filter(([key]) => !reservedKeys.includes(key))) + + if (reservedKeys.length > 0) { + return { success: false, reason: "reservedKeys", reservedKeys, data } + } + + return { success: true, data } +} + +const openAiExtraBodySchema = z + .string() + .superRefine((value, ctx) => { + const result = parseOpenAiExtraBody(value) + if (!result.success) { + ctx.addIssue({ code: "custom", message: result.reason }) + } + }) + .optional() + export const openAiProviderDefinition = createProviderDefinition({ apiProvider: providerIdentifiers.openai, modelIdKey: OPEN_AI_MODEL_ID_FIELD, @@ -22,5 +88,6 @@ export const openAiProviderDefinition = createProviderDefinition({ openAiStreamingEnabled: z.boolean().optional(), openAiHostHeader: z.string().optional(), // Keep temporarily for backward compatibility during migration. openAiHeaders: z.record(z.string(), z.string()).optional(), + openAiExtraBody: openAiExtraBodySchema, }, }) diff --git a/src/api/providers/__tests__/openai.spec.ts b/src/api/providers/__tests__/openai.spec.ts index 38550533a5..a41bd20648 100644 --- a/src/api/providers/__tests__/openai.spec.ts +++ b/src/api/providers/__tests__/openai.spec.ts @@ -261,6 +261,45 @@ describe("OpenAiHandler", () => { expect(textChunks[0].text).toBe("Test response") }) + it("adds Extra Body fields to streaming requests without allowing reserved field overrides", async () => { + const extraBodyHandler = new OpenAiHandler({ + ...mockOptions, + openAiExtraBody: JSON.stringify({ + metadata: { completion_window: "balanced" }, + model: "overridden-model", + messages: [], + stream: false, + }), + }) + + await collectStream(extraBodyHandler.createMessage(systemPrompt, messages)) + + expect(mockCreate).toHaveBeenCalledWith( + expect.objectContaining({ + metadata: { completion_window: "balanced" }, + model: mockOptions.openAiModelId, + stream: true, + messages: expect.arrayContaining([expect.objectContaining({ role: "user" })]), + }), + {}, + ) + }) + + it("adds Extra Body fields to non-streaming requests", async () => { + const extraBodyHandler = new OpenAiHandler({ + ...mockOptions, + openAiStreamingEnabled: false, + openAiExtraBody: JSON.stringify({ metadata: { completion_window: "balanced" } }), + }) + + await collectStream(extraBodyHandler.createMessage(systemPrompt, messages)) + + expect(mockCreate).toHaveBeenCalledWith( + expect.objectContaining({ metadata: { completion_window: "balanced" } }), + {}, + ) + }) + it("streams reasoning chunks from delta.reasoning_content", async () => { mockCreate.mockImplementationOnce(async () => asyncStreamFrom([ @@ -843,6 +882,20 @@ describe("OpenAiHandler", () => { ) }) + it("adds Extra Body fields to single-completion requests", async () => { + const extraBodyHandler = new OpenAiHandler({ + ...mockOptions, + openAiExtraBody: JSON.stringify({ metadata: { completion_window: "balanced" } }), + }) + + await extraBodyHandler.completePrompt("Test prompt") + + expect(mockCreate).toHaveBeenCalledWith( + expect.objectContaining({ metadata: { completion_window: "balanced" } }), + {}, + ) + }) + it("should handle API errors", async () => { mockCreate.mockRejectedValueOnce(new Error("API Error")) await expect(handler.completePrompt("Test prompt")).rejects.toThrow("OpenAI completion error: API Error") @@ -1103,6 +1156,21 @@ describe("OpenAiHandler", () => { ) }) + it.each([true, false])("adds Extra Body fields to O3 requests when streaming is %s", async (streaming) => { + const o3Handler = new OpenAiHandler({ + ...o3Options, + openAiStreamingEnabled: streaming, + openAiExtraBody: JSON.stringify({ metadata: { completion_window: "balanced" } }), + }) + + await collectStream(o3Handler.createMessage("system", [])) + + expect(mockCreate).toHaveBeenCalledWith( + expect.objectContaining({ metadata: { completion_window: "balanced" } }), + {}, + ) + }) + it("should handle tool calls with O3 model in streaming mode", async () => { const o3Handler = new OpenAiHandler(o3Options) diff --git a/src/api/providers/openai.ts b/src/api/providers/openai.ts index 5588dd37d6..5132b02450 100644 --- a/src/api/providers/openai.ts +++ b/src/api/providers/openai.ts @@ -10,6 +10,7 @@ import { openAiModelInfoSaneDefaults, DEEP_SEEK_DEFAULT_TEMPERATURE, OPENAI_AZURE_AI_INFERENCE_PATH, + parseOpenAiExtraBody, } from "@roo-code/types" import type { ApiHandlerOptions } from "../../shared/api" @@ -34,10 +35,12 @@ export class OpenAiHandler extends BaseProvider implements SingleCompletionHandl protected options: ApiHandlerOptions protected client: OpenAI private readonly providerName = "OpenAI" + private readonly extraBody: Record constructor(options: ApiHandlerOptions) { super() this.options = options + this.extraBody = parseOpenAiExtraBody(options.openAiExtraBody).data const baseURL = this.options.openAiBaseUrl || "https://api.openai.com/v1" const apiKey = this.options.openAiApiKey ?? NOT_PROVIDED @@ -152,7 +155,7 @@ export class OpenAiHandler extends BaseProvider implements SingleCompletionHandl const isGrokXAI = this._isGrokXAI(this.options.openAiBaseUrl) - const requestOptions: OpenAI.Chat.Completions.ChatCompletionCreateParamsStreaming = { + let requestOptions: OpenAI.Chat.Completions.ChatCompletionCreateParamsStreaming = { model: modelId, // Some OpenAI-Compatible models (e.g. claude-opus-4-7, claude-opus-4-8) reject // `temperature` as deprecated/unsupported, so honor the model's `supportsTemperature` @@ -174,6 +177,7 @@ export class OpenAiHandler extends BaseProvider implements SingleCompletionHandl // Add max_tokens if needed this.addMaxTokensIfNeeded(requestOptions, modelInfo) + requestOptions = this.withExtraBody(requestOptions) let stream try { @@ -227,7 +231,7 @@ export class OpenAiHandler extends BaseProvider implements SingleCompletionHandl yield this.processUsageMetrics(lastUsage, modelInfo) } } else { - const requestOptions: OpenAI.Chat.Completions.ChatCompletionCreateParamsNonStreaming = { + let requestOptions: OpenAI.Chat.Completions.ChatCompletionCreateParamsNonStreaming = { model: modelId, messages: deepseekReasoner ? convertToR1Format([{ role: "user", content: systemPrompt }, ...messages]) @@ -240,6 +244,7 @@ export class OpenAiHandler extends BaseProvider implements SingleCompletionHandl // Add max_tokens if needed this.addMaxTokensIfNeeded(requestOptions, modelInfo) + requestOptions = this.withExtraBody(requestOptions) let response try { @@ -304,13 +309,14 @@ export class OpenAiHandler extends BaseProvider implements SingleCompletionHandl const model = this.getModel() const modelInfo = model.info - const requestOptions: OpenAI.Chat.Completions.ChatCompletionCreateParamsNonStreaming = { + let requestOptions: OpenAI.Chat.Completions.ChatCompletionCreateParamsNonStreaming = { model: model.id, messages: [{ role: "user", content: prompt }], } // Add max_tokens if needed this.addMaxTokensIfNeeded(requestOptions, modelInfo) + requestOptions = this.withExtraBody(requestOptions) let response try { @@ -350,7 +356,7 @@ export class OpenAiHandler extends BaseProvider implements SingleCompletionHandl if (this.options.openAiStreamingEnabled ?? true) { const isGrokXAI = this._isGrokXAI(this.options.openAiBaseUrl) - const requestOptions: OpenAI.Chat.Completions.ChatCompletionCreateParamsStreaming = { + let requestOptions: OpenAI.Chat.Completions.ChatCompletionCreateParamsStreaming = { model: modelId, messages: [ { @@ -373,6 +379,7 @@ export class OpenAiHandler extends BaseProvider implements SingleCompletionHandl // but they do support max_completion_tokens (the modern OpenAI parameter) // This allows O3 models to limit response length when includeMaxTokens is enabled this.addMaxTokensIfNeeded(requestOptions, modelInfo) + requestOptions = this.withExtraBody(requestOptions) let stream try { @@ -386,7 +393,7 @@ export class OpenAiHandler extends BaseProvider implements SingleCompletionHandl yield* this.handleStreamResponse(stream) } else { - const requestOptions: OpenAI.Chat.Completions.ChatCompletionCreateParamsNonStreaming = { + let requestOptions: OpenAI.Chat.Completions.ChatCompletionCreateParamsNonStreaming = { model: modelId, messages: [ { @@ -407,6 +414,7 @@ export class OpenAiHandler extends BaseProvider implements SingleCompletionHandl // but they do support max_completion_tokens (the modern OpenAI parameter) // This allows O3 models to limit response length when includeMaxTokens is enabled this.addMaxTokensIfNeeded(requestOptions, modelInfo) + requestOptions = this.withExtraBody(requestOptions) let response try { @@ -543,6 +551,10 @@ export class OpenAiHandler extends BaseProvider implements SingleCompletionHandl requestOptions.max_completion_tokens = this.options.modelMaxTokens || modelInfo.maxTokens } } + + private withExtraBody(requestOptions: T): T { + return { ...this.extraBody, ...requestOptions } + } } export async function getOpenAiModels(baseUrl?: string, apiKey?: string, openAiHeaders?: Record) { diff --git a/src/core/config/__tests__/ProviderSettingsManager.spec.ts b/src/core/config/__tests__/ProviderSettingsManager.spec.ts index b7a0a9595c..798296eef5 100644 --- a/src/core/config/__tests__/ProviderSettingsManager.spec.ts +++ b/src/core/config/__tests__/ProviderSettingsManager.spec.ts @@ -484,6 +484,36 @@ describe("ProviderSettingsManager", () => { }, ) + it("persists OpenAI-compatible Extra Body only on OpenAI-compatible profiles", async () => { + mockSecrets.get.mockResolvedValue( + JSON.stringify({ + currentApiConfigName: "default", + apiConfigs: { default: {} }, + modeApiConfigs: {}, + }), + ) + const openAiExtraBody = JSON.stringify({ metadata: { completion_window: "balanced" } }) + + await providerSettingsManager.saveConfig("sail", { + apiProvider: providerIdentifiers.openai, + openAiModelId: "zai-org/GLM-5.2-FP8", + openAiExtraBody, + }) + + let storedProfiles = JSON.parse(mockSecrets.store.mock.calls.at(-1)?.[1]) + expect(storedProfiles.apiConfigs.sail).toMatchObject({ openAiExtraBody }) + + mockSecrets.get.mockResolvedValue(JSON.stringify(storedProfiles)) + await providerSettingsManager.saveConfig("anthropic", { + apiProvider: providerIdentifiers.anthropic, + apiKey: "test-key", + openAiExtraBody, + }) + + storedProfiles = JSON.parse(mockSecrets.store.mock.calls.at(-1)?.[1]) + expect(storedProfiles.apiConfigs.anthropic).not.toHaveProperty("openAiExtraBody") + }) + it("should only save provider relevant settings", async () => { mockSecrets.get.mockResolvedValue( JSON.stringify({ diff --git a/webview-ui/src/components/settings/__tests__/ApiOptions.spec.tsx b/webview-ui/src/components/settings/__tests__/ApiOptions.spec.tsx index c625650845..bd4fe95e0f 100644 --- a/webview-ui/src/components/settings/__tests__/ApiOptions.spec.tsx +++ b/webview-ui/src/components/settings/__tests__/ApiOptions.spec.tsx @@ -15,6 +15,9 @@ vi.mock("@vscode/webview-ui-toolkit/react", () => ({ ), + VSCodeTextArea: ({ value, onInput, ...props }: any) => ( +