From a83fc42433687fff66dd6973bfa8528d6c48e93c Mon Sep 17 00:00:00 2001 From: Roomote Date: Fri, 28 Aug 2026 13:10:52 +0000 Subject: [PATCH] fix(settings): scope provider model requests --- .../src/components/settings/ApiOptions.tsx | 14 ++++- .../ApiOptions.interactions.spec.tsx | 11 +++- .../ApiOptions.provider-filtering.spec.tsx | 62 ++++++++++++++++++- .../settings/providers/LMStudio.tsx | 3 +- .../components/settings/providers/LiteLLM.tsx | 2 +- .../settings/providers/Moonshot.tsx | 6 +- .../components/settings/providers/Ollama.tsx | 3 +- .../src/components/settings/providers/Poe.tsx | 2 +- .../providers/__tests__/LiteLLM.spec.tsx | 8 +++ .../providers/__tests__/Moonshot.spec.tsx | 1 + .../settings/providers/__tests__/Poe.spec.tsx | 15 ++++- 11 files changed, 114 insertions(+), 13 deletions(-) diff --git a/webview-ui/src/components/settings/ApiOptions.tsx b/webview-ui/src/components/settings/ApiOptions.tsx index 3e1495baff..714b507774 100644 --- a/webview-ui/src/components/settings/ApiOptions.tsx +++ b/webview-ui/src/components/settings/ApiOptions.tsx @@ -7,6 +7,7 @@ import { ExternalLinkIcon } from "@radix-ui/react-icons" import { type ProviderName, type ProviderSettings, + isDynamicProvider, isRetiredProvider, providerIdentifiers, DEFAULT_CONSECUTIVE_MISTAKE_LIMIT, @@ -180,8 +181,13 @@ const ApiOptions = ({ : selectedProvider const isRetiredSelectedProvider = typeof apiConfiguration.apiProvider === "string" && isRetiredProvider(apiConfiguration.apiProvider) + const routerProvider = + activeSelectedProvider && isDynamicProvider(activeSelectedProvider) ? activeSelectedProvider : undefined - const { data: routerModels, refetch: refetchRouterModels } = useRouterModels() + const { data: routerModels, refetch: refetchRouterModels } = useRouterModels({ + enabled: !!routerProvider, + provider: routerProvider, + }) useZooGatewayRouterModelsSync() const { data: openRouterModelProviders } = useOpenRouterModelProviders( @@ -242,12 +248,16 @@ const ApiOptions = ({ vscode.postMessage({ type: RouterModelsMessageType.requestRouterModels, values: { + provider: providerIdentifiers.litellm, litellmApiKey: apiConfiguration?.litellmApiKey, litellmBaseUrl: apiConfiguration?.litellmBaseUrl, }, }) } else if (selectedProvider === providerIdentifiers.poe) { - vscode.postMessage({ type: RouterModelsMessageType.requestRouterModels }) + vscode.postMessage({ + type: RouterModelsMessageType.requestRouterModels, + values: { provider: providerIdentifiers.poe }, + }) } }, 250, diff --git a/webview-ui/src/components/settings/__tests__/ApiOptions.interactions.spec.tsx b/webview-ui/src/components/settings/__tests__/ApiOptions.interactions.spec.tsx index 60057ded92..e96340ad4f 100644 --- a/webview-ui/src/components/settings/__tests__/ApiOptions.interactions.spec.tsx +++ b/webview-ui/src/components/settings/__tests__/ApiOptions.interactions.spec.tsx @@ -233,13 +233,20 @@ describe("ApiOptions interactions", () => { configuration: { litellmBaseUrl: "http://litellm:4000", litellmApiKey: "litellm-key" }, expectedMessage: { type: "requestRouterModels", - values: { litellmApiKey: "litellm-key", litellmBaseUrl: "http://litellm:4000" }, + values: { + provider: providerIdentifiers.litellm, + litellmApiKey: "litellm-key", + litellmBaseUrl: "http://litellm:4000", + }, }, }, { provider: providerIdentifiers.poe, configuration: { poeApiKey: "poe-key", poeBaseUrl: "https://api.poe.example/v1" }, - expectedMessage: { type: "requestRouterModels" }, + expectedMessage: { + type: "requestRouterModels", + values: { provider: providerIdentifiers.poe }, + }, }, ])("requests models for $provider", ({ provider, configuration, expectedMessage }) => { vi.useFakeTimers() diff --git a/webview-ui/src/components/settings/__tests__/ApiOptions.provider-filtering.spec.tsx b/webview-ui/src/components/settings/__tests__/ApiOptions.provider-filtering.spec.tsx index c9fb64272b..21a4cb865c 100644 --- a/webview-ui/src/components/settings/__tests__/ApiOptions.provider-filtering.spec.tsx +++ b/webview-ui/src/components/settings/__tests__/ApiOptions.provider-filtering.spec.tsx @@ -2,9 +2,10 @@ import { screen } from "@testing-library/react" import { renderWithExtensionState } from "@/utils/test-utils" -import type { ProviderSettings, OrganizationAllowList } from "@roo-code/types" +import { providerIdentifiers, type ProviderSettings, type OrganizationAllowList } from "@roo-code/types" import { useExtensionState } from "@src/context/ExtensionStateContext" +import { useRouterModels } from "@src/components/ui/hooks/useRouterModels" import { useSelectedModel } from "@src/components/ui/hooks/useSelectedModel" import ApiOptions from "../ApiOptions" @@ -35,10 +36,10 @@ vi.mock("@src/utils/vscode", () => ({ // Mock the router models hook vi.mock("@src/components/ui/hooks/useRouterModels", () => ({ - useRouterModels: () => ({ + useRouterModels: vi.fn(() => ({ data: null, refetch: vi.fn(), - }), + })), })) // Mock the selected model hook @@ -112,6 +113,61 @@ describe("ApiOptions Provider Filtering", () => { return renderWithExtensionState() } + beforeEach(() => { + vi.clearAllMocks() + vi.mocked(useSelectedModel).mockReturnValue({ + provider: "anthropic", + id: "claude-3-5-sonnet-20241022", + info: undefined, + isLoading: false, + isError: false, + }) + }) + + it("does not request router models for a static provider", () => { + renderWithProviders() + + expect(useRouterModels).toHaveBeenCalledWith({ enabled: false, provider: undefined }) + }) + + it("requests router models only for the selected dynamic provider", () => { + vi.mocked(useSelectedModel).mockReturnValue({ + provider: "kenari", + id: "glm-5-2", + info: undefined, + isLoading: false, + isError: false, + }) + + renderWithProviders({ + ...defaultProps, + apiConfiguration: { apiProvider: "kenari" } as ProviderSettings, + }) + + expect(useRouterModels).toHaveBeenCalledWith({ enabled: true, provider: "kenari" }) + }) + + it.each([providerIdentifiers.ollama, providerIdentifiers.lmstudio])( + "does not make an aggregate router request for local provider %s", + (provider) => { + vi.mocked(useSelectedModel).mockReturnValue({ + provider, + id: "local-model", + info: undefined, + isLoading: false, + isError: false, + }) + + renderWithProviders({ + ...defaultProps, + apiConfiguration: { apiProvider: provider } as ProviderSettings, + }) + + expect(useRouterModels).toHaveBeenCalledWith({ provider }) + expect(useRouterModels).not.toHaveBeenCalledWith() + }, + ) + it("should show all providers when no organization allow list is provided", () => { renderWithProviders() diff --git a/webview-ui/src/components/settings/providers/LMStudio.tsx b/webview-ui/src/components/settings/providers/LMStudio.tsx index 64c12606c1..4ae20b939f 100644 --- a/webview-ui/src/components/settings/providers/LMStudio.tsx +++ b/webview-ui/src/components/settings/providers/LMStudio.tsx @@ -9,6 +9,7 @@ import { type ExtensionMessage, type ModelRecord, LmStudioModelsMessageType, + providerIdentifiers, } from "@roo-code/types" import { useAppTranslation } from "@src/i18n/TranslationContext" @@ -27,7 +28,7 @@ export const LMStudio = ({ apiConfiguration, setApiConfigurationField }: LMStudi const { t } = useAppTranslation() const [lmStudioModels, setLmStudioModels] = useState({}) - const routerModels = useRouterModels() + const routerModels = useRouterModels({ provider: providerIdentifiers.lmstudio }) const initialBaseUrlRef = useRef(apiConfiguration?.lmStudioBaseUrl) const handleInputChange = useCallback( diff --git a/webview-ui/src/components/settings/providers/LiteLLM.tsx b/webview-ui/src/components/settings/providers/LiteLLM.tsx index 5f3b7dc27b..ab2b13dc17 100644 --- a/webview-ui/src/components/settings/providers/LiteLLM.tsx +++ b/webview-ui/src/components/settings/providers/LiteLLM.tsx @@ -115,7 +115,7 @@ export const LiteLLM = ({ vscode.postMessage({ type: RouterModelsMessageType.requestRouterModels, - values: { litellmApiKey: key, litellmBaseUrl: url }, + values: { provider: providerIdentifiers.litellm, litellmApiKey: key, litellmBaseUrl: url }, }) }, [apiConfiguration, setRefreshStatus, setRefreshError, t]) diff --git a/webview-ui/src/components/settings/providers/Moonshot.tsx b/webview-ui/src/components/settings/providers/Moonshot.tsx index ed561f2b9f..6e7d5aac33 100644 --- a/webview-ui/src/components/settings/providers/Moonshot.tsx +++ b/webview-ui/src/components/settings/providers/Moonshot.tsx @@ -101,7 +101,11 @@ export const Moonshot = ({ apiConfiguration, setApiConfigurationField, simplifyS vscode.postMessage({ type: RouterModelsMessageType.requestRouterModels, - values: { moonshotApiKey: key, moonshotBaseUrl: apiConfiguration.moonshotBaseUrl }, + values: { + provider: providerIdentifiers.moonshot, + moonshotApiKey: key, + moonshotBaseUrl: apiConfiguration.moonshotBaseUrl, + }, }) }, [apiConfiguration, t]) diff --git a/webview-ui/src/components/settings/providers/Ollama.tsx b/webview-ui/src/components/settings/providers/Ollama.tsx index 8d1e7348f4..80a8bf5904 100644 --- a/webview-ui/src/components/settings/providers/Ollama.tsx +++ b/webview-ui/src/components/settings/providers/Ollama.tsx @@ -8,6 +8,7 @@ import { type ModelRecord, ollamaDefaultModelInfo, OllamaModelsMessageType, + providerIdentifiers, } from "@roo-code/types" import { useAppTranslation } from "@src/i18n/TranslationContext" @@ -38,7 +39,7 @@ export const Ollama = ({ apiConfiguration, setApiConfigurationField }: OllamaPro const [refreshStatus, setRefreshStatus] = useState(RefreshStatus.Idle) const [refreshError, setRefreshError] = useState() const refreshStatusRef = useRef(refreshStatus) - const routerModels = useRouterModels() + const routerModels = useRouterModels({ provider: providerIdentifiers.ollama }) const handleInputChange = useCallback( ( diff --git a/webview-ui/src/components/settings/providers/Poe.tsx b/webview-ui/src/components/settings/providers/Poe.tsx index b549b8aae1..ef4cc51d4f 100644 --- a/webview-ui/src/components/settings/providers/Poe.tsx +++ b/webview-ui/src/components/settings/providers/Poe.tsx @@ -112,7 +112,7 @@ export const Poe = ({ vscode.postMessage({ type: RouterModelsMessageType.requestRouterModels, - values: { poeApiKey: key, poeBaseUrl: apiConfiguration.poeBaseUrl }, + values: { provider: providerIdentifiers.poe, poeApiKey: key, poeBaseUrl: apiConfiguration.poeBaseUrl }, }) }, [apiConfiguration, t]) diff --git a/webview-ui/src/components/settings/providers/__tests__/LiteLLM.spec.tsx b/webview-ui/src/components/settings/providers/__tests__/LiteLLM.spec.tsx index e4ff4521b1..1a2384eb44 100644 --- a/webview-ui/src/components/settings/providers/__tests__/LiteLLM.spec.tsx +++ b/webview-ui/src/components/settings/providers/__tests__/LiteLLM.spec.tsx @@ -72,6 +72,14 @@ describe("LiteLLM", () => { ) fireEvent.click(screen.getByTestId("refresh-button")) + expect(postMessageMock).toHaveBeenCalledWith({ + type: "requestRouterModels", + values: { + provider: providerIdentifiers.litellm, + litellmApiKey: "test-key", + litellmBaseUrl: "http://localhost:4000", + }, + }) act(() => { window.dispatchEvent(new MessageEvent("message", { data: { type: "routerModels" } })) }) 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 1cffc1d5ba..5c5d1b41b9 100644 --- a/webview-ui/src/components/settings/providers/__tests__/Moonshot.spec.tsx +++ b/webview-ui/src/components/settings/providers/__tests__/Moonshot.spec.tsx @@ -466,6 +466,7 @@ describe("Moonshot Component", () => { expect.objectContaining({ type: "requestRouterModels", values: { + provider: providerIdentifiers.moonshot, moonshotApiKey: "test-key", moonshotBaseUrl: "https://api.moonshot.cn/v1", }, diff --git a/webview-ui/src/components/settings/providers/__tests__/Poe.spec.tsx b/webview-ui/src/components/settings/providers/__tests__/Poe.spec.tsx index b8c9f6e254..34e5d5c2d8 100644 --- a/webview-ui/src/components/settings/providers/__tests__/Poe.spec.tsx +++ b/webview-ui/src/components/settings/providers/__tests__/Poe.spec.tsx @@ -10,8 +10,13 @@ import { import { Poe } from "../Poe" -const { mockUseExtensionState } = vi.hoisted(() => ({ +const { mockUseExtensionState, postMessageMock } = vi.hoisted(() => ({ mockUseExtensionState: vi.fn(), + postMessageMock: vi.fn(), +})) + +vi.mock("@src/utils/vscode", () => ({ + vscode: { postMessage: postMessageMock }, })) vi.mock("@src/context/ExtensionStateContext", () => ({ @@ -123,6 +128,14 @@ describe("Poe", () => { ) fireEvent.click(screen.getByTestId("refresh-button")) + expect(postMessageMock).toHaveBeenCalledWith({ + type: "requestRouterModels", + values: { + provider: providerIdentifiers.poe, + poeApiKey: "test-key", + poeBaseUrl: undefined, + }, + }) act(() => { window.dispatchEvent(new MessageEvent("message", { data: { type: "routerModels" } })) })