Skip to content
Open
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
3 changes: 2 additions & 1 deletion webview-ui/src/components/settings/providers/Kenari.tsx
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@ import {
type OrganizationAllowList,
type RouterModels,
kenariDefaultModelId,
providerIdentifiers,
} from "@roo-code/types"

import { useAppTranslation } from "@src/i18n/TranslationContext"
Expand Down Expand Up @@ -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"
Expand Down
7 changes: 4 additions & 3 deletions webview-ui/src/components/settings/providers/KimiCode.tsx
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@ import {
type KimiCodeAuthMethod,
type ModelRecord,
type ProviderSettings,
providerIdentifiers,
} from "@roo-code/types"

import { useAppTranslation } from "@src/i18n/TranslationContext"
Expand Down Expand Up @@ -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

Expand All @@ -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,
Expand Down
14 changes: 7 additions & 7 deletions webview-ui/src/components/settings/providers/LiteLLM.tsx
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@ import {
type OrganizationAllowList,
type ExtensionMessage,
litellmDefaultModelId,
providerIdentifiers,
} from "@roo-code/types"

import { RouterName } from "@roo/api"
Expand Down Expand Up @@ -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)
Expand All @@ -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".
}
Expand Down
13 changes: 8 additions & 5 deletions webview-ui/src/components/settings/providers/Moonshot.tsx
Original file line number Diff line number Diff line change
Expand Up @@ -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"

Expand All @@ -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"

Expand All @@ -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)
Expand Down Expand Up @@ -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)
}
/>
<Button
Expand Down
11 changes: 8 additions & 3 deletions webview-ui/src/components/settings/providers/OpenCodeGo.tsx
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@ import {
type RouterModels,
type ExtensionMessage,
opencodeGoDefaultModelId,
providerIdentifiers,
} from "@roo-code/types"

