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
33 changes: 33 additions & 0 deletions packages/types/src/__tests__/provider-settings.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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 = {
Expand Down
2 changes: 2 additions & 0 deletions packages/types/src/provider-settings.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
1 change: 1 addition & 0 deletions packages/types/src/provider-settings/index.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down
67 changes: 67 additions & 0 deletions packages/types/src/provider-settings/openai.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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<string, unknown> }
| {
success: false
reason: "invalidJson" | "objectRequired" | "reservedKeys"
data: Record<string, unknown>
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,
Expand All @@ -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,
},
})
68 changes: 68 additions & 0 deletions src/api/providers/__tests__/openai.spec.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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([
Expand Down Expand Up @@ -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")
Expand Down Expand Up @@ -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)

Expand Down
22 changes: 17 additions & 5 deletions src/api/providers/openai.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand All @@ -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<string, unknown>

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
Expand Down Expand Up @@ -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`
Expand All @@ -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 {
Expand Down Expand Up @@ -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])
Expand All @@ -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 {
Expand Down Expand Up @@ -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 {
Expand Down Expand Up @@ -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: [
{
Expand All @@ -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 {
Expand All @@ -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: [
{
Expand All @@ -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 {
Expand Down Expand Up @@ -543,6 +551,10 @@ export class OpenAiHandler extends BaseProvider implements SingleCompletionHandl
requestOptions.max_completion_tokens = this.options.modelMaxTokens || modelInfo.maxTokens
}
}

private withExtraBody<T extends object>(requestOptions: T): T {
return { ...this.extraBody, ...requestOptions }
}
}

export async function getOpenAiModels(baseUrl?: string, apiKey?: string, openAiHeaders?: Record<string, string>) {
Expand Down
30 changes: 30 additions & 0 deletions src/core/config/__tests__/ProviderSettingsManager.spec.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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({
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,9 @@ vi.mock("@vscode/webview-ui-toolkit/react", () => ({
<input type="text" value={value} onChange={onBlur} />
</div>
),
VSCodeTextArea: ({ value, onInput, ...props }: any) => (
<textarea value={value} onChange={(event) => onInput?.(event)} {...props} />
),
VSCodeLink: ({ children, href }: any) => <a href={href}>{children}</a>,
VSCodeRadio: ({ value, checked }: any) => <input type="radio" value={value} checked={checked} />,
VSCodeRadioGroup: ({ children }: any) => <div>{children}</div>,
Expand Down
Loading
Loading