From a38deb95f4cf7e1293e6c34675a1dd8f7746f9f2 Mon Sep 17 00:00:00 2001 From: gubin-dev Date: Wed, 5 Aug 2026 15:10:21 +0300 Subject: [PATCH] refactor(webview): canonicalize provider settings identifiers --- .../components/settings/providers/Kenari.tsx | 3 +- .../settings/providers/KimiCode.tsx | 7 +- .../components/settings/providers/LiteLLM.tsx | 14 +- .../settings/providers/Moonshot.tsx | 13 +- .../settings/providers/OpenCodeGo.tsx | 11 +- .../src/components/settings/providers/Poe.tsx | 6 +- .../settings/providers/Requesty.tsx | 8 +- .../components/settings/providers/Unbound.tsx | 6 +- .../settings/providers/VercelAiGateway.tsx | 3 +- .../settings/providers/ZooGateway.tsx | 3 +- .../providers/__tests__/KimiCode.spec.tsx | 75 ++++++++- .../providers/__tests__/LiteLLM.spec.tsx | 145 ++++++++++++++++++ .../providers/__tests__/Moonshot.spec.tsx | 93 +++++++++-- .../settings/providers/__tests__/Poe.spec.tsx | 115 ++++++++++++++ .../__tests__/ProviderRouting.spec.tsx | 82 ++++++++++ .../providers/__tests__/Requesty.spec.tsx | 57 +++++++ 16 files changed, 599 insertions(+), 42 deletions(-) create mode 100644 webview-ui/src/components/settings/providers/__tests__/LiteLLM.spec.tsx create mode 100644 webview-ui/src/components/settings/providers/__tests__/Poe.spec.tsx create mode 100644 webview-ui/src/components/settings/providers/__tests__/ProviderRouting.spec.tsx create mode 100644 webview-ui/src/components/settings/providers/__tests__/Requesty.spec.tsx diff --git a/webview-ui/src/components/settings/providers/Kenari.tsx b/webview-ui/src/components/settings/providers/Kenari.tsx index 577a44ac71..e8d2f5cdb9 100644 --- a/webview-ui/src/components/settings/providers/Kenari.tsx +++ b/webview-ui/src/components/settings/providers/Kenari.tsx @@ -6,6 +6,7 @@ import { type OrganizationAllowList, type RouterModels, kenariDefaultModelId, + providerIdentifiers, } from "@roo-code/types" import { useAppTranslation } from "@src/i18n/TranslationContext" @@ -69,7 +70,7 @@ export const Kenari = ({ apiConfiguration={apiConfiguration} setApiConfigurationField={setApiConfigurationField} defaultModelId={kenariDefaultModelId} - models={routerModels?.["kenari"] ?? {}} + models={routerModels?.[providerIdentifiers.kenari] ?? {}} modelIdKey="kenariModelId" serviceName="Kenari" serviceUrl="https://kenari.id/docs" diff --git a/webview-ui/src/components/settings/providers/KimiCode.tsx b/webview-ui/src/components/settings/providers/KimiCode.tsx index 4e9d3ff561..de060ad5c5 100644 --- a/webview-ui/src/components/settings/providers/KimiCode.tsx +++ b/webview-ui/src/components/settings/providers/KimiCode.tsx @@ -7,6 +7,7 @@ import { type KimiCodeAuthMethod, type ModelRecord, type ProviderSettings, + providerIdentifiers, } from "@roo-code/types" import { useAppTranslation } from "@src/i18n/TranslationContext" @@ -37,10 +38,10 @@ export const KimiCode = ({ const { t } = useAppTranslation() const authMethod = apiConfiguration.kimiCodeAuthMethod ?? "oauth" const { data, refetch, isFetching } = useRouterModels({ - provider: "kimi-code", + provider: providerIdentifiers.kimiCode, enabled: authMethod === "oauth" ? kimiCodeIsAuthenticated : !!apiConfiguration.kimiCodeApiKey, }) - const discoveredModels = data?.["kimi-code"] + const discoveredModels = data?.[providerIdentifiers.kimiCode] const models: ModelRecord = discoveredModels && Object.keys(discoveredModels).length > 0 ? discoveredModels : kimiCodeModels @@ -52,7 +53,7 @@ export const KimiCode = ({ vscode.postMessage({ type: "requestRouterModels", values: { - provider: "kimi-code", + provider: providerIdentifiers.kimiCode, refresh: true, kimiCodeAuthMethod: authMethod, kimiCodeApiKey: apiConfiguration.kimiCodeApiKey, diff --git a/webview-ui/src/components/settings/providers/LiteLLM.tsx b/webview-ui/src/components/settings/providers/LiteLLM.tsx index 2a8dcb8d67..9688a48184 100644 --- a/webview-ui/src/components/settings/providers/LiteLLM.tsx +++ b/webview-ui/src/components/settings/providers/LiteLLM.tsx @@ -7,6 +7,7 @@ import { type OrganizationAllowList, type ExtensionMessage, litellmDefaultModelId, + providerIdentifiers, } from "@roo-code/types" import { RouterName } from "@roo/api" @@ -46,7 +47,7 @@ export const LiteLLM = ({ const message = event.data if (message.type === "singleRouterModelFetchResponse" && !message.success) { const providerName = message.values?.provider as RouterName - if (providerName === "litellm") { + if (providerName === providerIdentifiers.litellm) { litellmErrorJustReceived.current = true setRefreshStatus("error") setRefreshError(message.error) @@ -57,12 +58,11 @@ export const LiteLLM = ({ if (refreshStatus === "loading") { if (!litellmErrorJustReceived.current) { setRefreshStatus("success") - // Invalidate only the LiteLLM router-models query so useSelectedModel - // picks up the refreshed list. useSelectedModel reads LiteLLM under the - // compound key ["routerModels", "litellm"] (see useRouterModels), so we - // target that exact key rather than the bare ["routerModels"] prefix, - // which would needlessly invalidate every other provider's query too. - queryClient.invalidateQueries({ queryKey: ["routerModels", "litellm"] }) + // Refresh the provider-scoped cache used by useSelectedModel and the shared cache used by + // ApiOptions. Target both exact keys rather than the bare ["routerModels"] prefix, which + // would needlessly invalidate every other provider's query too. + void queryClient.invalidateQueries({ queryKey: ["routerModels", providerIdentifiers.litellm] }) + void queryClient.invalidateQueries({ queryKey: ["routerModels", "all"] }) } // If litellmErrorJustReceived.current is true, status is already (or will be) "error". } diff --git a/webview-ui/src/components/settings/providers/Moonshot.tsx b/webview-ui/src/components/settings/providers/Moonshot.tsx index 2d6c7d849a..ee974c2502 100644 --- a/webview-ui/src/components/settings/providers/Moonshot.tsx +++ b/webview-ui/src/components/settings/providers/Moonshot.tsx @@ -2,8 +2,12 @@ import { useCallback, useState, useEffect, useRef } from "react" import { VSCodeTextField, VSCodeDropdown, VSCodeOption } from "@vscode/webview-ui-toolkit/react" import { useQueryClient } from "@tanstack/react-query" -import type { ProviderSettings, ExtensionMessage } from "@roo-code/types" -import { moonshotDefaultModelId } from "@roo-code/types" +import { + type ProviderSettings, + type ExtensionMessage, + moonshotDefaultModelId, + providerIdentifiers, +} from "@roo-code/types" import { RouterName } from "@roo/api" @@ -14,7 +18,6 @@ import { vscode } from "@src/utils/vscode" import { Button } from "@src/components/ui" import { ModelPicker } from "../ModelPicker" import { handleModelChangeSideEffects } from "../utils/providerModelConfig" -import type { ProviderName } from "@roo-code/types" import { inputEventTransform } from "../transforms" @@ -37,7 +40,7 @@ export const Moonshot = ({ apiConfiguration, setApiConfigurationField, simplifyS const message = event.data if (message.type === "singleRouterModelFetchResponse" && !message.success) { const providerName = message.values?.provider as RouterName - if (providerName === "moonshot" && refreshStatus === "loading") { + if (providerName === providerIdentifiers.moonshot && refreshStatus === "loading") { moonshotErrorJustReceived.current = true setRefreshStatus("error") setRefreshError(message.error) @@ -138,7 +141,7 @@ export const Moonshot = ({ apiConfiguration, setApiConfigurationField, simplifyS serviceUrl="https://platform.moonshot.ai" simplifySettings={simplifySettings} onModelChange={(modelId) => - handleModelChangeSideEffects("moonshot" as ProviderName, modelId, setApiConfigurationField) + handleModelChangeSideEffects(providerIdentifiers.moonshot, modelId, setApiConfigurationField) } /> + ), +})) + +vi.mock("../../ModelPicker", () => ({ + ModelPicker: () =>
, +})) + +describe("LiteLLM", () => { + const organizationAllowList: OrganizationAllowList = { allowAll: true, providers: {} } + + beforeEach(() => { + vi.clearAllMocks() + mockUseExtensionState.mockReturnValue({ routerModels: { [providerIdentifiers.litellm]: {} } }) + }) + + it("invalidates both LiteLLM caches after a successful model refresh", async () => { + const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } }) + const invalidateQueries = vi.spyOn(queryClient, "invalidateQueries") + const apiConfiguration: ProviderSettings = { + apiProvider: providerIdentifiers.litellm, + litellmApiKey: "test-key", + litellmBaseUrl: "http://localhost:4000", + } + + render( + + + , + ) + + fireEvent.click(screen.getByTestId("refresh-button")) + act(() => { + window.dispatchEvent(new MessageEvent("message", { data: { type: "routerModels" } })) + }) + + await waitFor(() => { + expect(invalidateQueries).toHaveBeenCalledWith({ + queryKey: ["routerModels", providerIdentifiers.litellm], + }) + expect(invalidateQueries).toHaveBeenCalledWith({ queryKey: ["routerModels", "all"] }) + }) + }) + + it("recognizes failed refresh responses for the canonical LiteLLM provider", () => { + const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } }) + render( + + + , + ) + + act(() => { + window.dispatchEvent( + new MessageEvent("message", { + data: { + type: "singleRouterModelFetchResponse", + success: false, + values: { provider: providerIdentifiers.litellm }, + error: "LiteLLM unavailable", + }, + }), + ) + }) + + expect(screen.getByText("LiteLLM unavailable")).toBeInTheDocument() + }) + + it("ignores failed refresh responses for another provider", () => { + const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } }) + render( + + + , + ) + + fireEvent.click(screen.getByTestId("refresh-button")) + expect(screen.getByText("settings:providers.refreshModels.loading")).toBeInTheDocument() + + act(() => { + window.dispatchEvent( + new MessageEvent("message", { + data: { + type: "singleRouterModelFetchResponse", + success: false, + values: { provider: providerIdentifiers.openrouter }, + error: "OpenRouter unavailable", + }, + }), + ) + }) + + expect(screen.queryByText("OpenRouter unavailable")).not.toBeInTheDocument() + expect(screen.getByText("settings:providers.refreshModels.loading")).toBeInTheDocument() + }) +}) diff --git a/webview-ui/src/components/settings/providers/__tests__/Moonshot.spec.tsx b/webview-ui/src/components/settings/providers/__tests__/Moonshot.spec.tsx index 1ffcf9b683..3265aef30d 100644 --- a/webview-ui/src/components/settings/providers/__tests__/Moonshot.spec.tsx +++ b/webview-ui/src/components/settings/providers/__tests__/Moonshot.spec.tsx @@ -2,7 +2,7 @@ import React from "react" import { render, screen, fireEvent, waitFor, act } from "@/utils/test-utils" -import type { ProviderSettings } from "@roo-code/types" +import { providerIdentifiers, type ProviderSettings } from "@roo-code/types" import { Moonshot } from "../Moonshot" @@ -47,13 +47,18 @@ vi.mock("@vscode/webview-ui-toolkit/react", async (importOriginal) => { }) // Mock the ModelPicker - must be a simple component that doesn't import anything -vi.mock("../ModelPicker", () => ({ - ModelPicker: function MockModelPicker() { +vi.mock("../../ModelPicker", () => ({ + ModelPicker: function MockModelPicker({ onModelChange }: { onModelChange?: (modelId: string) => void }) { return React.createElement( "div", { "data-testid": "model-picker" }, React.createElement("span", { "data-testid": "model-picker-default" }, "mock-default"), React.createElement("span", { "data-testid": "model-picker-count" }, "0"), + React.createElement( + "button", + { "data-testid": "change-model", onClick: () => onModelChange?.("moonshot-v1-128k") }, + "Change model", + ), ) }, })) @@ -103,11 +108,6 @@ vi.mock("@src/components/common/VSCodeButtonLink", () => ({ ), })) -// Mock handleModelChangeSideEffects -vi.mock("../utils/providerModelConfig", () => ({ - handleModelChangeSideEffects: vi.fn(), -})) - import { useExtensionState } from "@src/context/ExtensionStateContext" import { vscode } from "@src/utils/vscode" @@ -117,7 +117,7 @@ describe("Moonshot Component", () => { const mockSetApiConfigurationField = vi.fn() const createDefaultApiConfiguration = (overrides?: Partial): ProviderSettings => ({ - apiProvider: "moonshot", + apiProvider: providerIdentifiers.moonshot, moonshotBaseUrl: "https://api.moonshot.ai/v1", ...overrides, }) @@ -258,7 +258,7 @@ describe("Moonshot Component", () => { type: "singleRouterModelFetchResponse", success: false, error: "API connection failed", - values: { provider: "moonshot" }, + values: { provider: providerIdentifiers.moonshot }, }, "*", ) @@ -274,6 +274,77 @@ describe("Moonshot Component", () => { }) }) + it("ignores another provider's failed refresh response while loading", async () => { + render( + , + ) + + const refreshButton = screen + .getAllByTestId("button") + .find((button) => button.getAttribute("data-variant") === "outline")! + fireEvent.click(refreshButton) + await waitFor(() => expect(screen.getByText("settings:providers.refreshModels.loading")).toBeInTheDocument()) + + act(() => { + window.dispatchEvent( + new MessageEvent("message", { + data: { + type: "singleRouterModelFetchResponse", + success: false, + error: "OpenRouter unavailable", + values: { provider: providerIdentifiers.openrouter }, + }, + }), + ) + }) + + expect(screen.queryByText("OpenRouter unavailable")).not.toBeInTheDocument() + expect(screen.getByText("settings:providers.refreshModels.loading")).toBeInTheDocument() + }) + + it("ignores a Moonshot failure response before refresh starts", () => { + render( + , + ) + + act(() => { + window.dispatchEvent( + new MessageEvent("message", { + data: { + type: "singleRouterModelFetchResponse", + success: false, + error: "Moonshot unavailable", + values: { provider: providerIdentifiers.moonshot }, + }, + }), + ) + }) + + expect(screen.queryByText("Moonshot unavailable")).not.toBeInTheDocument() + expect(screen.queryByText("settings:providers.refreshModels.loading")).not.toBeInTheDocument() + }) + + it("resets model-specific settings when the selected model changes", () => { + render( + , + ) + + fireEvent.click(screen.getByTestId("change-model")) + + expect(mockSetApiConfigurationField).toHaveBeenCalledWith("reasoningEffort", undefined) + expect(mockSetApiConfigurationField).toHaveBeenCalledWith("modelMaxTokens", undefined) + expect(mockSetApiConfigurationField).toHaveBeenCalledWith("modelMaxThinkingTokens", undefined) + }) + it("race condition: error arrives before routerModels success — stays in error state", async () => { mockUseExtensionState.mockReturnValue({ routerModels: {}, @@ -308,7 +379,7 @@ describe("Moonshot Component", () => { type: "singleRouterModelFetchResponse", success: false, error: "API connection failed", - values: { provider: "moonshot" }, + values: { provider: providerIdentifiers.moonshot }, }, "*", ) diff --git a/webview-ui/src/components/settings/providers/__tests__/Poe.spec.tsx b/webview-ui/src/components/settings/providers/__tests__/Poe.spec.tsx new file mode 100644 index 0000000000..6595750cf9 --- /dev/null +++ b/webview-ui/src/components/settings/providers/__tests__/Poe.spec.tsx @@ -0,0 +1,115 @@ +import { QueryClient, QueryClientProvider } from "@tanstack/react-query" +import { act, fireEvent, render, screen } from "@testing-library/react" + +import { providerIdentifiers, type OrganizationAllowList, type ProviderSettings } from "@roo-code/types" + +import { Poe } from "../Poe" + +const { mockUseExtensionState } = vi.hoisted(() => ({ + mockUseExtensionState: vi.fn(), +})) + +vi.mock("@src/context/ExtensionStateContext", () => ({ + useExtensionState: mockUseExtensionState, +})) + +vi.mock("@src/i18n/TranslationContext", () => ({ + useAppTranslation: () => ({ t: (key: string) => key }), +})) + +vi.mock("@vscode/webview-ui-toolkit/react", () => ({ + VSCodeTextField: ({ children }: { children: React.ReactNode }) =>
{children}
, +})) + +vi.mock("@src/components/common/VSCodeButtonLink", () => ({ + VSCodeButtonLink: ({ children }: { children: React.ReactNode }) =>
{children}
, +})) + +vi.mock("@src/components/ui", () => ({ + Button: ({ children, onClick, disabled }: React.ComponentProps<"button">) => ( + + ), +})) + +vi.mock("../../ModelPicker", () => ({ + ModelPicker: ({ onModelChange }: { onModelChange?: (modelId: string) => void }) => ( + + ), +})) + +describe("Poe", () => { + const organizationAllowList: OrganizationAllowList = { allowAll: true, providers: {} } + const setApiConfigurationField = vi.fn() + + const renderComponent = (apiConfiguration: ProviderSettings = { apiProvider: providerIdentifiers.poe }) => { + const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } }) + + return render( + + + , + ) + } + + beforeEach(() => { + vi.clearAllMocks() + mockUseExtensionState.mockReturnValue({ routerModels: { [providerIdentifiers.poe]: {} } }) + }) + + it("shows the Poe refresh error returned by the extension", () => { + renderComponent({ apiProvider: providerIdentifiers.poe, poeApiKey: "test-key" }) + + act(() => { + window.dispatchEvent( + new MessageEvent("message", { + data: { + type: "singleRouterModelFetchResponse", + success: false, + values: { provider: providerIdentifiers.poe }, + error: "Poe authentication failed", + }, + }), + ) + }) + + expect(screen.getByText("Poe authentication failed")).toBeInTheDocument() + }) + + it("ignores failed refresh responses for another provider", () => { + renderComponent({ apiProvider: providerIdentifiers.poe, poeApiKey: "test-key" }) + + act(() => { + window.dispatchEvent( + new MessageEvent("message", { + data: { + type: "singleRouterModelFetchResponse", + success: false, + values: { provider: providerIdentifiers.openrouter }, + error: "OpenRouter authentication failed", + }, + }), + ) + }) + + expect(screen.queryByText("OpenRouter authentication failed")).not.toBeInTheDocument() + expect(screen.queryByText("settings:providers.refreshModels.error")).not.toBeInTheDocument() + }) + + it("clears model-specific reasoning settings when the Poe model changes", () => { + renderComponent() + + fireEvent.click(screen.getByTestId("model-picker")) + + expect(setApiConfigurationField).toHaveBeenCalledWith("reasoningEffort", undefined) + expect(setApiConfigurationField).toHaveBeenCalledWith("modelMaxTokens", undefined) + expect(setApiConfigurationField).toHaveBeenCalledWith("modelMaxThinkingTokens", undefined) + }) +}) diff --git a/webview-ui/src/components/settings/providers/__tests__/ProviderRouting.spec.tsx b/webview-ui/src/components/settings/providers/__tests__/ProviderRouting.spec.tsx new file mode 100644 index 0000000000..e355baa08c --- /dev/null +++ b/webview-ui/src/components/settings/providers/__tests__/ProviderRouting.spec.tsx @@ -0,0 +1,82 @@ +import { fireEvent, render, screen } from "@testing-library/react" + +import { providerIdentifiers, type OrganizationAllowList, type RouterModels } from "@roo-code/types" + +import { vscode } from "@src/utils/vscode" + +import { Unbound } from "../Unbound" +import { VercelAiGateway } from "../VercelAiGateway" + +const { modelPickerMock } = vi.hoisted(() => ({ modelPickerMock: vi.fn(() => null) })) + +vi.mock("@src/i18n/TranslationContext", () => ({ + useAppTranslation: () => ({ t: (key: string) => key }), +})) + +vi.mock("@vscode/webview-ui-toolkit/react", () => ({ + VSCodeTextField: ({ children }: { children: React.ReactNode }) =>
{children}
, +})) + +vi.mock("@src/components/ui", () => ({ + Button: ({ children, onClick }: React.ComponentProps<"button">) => , +})) + +vi.mock("../../ModelPicker", () => ({ ModelPicker: modelPickerMock })) +vi.mock("@src/components/common/VSCodeButtonLink", () => ({ VSCodeButtonLink: () => null })) + +describe("provider model routing", () => { + const organizationAllowList: OrganizationAllowList = { allowAll: true, providers: {} } + + beforeEach(() => vi.clearAllMocks()) + + it("requests fresh Unbound models when the refresh button is clicked", () => { + const postMessage = vi.spyOn(vscode, "postMessage").mockImplementation(() => undefined) + + render( + , + ) + + fireEvent.click(screen.getByRole("button", { name: "settings:providers.refreshModels.label" })) + + expect(postMessage).toHaveBeenCalledWith({ + type: "requestRouterModels", + values: { provider: providerIdentifiers.unbound, refresh: true }, + }) + }) + + it("passes Vercel AI Gateway models selected by its provider identifier to the model picker", () => { + const models = { "anthropic/claude": { contextWindow: 1, supportsPromptCache: false } } + const routerModels = Object.fromEntries( + Object.values(providerIdentifiers).map((provider) => [provider, {}]), + ) as RouterModels + routerModels[providerIdentifiers.vercelAiGateway] = models + + render( + , + ) + + expect(modelPickerMock).toHaveBeenCalledWith(expect.objectContaining({ models }), expect.anything()) + }) + + it("passes an empty model set when Vercel AI Gateway models are unavailable", () => { + render( + , + ) + + expect(modelPickerMock).toHaveBeenCalledWith(expect.objectContaining({ models: {} }), expect.anything()) + }) +}) diff --git a/webview-ui/src/components/settings/providers/__tests__/Requesty.spec.tsx b/webview-ui/src/components/settings/providers/__tests__/Requesty.spec.tsx new file mode 100644 index 0000000000..bb7d4e234f --- /dev/null +++ b/webview-ui/src/components/settings/providers/__tests__/Requesty.spec.tsx @@ -0,0 +1,57 @@ +import { fireEvent, render, screen } from "@testing-library/react" + +import { providerIdentifiers, type OrganizationAllowList } from "@roo-code/types" + +import { vscode } from "@src/utils/vscode" + +import { Requesty } from "../Requesty" + +vi.mock("@src/i18n/TranslationContext", () => ({ + useAppTranslation: () => ({ t: (key: string) => key }), +})) + +vi.mock("@vscode/webview-ui-toolkit/react", () => ({ + VSCodeTextField: ({ children }: { children: React.ReactNode }) =>
{children}
, + VSCodeCheckbox: ({ children }: { children: React.ReactNode }) =>
{children}
, +})) + +vi.mock("@src/components/ui", () => ({ + Button: ({ children, onClick }: React.ComponentProps<"button">) => ( + + ), +})) + +vi.mock("../../ModelPicker", () => ({ ModelPicker: () => null })) +vi.mock("../RequestyBalanceDisplay", () => ({ RequestyBalanceDisplay: () => null })) + +describe("Requesty", () => { + const organizationAllowList: OrganizationAllowList = { allowAll: true, providers: {} } + + it("uses the canonical Requesty identifier for OAuth and model refresh", () => { + const postMessage = vi.spyOn(vscode, "postMessage").mockImplementation(() => undefined) + + render( + , + ) + + const href = screen.getByRole("link").getAttribute("href") + expect(href).not.toBeNull() + const callbackUrl = new URL(href!).searchParams.get("callback_url") + expect(callbackUrl).not.toBeNull() + expect(new URL(callbackUrl!).pathname).toMatch(new RegExp(`/${providerIdentifiers.requesty}$`)) + + fireEvent.click(screen.getByTestId("refresh-button")) + expect(postMessage).toHaveBeenCalledWith({ + type: "requestRouterModels", + values: { provider: providerIdentifiers.requesty, refresh: true }, + }) + }) +})