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
14 changes: 12 additions & 2 deletions webview-ui/src/components/settings/ApiOptions.tsx
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@ import { ExternalLinkIcon } from "@radix-ui/react-icons"
import {
type ProviderName,
type ProviderSettings,
isDynamicProvider,
isRetiredProvider,
providerIdentifiers,
DEFAULT_CONSECUTIVE_MISTAKE_LIMIT,
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -112,6 +113,61 @@ describe("ApiOptions Provider Filtering", () => {
return renderWithExtensionState(<ApiOptions {...props} />)
}

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()

Expand Down
3 changes: 2 additions & 1 deletion webview-ui/src/components/settings/providers/LMStudio.tsx
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@ import {
type ExtensionMessage,
type ModelRecord,
LmStudioModelsMessageType,
providerIdentifiers,
} from "@roo-code/types"

import { useAppTranslation } from "@src/i18n/TranslationContext"
Expand All @@ -27,7 +28,7 @@ export const LMStudio = ({ apiConfiguration, setApiConfigurationField }: LMStudi
const { t } = useAppTranslation()

const [lmStudioModels, setLmStudioModels] = useState<ModelRecord>({})
const routerModels = useRouterModels()
const routerModels = useRouterModels({ provider: providerIdentifiers.lmstudio })
const initialBaseUrlRef = useRef(apiConfiguration?.lmStudioBaseUrl)

const handleInputChange = useCallback(
Expand Down
2 changes: 1 addition & 1 deletion webview-ui/src/components/settings/providers/LiteLLM.tsx
Original file line number Diff line number Diff line change
Expand Up @@ -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])

Expand Down
6 changes: 5 additions & 1 deletion webview-ui/src/components/settings/providers/Moonshot.tsx
Original file line number Diff line number Diff line change
Expand Up @@ -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])

Expand Down
3 changes: 2 additions & 1 deletion webview-ui/src/components/settings/providers/Ollama.tsx
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@ import {
type ModelRecord,
ollamaDefaultModelInfo,
OllamaModelsMessageType,
providerIdentifiers,
} from "@roo-code/types"

import { useAppTranslation } from "@src/i18n/TranslationContext"
Expand Down Expand Up @@ -38,7 +39,7 @@ export const Ollama = ({ apiConfiguration, setApiConfigurationField }: OllamaPro
const [refreshStatus, setRefreshStatus] = useState(RefreshStatus.Idle)
const [refreshError, setRefreshError] = useState<string | undefined>()
const refreshStatusRef = useRef(refreshStatus)
const routerModels = useRouterModels()
const routerModels = useRouterModels({ provider: providerIdentifiers.ollama })

const handleInputChange = useCallback(
<K extends keyof ProviderSettings, E>(
Expand Down
2 changes: 1 addition & 1 deletion webview-ui/src/components/settings/providers/Poe.tsx
Original file line number Diff line number Diff line change
Expand Up @@ -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])

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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" } }))
})
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -466,6 +466,7 @@ describe("Moonshot Component", () => {
expect.objectContaining({
type: "requestRouterModels",
values: {
provider: providerIdentifiers.moonshot,
moonshotApiKey: "test-key",
moonshotBaseUrl: "https://api.moonshot.cn/v1",
},
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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", () => ({
Expand Down Expand Up @@ -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" } }))
})
Expand Down
Loading