diff --git a/webview-ui/src/components/settings/ApiOptions.tsx b/webview-ui/src/components/settings/ApiOptions.tsx index c5e69978ff..78161455c7 100644 --- a/webview-ui/src/components/settings/ApiOptions.tsx +++ b/webview-ui/src/components/settings/ApiOptions.tsx @@ -8,6 +8,7 @@ import { type ProviderName, type ProviderSettings, isRetiredProvider, + providerIdentifiers, DEFAULT_CONSECUTIVE_MISTAKE_LIMIT, } from "@roo-code/types" @@ -207,7 +208,7 @@ const ApiOptions = ({ // stops typing. useDebounce( () => { - if (selectedProvider === "openai") { + if (selectedProvider === providerIdentifiers.openai) { // Use our custom headers state to build the headers object. const headerObject = convertHeadersToObject(customHeaders) @@ -220,7 +221,7 @@ const ApiOptions = ({ openAiHeaders: headerObject, }, }) - } else if (selectedProvider === "ollama") { + } else if (selectedProvider === providerIdentifiers.ollama) { vscode.postMessage({ type: "requestOllamaModels", values: { @@ -228,11 +229,11 @@ const ApiOptions = ({ apiKey: apiConfiguration?.ollamaApiKey, }, }) - } else if (selectedProvider === "lmstudio") { + } else if (selectedProvider === providerIdentifiers.lmstudio) { requestLmStudioModels(apiConfiguration?.lmStudioBaseUrl) - } else if (selectedProvider === "vscode-lm") { + } else if (selectedProvider === providerIdentifiers.vscodeLm) { vscode.postMessage({ type: "requestVsCodeLmModels" }) - } else if (selectedProvider === "litellm") { + } else if (selectedProvider === providerIdentifiers.litellm) { vscode.postMessage({ type: "requestRouterModels", values: { @@ -240,7 +241,7 @@ const ApiOptions = ({ litellmBaseUrl: apiConfiguration?.litellmBaseUrl, }, }) - } else if (selectedProvider === "poe") { + } else if (selectedProvider === providerIdentifiers.poe) { vscode.postMessage({ type: "requestRouterModels" }) } }, @@ -270,7 +271,7 @@ const ApiOptions = ({ // Zoo Gateway renders its own auth-state error inline (sign-in card in // ZooGateway.tsx) so it can react to zooCodeIsAuthenticated changes // without re-running this effect or threading auth state through validation. - if (apiConfiguration.apiProvider === "zoo-gateway") { + if (apiConfiguration.apiProvider === providerIdentifiers.zooGateway) { setErrorMessage(undefined) return } @@ -322,7 +323,7 @@ const ApiOptions = ({ } // Bedrock has a special “custom-arn” pseudo-model that isn't part of MODELS_BY_PROVIDER. - if (provider === "bedrock" && modelId === "custom-arn") { + if (provider === providerIdentifiers.bedrock && modelId === "custom-arn") { return } @@ -441,7 +442,7 @@ const ApiOptions = ({ ) : ( <> - {selectedProvider === "openrouter" && ( + {selectedProvider === providerIdentifiers.openrouter && ( )} - {selectedProvider === "requesty" && ( + {selectedProvider === providerIdentifiers.requesty && ( )} - {selectedProvider === "unbound" && ( + {selectedProvider === providerIdentifiers.unbound && ( )} - {selectedProvider === "anthropic" && ( + {selectedProvider === providerIdentifiers.anthropic && ( )} - {selectedProvider === "openai-codex" && ( + {selectedProvider === providerIdentifiers.openaiCodex && ( )} - {selectedProvider === "openai-native" && ( + {selectedProvider === providerIdentifiers.openaiNative && ( )} - {selectedProvider === "mistral" && ( + {selectedProvider === providerIdentifiers.mistral && ( )} - {selectedProvider === "baseten" && ( + {selectedProvider === providerIdentifiers.baseten && ( )} - {selectedProvider === "bedrock" && ( + {selectedProvider === providerIdentifiers.bedrock && ( )} - {selectedProvider === "vertex" && ( + {selectedProvider === providerIdentifiers.vertex && ( )} - {selectedProvider === "gemini" && ( + {selectedProvider === providerIdentifiers.gemini && ( )} - {selectedProvider === "openai" && ( + {selectedProvider === providerIdentifiers.openai && ( )} - {selectedProvider === "lmstudio" && ( + {selectedProvider === providerIdentifiers.lmstudio && ( )} - {selectedProvider === "deepseek" && ( + {selectedProvider === providerIdentifiers.deepseek && ( )} - {selectedProvider === "qwen-code" && ( + {selectedProvider === providerIdentifiers.qwenCode && ( )} - {selectedProvider === "moonshot" && ( + {selectedProvider === providerIdentifiers.moonshot && ( )} - {selectedProvider === "kimi-code" && ( + {selectedProvider === providerIdentifiers.kimiCode && ( )} - {selectedProvider === "minimax" && ( + {selectedProvider === providerIdentifiers.minimax && ( )} - {selectedProvider === "mimo" && ( + {selectedProvider === providerIdentifiers.mimo && ( )} - {selectedProvider === "vscode-lm" && ( + {selectedProvider === providerIdentifiers.vscodeLm && ( )} - {selectedProvider === "ollama" && ( + {selectedProvider === providerIdentifiers.ollama && ( )} - {selectedProvider === "xai" && ( + {selectedProvider === providerIdentifiers.xai && ( )} - {selectedProvider === "litellm" && ( + {selectedProvider === providerIdentifiers.litellm && ( )} - {selectedProvider === "sambanova" && ( + {selectedProvider === providerIdentifiers.sambanova && ( )} - {selectedProvider === "zai" && ( + {selectedProvider === providerIdentifiers.zai && ( )} - {selectedProvider === "vercel-ai-gateway" && ( + {selectedProvider === providerIdentifiers.vercelAiGateway && ( )} - {selectedProvider === "opencode-go" && ( + {selectedProvider === providerIdentifiers.opencodeGo && ( )} - {selectedProvider === "kenari" && ( + {selectedProvider === providerIdentifiers.kenari && ( )} - {selectedProvider === "zoo-gateway" && ( + {selectedProvider === providerIdentifiers.zooGateway && ( )} - {selectedProvider === "fireworks" && ( + {selectedProvider === providerIdentifiers.fireworks && ( )} - {selectedProvider === "friendli" && ( + {selectedProvider === providerIdentifiers.friendli && ( )} - {selectedProvider === "poe" && ( + {selectedProvider === providerIdentifiers.poe && ( - {selectedProvider === "bedrock" && selectedModelId === "custom-arn" && ( + {selectedProvider === providerIdentifiers.bedrock && selectedModelId === "custom-arn" && ( setApiConfigurationField("consecutiveMistakeLimit", value)} /> - {selectedProvider === "poe" && ( + {selectedProvider === providerIdentifiers.poe && ( )} - {selectedProvider === "openrouter" && + {selectedProvider === providerIdentifiers.openrouter && openRouterModelProviders && Object.keys(openRouterModelProviders).length > 0 && (
diff --git a/webview-ui/src/components/settings/__tests__/ApiOptions.interactions.spec.tsx b/webview-ui/src/components/settings/__tests__/ApiOptions.interactions.spec.tsx new file mode 100644 index 0000000000..02c6c1af02 --- /dev/null +++ b/webview-ui/src/components/settings/__tests__/ApiOptions.interactions.spec.tsx @@ -0,0 +1,388 @@ +import { act, fireEvent, render, screen, within } from "@/utils/test-utils" +import { bedrockDefaultModelId, providerIdentifiers, type ProviderSettings } from "@roo-code/types" +import type { ChangeEventHandler, InputHTMLAttributes, ReactNode } from "react" + +import { requestLmStudioModels } from "@src/components/ui/hooks/useLmStudioModels" +import type { useOpenRouterModelProviders } from "@src/components/ui/hooks/useOpenRouterModelProviders" +import { vscode } from "@src/utils/vscode" + +import ApiOptions, { type ApiOptionsProps } from "../ApiOptions" + +type OpenRouterModelProvidersQueryResult = Pick, "data"> + +const { useOpenRouterModelProvidersMock } = vi.hoisted(() => ({ + useOpenRouterModelProvidersMock: vi.fn<() => OpenRouterModelProvidersQueryResult>(() => ({ data: undefined })), +})) + +type ChildrenProps = { children?: ReactNode } + +type VSCodeTextFieldMockProps = ChildrenProps & + Pick, "value" | "placeholder"> & { + onInput?: ChangeEventHandler + } + +type SearchableSelectMockProps = { + value?: string + onValueChange: (value: string) => void + options: Array<{ value: string; label: string }> + "data-testid"?: string +} + +vi.mock("@src/context/ExtensionStateContext", () => ({ + useExtensionState: () => ({ + organizationAllowList: { allowAll: true, providers: {} }, + openAiCodexIsAuthenticated: false, + kimiCodeIsAuthenticated: false, + kimiCodeOAuthState: undefined, + }), +})) + +vi.mock("@src/components/ui/hooks/useRouterModels", () => ({ + useRouterModels: () => ({ data: {}, refetch: vi.fn() }), +})) + +vi.mock("@src/components/ui/hooks/useZooGatewayRouterModelsSync", () => ({ + useZooGatewayRouterModelsSync: vi.fn(), +})) + +vi.mock("@src/components/ui/hooks/useOpenRouterModelProviders", () => ({ + useOpenRouterModelProviders: useOpenRouterModelProvidersMock, + OPENROUTER_DEFAULT_PROVIDER_NAME: "Auto", +})) + +vi.mock("@src/components/ui/hooks/useSelectedModel", () => ({ + useSelectedModel: (configuration: ProviderSettings) => ({ + provider: configuration.apiProvider, + id: configuration.apiModelId, + info: {}, + }), +})) + +vi.mock("@src/components/ui/hooks/useLmStudioModels", () => ({ + requestLmStudioModels: vi.fn(), +})) + +vi.mock("../providers", () => { + const Provider = () => null + return { + Anthropic: Provider, + Baseten: Provider, + Bedrock: Provider, + DeepSeek: Provider, + Gemini: Provider, + LMStudio: Provider, + LiteLLM: Provider, + Mistral: Provider, + Moonshot: Provider, + KimiCode: Provider, + Ollama: Provider, + OpenAI: Provider, + OpenAICompatible: Provider, + OpenAICodex: Provider, + OpenRouter: Provider, + Poe: Provider, + QwenCode: Provider, + Requesty: Provider, + SambaNova: Provider, + Unbound: Provider, + Vertex: Provider, + VSCodeLM: Provider, + XAI: Provider, + ZAi: Provider, + Fireworks: Provider, + Friendli: Provider, + VercelAiGateway: Provider, + OpenCodeGo: Provider, + Kenari: Provider, + ZooGateway: Provider, + MiniMax: Provider, + Mimo: Provider, + } +}) + +vi.mock("../providers/BedrockCustomArn", () => ({ + BedrockCustomArn: () =>
, +})) +vi.mock("../ModelPicker", () => ({ ModelPicker: () => null })) +vi.mock("../ApiErrorMessage", () => ({ ApiErrorMessage: () => null })) +vi.mock("../ThinkingBudget", () => ({ ThinkingBudget: () => null })) +vi.mock("../Verbosity", () => ({ Verbosity: () => null })) +vi.mock("../TodoListSettingsControl", () => ({ TodoListSettingsControl: () => null })) +vi.mock("../TemperatureControl", () => ({ TemperatureControl: () => null })) +vi.mock("../RateLimitSecondsControl", () => ({ RateLimitSecondsControl: () => null })) +vi.mock("../ConsecutiveMistakeLimitControl", () => ({ + ConsecutiveMistakeLimitControl: ({ value, onChange }: { value: number; onChange: (value: number) => void }) => ( +
+ onChange(Number(event.target.value))} /> +
+ ), +})) + +vi.mock("@vscode/webview-ui-toolkit/react", () => ({ + VSCodeTextField: ({ children, value, onInput, placeholder }: VSCodeTextFieldMockProps) => ( + + ), + VSCodeLink: ({ children }: ChildrenProps) => {children}, +})) + +vi.mock("@/components/ui", () => ({ + SearchableSelect: ({ value, onValueChange, options, "data-testid": testId }: SearchableSelectMockProps) => ( +
+ +
+ ), + Collapsible: ({ children }: ChildrenProps) =>
{children}
, + CollapsibleTrigger: ({ children }: ChildrenProps) =>
{children}
, + CollapsibleContent: ({ children }: ChildrenProps) =>
{children}
, + Select: ({ children }: ChildrenProps) =>
{children}
, + SelectTrigger: ({ children }: ChildrenProps) =>
{children}
, + SelectValue: () => null, + SelectContent: ({ children }: ChildrenProps) =>
{children}
, + SelectItem: ({ children }: ChildrenProps) =>
{children}
, +})) + +const renderApiOptions = (props: Partial = {}) => + render( + undefined} + uriScheme={undefined} + apiConfiguration={{}} + setApiConfigurationField={() => undefined} + {...props} + />, + ) + +describe("ApiOptions interactions", () => { + afterEach(() => { + vi.useRealTimers() + vi.restoreAllMocks() + }) + + describe("debounced provider model refresh", () => { + it.each([ + { + provider: providerIdentifiers.openai, + configuration: { + openAiBaseUrl: "https://openai.example/v1", + openAiApiKey: "openai-key", + openAiHeaders: { "X-Custom": "header-value" }, + }, + expectedMessage: { + type: "requestOpenAiModels", + values: { + baseUrl: "https://openai.example/v1", + apiKey: "openai-key", + customHeaders: {}, + openAiHeaders: { "X-Custom": "header-value" }, + }, + }, + }, + { + provider: providerIdentifiers.ollama, + configuration: { ollamaBaseUrl: "http://ollama:11434", ollamaApiKey: "ollama-key" }, + expectedMessage: { + type: "requestOllamaModels", + values: { baseUrl: "http://ollama:11434", apiKey: "ollama-key" }, + }, + }, + { + provider: providerIdentifiers.vscodeLm, + configuration: {}, + expectedMessage: { type: "requestVsCodeLmModels" }, + }, + { + provider: providerIdentifiers.litellm, + configuration: { litellmBaseUrl: "http://litellm:4000", litellmApiKey: "litellm-key" }, + expectedMessage: { + type: "requestRouterModels", + values: { litellmApiKey: "litellm-key", litellmBaseUrl: "http://litellm:4000" }, + }, + }, + { + provider: providerIdentifiers.poe, + configuration: { poeApiKey: "poe-key", poeBaseUrl: "https://api.poe.example/v1" }, + expectedMessage: { type: "requestRouterModels" }, + }, + ])("requests models for $provider", ({ provider, configuration, expectedMessage }) => { + vi.useFakeTimers() + const postMessage = vi.spyOn(vscode, "postMessage").mockImplementation(() => undefined) + + renderApiOptions({ apiConfiguration: { apiProvider: provider, ...configuration } }) + act(() => vi.advanceTimersByTime(249)) + expect(postMessage).not.toHaveBeenCalledWith(expectedMessage) + + act(() => vi.advanceTimersByTime(1)) + expect(postMessage).toHaveBeenCalledTimes(1) + expect(postMessage).toHaveBeenCalledWith(expectedMessage) + }) + + it("requests LM Studio models using its configured base URL", () => { + vi.useFakeTimers() + renderApiOptions({ + apiConfiguration: { + apiProvider: providerIdentifiers.lmstudio, + lmStudioBaseUrl: "http://lmstudio:1234", + }, + }) + + act(() => vi.advanceTimersByTime(249)) + expect(requestLmStudioModels).not.toHaveBeenCalledWith("http://lmstudio:1234") + + act(() => vi.advanceTimersByTime(1)) + expect(requestLmStudioModels).toHaveBeenCalledTimes(1) + expect(requestLmStudioModels).toHaveBeenCalledWith("http://lmstudio:1234") + }) + + it("does not request dynamic models for a static provider", () => { + vi.useFakeTimers() + const postMessage = vi.spyOn(vscode, "postMessage").mockImplementation(() => undefined) + + renderApiOptions({ apiConfiguration: { apiProvider: providerIdentifiers.anthropic } }) + act(() => vi.advanceTimersByTime(250)) + + expect(postMessage).not.toHaveBeenCalled() + }) + }) + + it.each([ + providerIdentifiers.requesty, + providerIdentifiers.unbound, + providerIdentifiers.anthropic, + providerIdentifiers.openaiCodex, + providerIdentifiers.openaiNative, + providerIdentifiers.mistral, + providerIdentifiers.baseten, + providerIdentifiers.bedrock, + providerIdentifiers.gemini, + providerIdentifiers.lmstudio, + providerIdentifiers.deepseek, + providerIdentifiers.qwenCode, + providerIdentifiers.moonshot, + providerIdentifiers.kimiCode, + providerIdentifiers.minimax, + providerIdentifiers.mimo, + providerIdentifiers.ollama, + providerIdentifiers.litellm, + providerIdentifiers.sambanova, + providerIdentifiers.zai, + providerIdentifiers.xai, + providerIdentifiers.fireworks, + providerIdentifiers.friendli, + providerIdentifiers.vercelAiGateway, + providerIdentifiers.opencodeGo, + ])("renders the canonical %s provider branch", (apiProvider) => { + const { unmount } = renderApiOptions({ apiConfiguration: { apiProvider } }) + unmount() + }) + + it("clears parent validation errors for Zoo Gateway", () => { + const setErrorMessage = vi.fn() + renderApiOptions({ apiConfiguration: { apiProvider: providerIdentifiers.zooGateway }, setErrorMessage }) + + expect(setErrorMessage).toHaveBeenCalledWith(undefined) + }) + + it("renders OpenRouter provider routing when provider metadata is available", () => { + useOpenRouterModelProvidersMock.mockReturnValue({ + data: { preferred: { label: "Preferred", contextWindow: 1, supportsPromptCache: false } }, + }) + + renderApiOptions({ + apiConfiguration: { + apiProvider: providerIdentifiers.openrouter, + openRouterModelId: "anthropic/claude-sonnet-4.5", + }, + }) + + expect(screen.getByText("settings:providers.openRouter.providerRouting.title")).toBeInTheDocument() + }) + + it("preserves the Bedrock custom ARN pseudo-model when switching to Bedrock", () => { + const setApiConfigurationField = vi.fn() + renderApiOptions({ + apiConfiguration: { apiProvider: providerIdentifiers.anthropic, apiModelId: "custom-arn" }, + setApiConfigurationField, + }) + + const providerSelect = screen.getByTestId("provider-select").querySelector("select") as HTMLSelectElement + fireEvent.change(providerSelect, { target: { value: providerIdentifiers.bedrock } }) + + expect(setApiConfigurationField).toHaveBeenCalledWith("apiProvider", providerIdentifiers.bedrock) + expect(setApiConfigurationField.mock.calls.filter(([field]) => field === "apiModelId")).toEqual([]) + }) + + it("resets an invalid ordinary model to the Bedrock default when switching providers", () => { + const setApiConfigurationField = vi.fn() + renderApiOptions({ + apiConfiguration: { apiProvider: providerIdentifiers.anthropic, apiModelId: "not-a-bedrock-model" }, + setApiConfigurationField, + }) + + const providerSelect = screen.getByTestId("provider-select").querySelector("select") as HTMLSelectElement + fireEvent.change(providerSelect, { target: { value: providerIdentifiers.bedrock } }) + + expect(setApiConfigurationField).toHaveBeenCalledWith("apiProvider", providerIdentifiers.bedrock) + expect(setApiConfigurationField).toHaveBeenCalledWith("apiModelId", bedrockDefaultModelId, false) + }) + + it("renders the custom ARN settings only for Bedrock's custom ARN pseudo-model", () => { + const { rerender } = render( + undefined} + uriScheme={undefined} + apiConfiguration={{ apiProvider: providerIdentifiers.bedrock, apiModelId: "custom-arn" }} + setApiConfigurationField={() => undefined} + />, + ) + + expect(screen.getByTestId("bedrock-custom-arn")).toBeInTheDocument() + + rerender( + undefined} + uriScheme={undefined} + apiConfiguration={{ apiProvider: providerIdentifiers.bedrock, apiModelId: bedrockDefaultModelId }} + setApiConfigurationField={() => undefined} + />, + ) + + expect(screen.queryByTestId("bedrock-custom-arn")).not.toBeInTheDocument() + }) + + it("updates the consecutive mistake limit from advanced settings", () => { + const setApiConfigurationField = vi.fn() + renderApiOptions({ apiConfiguration: {}, setApiConfigurationField }) + + fireEvent.change(within(screen.getByTestId("consecutive-mistake-limit-control")).getByRole("slider"), { + target: { value: "7" }, + }) + + expect(setApiConfigurationField).toHaveBeenCalledWith("consecutiveMistakeLimit", 7) + }) + + it("renders and updates the Poe base URL in advanced settings", () => { + const setApiConfigurationField = vi.fn() + renderApiOptions({ + apiConfiguration: { apiProvider: providerIdentifiers.poe, poeBaseUrl: "https://api.poe.example/v1" }, + setApiConfigurationField, + }) + + const poeBaseUrl = screen.getByPlaceholderText("https://api.poe.com/v1") + expect(poeBaseUrl).toHaveValue("https://api.poe.example/v1") + + fireEvent.change(poeBaseUrl, { target: { value: "https://new.poe.example/v1" } }) + expect(setApiConfigurationField).toHaveBeenCalledWith("poeBaseUrl", "https://new.poe.example/v1") + }) +})