Skip to content
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@ import type { Mock } from "vitest"
import {
ProviderSettings,
ModelInfo,
anthropicModels,
BEDROCK_1M_CONTEXT_MODEL_IDS,
litellmDefaultModelInfo,
kenariDefaultModelId,
Expand Down Expand Up @@ -440,6 +441,23 @@ describe("useSelectedModel", () => {
expect(result.current.info?.inputPrice).toBe(6.0)
expect(result.current.info?.outputPrice).toBe(22.5)
})

it.each(["anthropic", "gemini-cli", "fake-ai"] as const)(
"should explicitly resolve configured models for %s",
(apiProvider) => {
const apiConfiguration: ProviderSettings = {
apiProvider,
apiModelId: "claude-sonnet-4-6",
}

const wrapper = createWrapper()
const { result } = renderHook(() => useSelectedModel(apiConfiguration), { wrapper })

expect(result.current.provider).toBe(apiProvider)
expect(result.current.id).toBe("claude-sonnet-4-6")
expect(result.current.info).toEqual(anthropicModels["claude-sonnet-4-6"])
},
)
})

describe("bedrock provider with 1M context", () => {
Expand Down
74 changes: 39 additions & 35 deletions webview-ui/src/components/ui/hooks/useSelectedModel.ts
Original file line number Diff line number Diff line change
Expand Up @@ -36,6 +36,7 @@ import {
isDynamicProvider,
isRetiredProvider,
getProviderDefaultModelId,
providerIdentifiers,
} from "@roo-code/types"

import { useRouterModels } from "./useRouterModels"
Expand Down Expand Up @@ -148,7 +149,7 @@ function getSelectedModel({
// this gives a better UX than showing the default model
const defaultModelId = getProviderDefaultModelId(provider)
switch (provider) {
case "openrouter": {
case providerIdentifiers.openrouter: {
const id = getValidatedModelId(apiConfiguration.openRouterModelId, routerModels.openrouter, defaultModelId)
let info = routerModels.openrouter?.[id]
const specificProvider = apiConfiguration.openRouterSpecificProvider
Expand All @@ -164,17 +165,17 @@ function getSelectedModel({

return { id, info }
}
case "requesty": {
case providerIdentifiers.requesty: {
const id = getValidatedModelId(apiConfiguration.requestyModelId, routerModels.requesty, defaultModelId)
const routerInfo = routerModels.requesty?.[id]
return { id, info: routerInfo }
}
case "unbound": {
case providerIdentifiers.unbound: {
const id = getValidatedModelId(apiConfiguration.unboundModelId, routerModels.unbound, defaultModelId)
const routerInfo = routerModels.unbound?.[id]
return { id, info: routerInfo }
}
case "litellm": {
case providerIdentifiers.litellm: {
// When the model list is empty (not yet loaded or still loading),
// preserve the configured model ID. LiteLLM is a proxy with no inherent
// default model, so we never substitute a hardcoded default here -- when
Expand All @@ -187,17 +188,17 @@ function getSelectedModel({
const routerInfo = routerModels.litellm?.[id]
return { id, info: routerInfo ?? litellmDefaultModelInfo }
}
case "xai": {
case providerIdentifiers.xai: {
const id = apiConfiguration.apiModelId ?? defaultModelId
const info = xaiModels[id as keyof typeof xaiModels]
return info ? { id, info } : { id, info: undefined }
}
case "baseten": {
case providerIdentifiers.baseten: {
const id = apiConfiguration.apiModelId ?? defaultModelId
const info = basetenModels[id as keyof typeof basetenModels]
return { id, info }
}
case "bedrock": {
case providerIdentifiers.bedrock: {
const id = apiConfiguration.apiModelId ?? defaultModelId
const baseInfo = bedrockModels[id as keyof typeof bedrockModels]

Expand All @@ -221,7 +222,7 @@ function getSelectedModel({

return { id, info: baseInfo }
}
case "vertex": {
case providerIdentifiers.vertex: {
const id = apiConfiguration.apiModelId ?? defaultModelId
const baseInfo = vertexModels[id as keyof typeof vertexModels]

Expand All @@ -244,12 +245,12 @@ function getSelectedModel({

return { id, info: baseInfo }
}
case "gemini": {
case providerIdentifiers.gemini: {
const id = apiConfiguration.apiModelId ?? defaultModelId
const info = geminiModels[id as keyof typeof geminiModels]
return { id, info }
}
case "deepseek": {
case providerIdentifiers.deepseek: {
const availableModels = routerModels.deepseek
? { ...deepSeekModels, ...routerModels.deepseek }
: deepSeekModels
Expand All @@ -258,7 +259,7 @@ function getSelectedModel({
const staticInfo = deepSeekModels[id as keyof typeof deepSeekModels]
return { id, info: routerInfo ?? staticInfo }
}
case "moonshot": {
case providerIdentifiers.moonshot: {
const availableModels = routerModels.moonshot
? { ...moonshotModels, ...routerModels.moonshot }
: moonshotModels
Expand All @@ -267,47 +268,47 @@ function getSelectedModel({
const staticInfo = moonshotModels[id as keyof typeof moonshotModels]
return { id, info: routerInfo ?? staticInfo }
}
case "kimi-code": {
case providerIdentifiers.kimiCode: {
const configuredId = apiConfiguration.apiModelId
const availableModels = routerModels["kimi-code"]
const id = configuredId || defaultModelId
return { id, info: availableModels?.[id] ?? kimiCodeDefaultModelInfo }
}
case "minimax": {
case providerIdentifiers.minimax: {
const id = apiConfiguration.apiModelId ?? defaultModelId
const info = minimaxModels[id as keyof typeof minimaxModels]
return { id, info }
}
case "mimo": {
case providerIdentifiers.mimo: {
const id = apiConfiguration.apiModelId ?? defaultModelId
const info = mimoModels[id as keyof typeof mimoModels] ?? mimoModels["mimo-v2.5-pro"]
return { id, info }
}
case "zai": {
case providerIdentifiers.zai: {
const isChina = apiConfiguration.zaiApiLine === "china_coding"
const models = isChina ? mainlandZAiModels : internationalZAiModels
const defaultModelId = getProviderDefaultModelId(provider, { isChina })
const id = apiConfiguration.apiModelId ?? defaultModelId
const info = models[id as keyof typeof models]
return { id, info }
}
case "openai-native": {
case providerIdentifiers.openaiNative: {
const id = apiConfiguration.apiModelId ?? defaultModelId
const info = openAiNativeModels[id as keyof typeof openAiNativeModels]
return { id, info }
}
case "mistral": {
case providerIdentifiers.mistral: {
const id = apiConfiguration.apiModelId ?? defaultModelId
const info = mistralModels[id as keyof typeof mistralModels]
return { id, info }
}
case "openai": {
case providerIdentifiers.openai: {
const id = apiConfiguration.openAiModelId ?? ""
const customInfo = apiConfiguration?.openAiCustomModelInfo
const info = customInfo ?? openAiModelInfoSaneDefaults
return { id, info }
}
case "ollama": {
case providerIdentifiers.ollama: {
const id = apiConfiguration.ollamaModelId ?? ""
const info = ollamaModels && ollamaModels[apiConfiguration.ollamaModelId!]

Expand All @@ -323,15 +324,15 @@ function getSelectedModel({
info: adjustedInfo || undefined,
}
}
case "lmstudio": {
case providerIdentifiers.lmstudio: {
const id = apiConfiguration.lmStudioModelId ?? ""
const modelInfo = lmStudioModels && lmStudioModels[apiConfiguration.lmStudioModelId!]
return {
id,
info: modelInfo ? { ...lMStudioDefaultModelInfo, ...modelInfo } : undefined,
}
}
case "vscode-lm": {
case providerIdentifiers.vscodeLm: {
const id = apiConfiguration?.vsCodeLmModelSelector
? `${apiConfiguration.vsCodeLmModelSelector.vendor}/${apiConfiguration.vsCodeLmModelSelector.family}`
: vscodeLlmDefaultModelId
Expand All @@ -351,37 +352,37 @@ function getSelectedModel({
}
return { id, info }
}
case "sambanova": {
case providerIdentifiers.sambanova: {
const id = apiConfiguration.apiModelId ?? defaultModelId
const info = sambaNovaModels[id as keyof typeof sambaNovaModels]
return { id, info }
}
case "fireworks": {
case providerIdentifiers.fireworks: {
const id = apiConfiguration.apiModelId ?? defaultModelId
const info = fireworksModels[id as keyof typeof fireworksModels]
return { id, info }
}
case "friendli": {
case providerIdentifiers.friendli: {
const id = apiConfiguration.apiModelId ?? defaultModelId
const info = friendliModels[id as keyof typeof friendliModels]
return { id, info }
}
case "poe": {
case providerIdentifiers.poe: {
const id = apiConfiguration.apiModelId ?? defaultModelId
const info = routerModels.poe?.[id]
return { id, info }
}
case "qwen-code": {
case providerIdentifiers.qwenCode: {
const id = apiConfiguration.apiModelId ?? defaultModelId
const info = qwenCodeModels[id as keyof typeof qwenCodeModels]
return { id, info }
}
case "openai-codex": {
case providerIdentifiers.openaiCodex: {
const id = apiConfiguration.apiModelId ?? defaultModelId
const info = openAiCodexModels[id as keyof typeof openAiCodexModels]
return { id, info }
}
case "vercel-ai-gateway": {
case providerIdentifiers.vercelAiGateway: {
const id = getValidatedModelId(
apiConfiguration.vercelAiGatewayModelId,
routerModels["vercel-ai-gateway"],
Expand All @@ -390,7 +391,7 @@ function getSelectedModel({
const info = routerModels["vercel-ai-gateway"]?.[id]
return { id, info }
}
case "opencode-go": {
case providerIdentifiers.opencodeGo: {
const id = getValidatedModelId(
apiConfiguration.opencodeGoModelId,
routerModels["opencode-go"],
Expand All @@ -401,14 +402,14 @@ function getSelectedModel({
const info = routerModels["opencode-go"]?.[id] ?? opencodeGoDefaultModelInfo
return { id, info }
}
case "kenari": {
case providerIdentifiers.kenari: {
const id = getValidatedModelId(apiConfiguration.kenariModelId, routerModels["kenari"], defaultModelId)
// Fall back to the provider's default ModelInfo so capability-driven UI
// keeps working when the /models list is empty or unavailable.
const info = routerModels["kenari"]?.[id] ?? kenariDefaultModelInfo
return { id, info }
}
case "zoo-gateway": {
case providerIdentifiers.zooGateway: {
const id = getValidatedModelId(
apiConfiguration.zooGatewayModelId,
routerModels["zoo-gateway"],
Expand All @@ -417,10 +418,9 @@ function getSelectedModel({
const info = routerModels["zoo-gateway"]?.[id]
return { id, info }
}
// case "anthropic":
// case "fake-ai":
default: {
provider satisfies "anthropic" | "gemini-cli" | "fake-ai"
case providerIdentifiers.anthropic:
Comment thread
WebMad marked this conversation as resolved.
case providerIdentifiers.geminiCli:
case providerIdentifiers.fakeAi: {
Comment thread
coderabbitai[bot] marked this conversation as resolved.
const id = apiConfiguration.apiModelId ?? defaultModelId
const baseInfo = anthropicModels[id as keyof typeof anthropicModels]

Expand Down Expand Up @@ -461,5 +461,9 @@ function getSelectedModel({

return { id, info: baseInfo }
}
default: {
provider satisfies never
throw new Error(`Unsupported provider: ${provider}`)
}
}
}
2 changes: 1 addition & 1 deletion webview-ui/src/utils/__tests__/validate.spec.ts
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
import type { ProviderSettings, OrganizationAllowList, RouterModels } from "@roo-code/types"
import { type ProviderSettings, type OrganizationAllowList, type RouterModels } from "@roo-code/types"

// Mock i18next to return translation keys with interpolated values
vi.mock("i18next", () => ({
Expand Down
Loading
Loading