import type { RouterName } from "@roo/api"
Expand Down Expand Up @@ -46,7 +47,7 @@ export const OpenCodeGo = ({
const message = event.data
if (message.type === "singleRouterModelFetchResponse" && !message.success) {
const providerName = message.values?.provider as RouterName
if (providerName === "opencode-go") {
if (providerName === providerIdentifiers.opencodeGo) {
errorJustReceived.current = true
setRefreshStatus("error")
setRefreshError(message.error)
Expand Down Expand Up @@ -83,7 +84,11 @@ export const OpenCodeGo = ({
setRefreshError(undefined)
vscode.postMessage({
type: "requestRouterModels",
values: { provider: "opencode-go", refresh: true, opencodeGoApiKey: apiConfiguration.opencodeGoApiKey },
values: {
provider: providerIdentifiers.opencodeGo,
refresh: true,
opencodeGoApiKey: apiConfiguration.opencodeGoApiKey,
},
})
}, [apiConfiguration.opencodeGoApiKey])

Expand Down Expand Up @@ -136,7 +141,7 @@ export const OpenCodeGo = ({
apiConfiguration={apiConfiguration}
setApiConfigurationField={setApiConfigurationField}
defaultModelId={opencodeGoDefaultModelId}
models={routerModels?.["opencode-go"] ?? {}}
models={routerModels?.[providerIdentifiers.opencodeGo] ?? {}}
modelIdKey="opencodeGoModelId"
serviceName="Opencode Go"
serviceUrl="https://opencode.ai/docs/go/"
Expand Down
6 changes: 3 additions & 3 deletions webview-ui/src/components/settings/providers/Poe.tsx
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@ import {
type OrganizationAllowList,
type ExtensionMessage,
poeDefaultModelId,
type ProviderName,
providerIdentifiers,
} from "@roo-code/types"

import { RouterName } from "@roo/api"
Expand Down Expand Up @@ -49,7 +49,7 @@ export const Poe = ({
const message = event.data
if (message.type === "singleRouterModelFetchResponse" && !message.success) {
const providerName = message.values?.provider as RouterName
if (providerName === "poe") {
if (providerName === providerIdentifiers.poe) {
poeErrorJustReceived.current = true
setRefreshStatus("error")
setRefreshError(message.error)
Expand Down Expand Up @@ -158,7 +158,7 @@ export const Poe = ({
errorMessage={modelValidationError}
simplifySettings={simplifySettings}
onModelChange={(modelId) =>
handleModelChangeSideEffects("poe" as ProviderName, modelId, setApiConfigurationField)
handleModelChangeSideEffects(providerIdentifiers.poe, modelId, setApiConfigurationField)
}
/>
</>
Expand Down
8 changes: 6 additions & 2 deletions webview-ui/src/components/settings/providers/Requesty.tsx
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@ import {
type OrganizationAllowList,
type RouterModels,
requestyDefaultModelId,
providerIdentifiers,
} from "@roo-code/types"

import { vscode } from "@src/utils/vscode"
Expand Down Expand Up @@ -59,7 +60,7 @@ export const Requesty = ({
)

const getApiKeyUrl = () => {
const callbackUrl = getCallbackUrl("requesty", uriScheme)
const callbackUrl = getCallbackUrl(providerIdentifiers.requesty, uriScheme)
const baseUrl = toRequestyServiceUrl(apiConfiguration.requestyBaseUrl, "app")

const authUrl = new URL(`oauth/authorize?callback_url=${callbackUrl}`, baseUrl)
Expand Down Expand Up @@ -129,7 +130,10 @@ export const Requesty = ({
<Button
variant="outline"
onClick={() => {
vscode.postMessage({ type: "requestRouterModels", values: { provider: "requesty", refresh: true } })
vscode.postMessage({
type: "requestRouterModels",
values: { provider: providerIdentifiers.requesty, refresh: true },
})
}}>
<div className="flex items-center gap-2">
<span className="codicon codicon-refresh" />
Expand Down
6 changes: 5 additions & 1 deletion webview-ui/src/components/settings/providers/Unbound.tsx
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@ import {
type OrganizationAllowList,
type RouterModels,
unboundDefaultModelId,
providerIdentifiers,
} from "@roo-code/types"

import { vscode } from "@src/utils/vscode"
Expand Down Expand Up @@ -77,7 +78,10 @@ export const Unbound = ({
<Button
variant="outline"
onClick={() => {
vscode.postMessage({ type: "requestRouterModels", values: { provider: "unbound", refresh: true } })
vscode.postMessage({
type: "requestRouterModels",
values: { provider: providerIdentifiers.unbound, refresh: true },
})
}}>
<div className="flex items-center gap-2">
<span className="codicon codicon-refresh" />
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@ import {
type OrganizationAllowList,
type RouterModels,
vercelAiGatewayDefaultModelId,
providerIdentifiers,
} from "@roo-code/types"

import { useAppTranslation } from "@src/i18n/TranslationContext"
Expand Down Expand Up @@ -69,7 +70,7 @@ export const VercelAiGateway = ({
apiConfiguration={apiConfiguration}
setApiConfigurationField={setApiConfigurationField}
defaultModelId={vercelAiGatewayDefaultModelId}
models={routerModels?.["vercel-ai-gateway"] ?? {}}
models={routerModels?.[providerIdentifiers.vercelAiGateway] ?? {}}
modelIdKey="vercelAiGatewayModelId"
serviceName="Vercel AI Gateway"
serviceUrl="https://vercel.com/ai-gateway/models"
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@ import {
type OrganizationAllowList,
type RouterModels,
zooGatewayDefaultModelId,
providerIdentifiers,
} from "@roo-code/types"

import { useExtensionState } from "@src/context/ExtensionStateContext"
Expand Down Expand Up @@ -63,7 +64,7 @@ export const ZooGateway = ({
const authUrl = getZooCodeAuthUrl(uriScheme, zooCodeBaseUrl, deviceName)
const resolvedDashboardBase = zooCodeBaseUrl?.replace(/\/$/, "") || "https://www.zoocode.dev"

const zooModels = useMemo(() => routerModels?.["zoo-gateway"] ?? {}, [routerModels])
const zooModels = useMemo(() => routerModels?.[providerIdentifiers.zooGateway] ?? {}, [routerModels])
const modelIds = useMemo(() => Object.keys(zooModels), [zooModels])
const resolvedDefaultModelId = useMemo(() => pickZooGatewayDefaultModelId(modelIds), [modelIds])

Expand Down
Original file line number Diff line number Diff line change
@@ -1,21 +1,88 @@
import { fireEvent, render, screen } from "@testing-library/react"

import { kimiCodeModels, providerIdentifiers, type ModelRecord } from "@roo-code/types"

import { KimiCode } from "../KimiCode"

const { mockUseRouterModels } = vi.hoisted(() => ({
mockUseRouterModels: vi.fn(),
}))

vi.mock("@src/components/ui/hooks/useRouterModels", () => ({
useRouterModels: () => ({ data: { "kimi-code": {} }, refetch: vi.fn(), isFetching: false }),
useRouterModels: mockUseRouterModels,
}))

vi.mock("../../ModelPicker", () => ({
ModelPicker: () => <div data-testid="kimi-code-model-picker" />,
ModelPicker: ({ models }: { models: ModelRecord }) => (
<div data-testid="kimi-code-model-picker" data-model-ids={JSON.stringify(Object.keys(models))} />
),
}))

describe("KimiCode settings", () => {
beforeEach(() => {
vi.clearAllMocks()
mockUseRouterModels.mockReturnValue({
data: { [providerIdentifiers.kimiCode]: {} },
refetch: vi.fn(),
isFetching: false,
})
})

it("defaults to OAuth when no authentication method is configured", () => {
render(
<KimiCode
apiConfiguration={{ apiProvider: providerIdentifiers.kimiCode }}
setApiConfigurationField={vi.fn()}
/>,
)

expect(screen.getByText("settings:providers.kimiCode.signIn")).toBeInTheDocument()
expect(screen.queryByTestId("kimi-code-api-key")).not.toBeInTheDocument()
expect(mockUseRouterModels).toHaveBeenCalledWith({
provider: providerIdentifiers.kimiCode,
enabled: false,
})
})

it("uses discovered Kimi Code models when the extension returns models", () => {
mockUseRouterModels.mockReturnValue({
data: { [providerIdentifiers.kimiCode]: { "discovered-model": {} } },
refetch: vi.fn(),
isFetching: false,
})

render(
<KimiCode
apiConfiguration={{ apiProvider: providerIdentifiers.kimiCode }}
setApiConfigurationField={vi.fn()}
/>,
)

expect(screen.getByTestId("kimi-code-model-picker")).toHaveAttribute(
"data-model-ids",
JSON.stringify(["discovered-model"]),
)
})

it("falls back to static Kimi Code models when no models are discovered", () => {
render(
<KimiCode
apiConfiguration={{ apiProvider: providerIdentifiers.kimiCode }}
setApiConfigurationField={vi.fn()}
/>,
)

expect(screen.getByTestId("kimi-code-model-picker")).toHaveAttribute(
"data-model-ids",
JSON.stringify(Object.keys(kimiCodeModels)),
)
})

it("binds the API key input through the buffered settings setter", () => {
const setField = vi.fn()
render(
<KimiCode
apiConfiguration={{ apiProvider: "kimi-code", kimiCodeAuthMethod: "api-key" }}
apiConfiguration={{ apiProvider: providerIdentifiers.kimiCode, kimiCodeAuthMethod: "api-key" }}
setApiConfigurationField={setField}
/>,
)
Expand All @@ -27,7 +94,7 @@ describe("KimiCode settings", () => {
it("shows device-code polling state", () => {
render(
<KimiCode
apiConfiguration={{ apiProvider: "kimi-code", kimiCodeAuthMethod: "oauth" }}
apiConfiguration={{ apiProvider: providerIdentifiers.kimiCode, kimiCodeAuthMethod: "oauth" }}
setApiConfigurationField={vi.fn()}
kimiCodeOAuthState={{
status: "polling",
Expand Down
Loading
Loading