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" } }))
})