diff --git a/webview-ui/src/components/ui/hooks/__tests__/useSelectedModel.spec.ts b/webview-ui/src/components/ui/hooks/__tests__/useSelectedModel.spec.ts index 64d3067092..7245faf684 100644 --- a/webview-ui/src/components/ui/hooks/__tests__/useSelectedModel.spec.ts +++ b/webview-ui/src/components/ui/hooks/__tests__/useSelectedModel.spec.ts @@ -24,6 +24,12 @@ import { minimaxModels, friendliDefaultModelId, friendliModels, + fireworksDefaultModelId, + fireworksModels, + basetenDefaultModelId, + basetenModels, + sambaNovaDefaultModelId, + sambaNovaModels, deepSeekDefaultModelId, deepSeekModels, openRouterDefaultModelId, @@ -1372,6 +1378,109 @@ describe("useSelectedModel", () => { }) }) + describe("OpenAI-compatible static catalog providers (fireworks, baseten, sambanova)", () => { + beforeEach(() => { + mockUseRouterModels.mockReturnValue(createRouterModelsResult({ openrouter: {}, requesty: {}, litellm: {} })) + mockUseOpenRouterModelProviders.mockReturnValue(createOpenRouterModelProvidersResult({})) + }) + + it("returns the default Fireworks model when no model is configured", () => { + const { result } = renderHook(() => useSelectedModel({ apiProvider: providerIdentifiers.fireworks }), { + wrapper: createWrapper(), + }) + + expect(result.current.id).toBe(fireworksDefaultModelId) + expect(result.current.info).toEqual(fireworksModels[fireworksDefaultModelId]) + }) + + it("returns known Fireworks model metadata", () => { + const id = "accounts/fireworks/models/kimi-k2-instruct" + const { result } = renderHook( + () => useSelectedModel({ apiProvider: providerIdentifiers.fireworks, apiModelId: id }), + { wrapper: createWrapper() }, + ) + + expect(result.current.id).toBe(id) + expect(result.current.info).toEqual(fireworksModels[id]) + }) + + const staticCatalogCases = [ + { + apiProvider: providerIdentifiers.fireworks, + defaultId: fireworksDefaultModelId, + models: fireworksModels as Record, + knownId: "accounts/fireworks/models/kimi-k2-instruct", + }, + { + apiProvider: providerIdentifiers.baseten, + defaultId: basetenDefaultModelId, + models: basetenModels as Record, + knownId: Object.keys(basetenModels).find((id) => id !== basetenDefaultModelId)!, + }, + { + apiProvider: providerIdentifiers.sambanova, + defaultId: sambaNovaDefaultModelId, + models: sambaNovaModels as Record, + knownId: Object.keys(sambaNovaModels).find((id) => id !== sambaNovaDefaultModelId)!, + }, + ] + + it.each(staticCatalogCases)( + "returns the default $apiProvider model when no model is configured", + ({ apiProvider, defaultId, models }) => { + const { result } = renderHook(() => useSelectedModel({ apiProvider }), { wrapper: createWrapper() }) + + expect(result.current.provider).toBe(apiProvider) + expect(result.current.id).toBe(defaultId) + expect(result.current.info).toEqual(models[defaultId]) + expect(result.current.info).toBeDefined() + }, + ) + + it.each(staticCatalogCases)( + "falls back to the default $apiProvider model when apiModelId is an empty string", + ({ apiProvider, defaultId, models }) => { + const { result } = renderHook(() => useSelectedModel({ apiProvider, apiModelId: "" }), { + wrapper: createWrapper(), + }) + + expect(result.current.id).toBe(defaultId) + expect(result.current.info).toEqual(models[defaultId]) + expect(result.current.info).toBeDefined() + }, + ) + + it.each(staticCatalogCases)( + "returns known $apiProvider catalog metadata", + ({ apiProvider, models, knownId }) => { + expect(knownId).toBeDefined() + const { result } = renderHook(() => useSelectedModel({ apiProvider, apiModelId: knownId }), { + wrapper: createWrapper(), + }) + + expect(result.current.provider).toBe(apiProvider) + expect(result.current.id).toBe(knownId) + expect(result.current.info).toEqual(models[knownId]) + }, + ) + + it.each([providerIdentifiers.fireworks, providerIdentifiers.baseten, providerIdentifiers.sambanova])( + "honors a custom %s model ID with sane default metadata (not undefined info)", + (apiProvider) => { + // Regression: custom models had undefined info, so the task header fell back to a + // context window of 1. The backend uses openAiModelInfoSaneDefaults, so the UI must too. + const { result } = renderHook( + () => useSelectedModel({ apiProvider, apiModelId: "some/custom-model-not-in-list" }), + { wrapper: createWrapper() }, + ) + + expect(result.current.id).toBe("some/custom-model-not-in-list") + expect(result.current.info).toEqual(openAiModelInfoSaneDefaults) + expect(result.current.info?.contextWindow).toBeGreaterThan(1) + }, + ) + }) + describe("Z AI provider", () => { it("uses the International Coding catalog when no API line is configured", () => { const apiConfiguration: ProviderSettings = { diff --git a/webview-ui/src/components/ui/hooks/useSelectedModel.ts b/webview-ui/src/components/ui/hooks/useSelectedModel.ts index 8d4b70ad4a..859da7f957 100644 --- a/webview-ui/src/components/ui/hooks/useSelectedModel.ts +++ b/webview-ui/src/components/ui/hooks/useSelectedModel.ts @@ -57,6 +57,27 @@ function getValidatedModelId( return configuredId && availableModels?.[configuredId] ? configuredId : defaultModelId } +/** + * Resolves a static-catalog model for providers backed by `BaseOpenAiCompatibleProvider` + * (Fireworks, Baseten, SambaNova). Mirrors the backend `getModel()`: a user-supplied custom + * model ID that isn't in the static list is honored with `openAiModelInfoSaneDefaults`, so the + * UI shows sensible limits (e.g. context window) instead of treating the selection as unknown. + */ +function getStaticCatalogModel( + configuredId: string | undefined, + models: Record, + defaultModelId: string, +): { id: string; info: ModelInfo | undefined } { + // An empty configured ID means "unset": fall back to the provider default. + if (!configuredId) { + return { id: defaultModelId, info: models[defaultModelId] } + } + if (!(configuredId in models)) { + return { id: configuredId, info: { ...openAiModelInfoSaneDefaults } } + } + return { id: configuredId, info: models[configuredId] } +} + /** * Resolves the model currently selected for the active API provider. * @@ -230,11 +251,8 @@ function getSelectedModel({ const info = xaiModels[id as keyof typeof xaiModels] return info ? { id, info } : { id, info: undefined } } - case providerIdentifiers.baseten: { - const id = apiConfiguration.apiModelId ?? defaultModelId - const info = basetenModels[id as keyof typeof basetenModels] - return { id, info } - } + case providerIdentifiers.baseten: + return getStaticCatalogModel(apiConfiguration.apiModelId, basetenModels, defaultModelId) case providerIdentifiers.bedrock: { const id = apiConfiguration.apiModelId ?? defaultModelId const baseInfo = bedrockModels[id as keyof typeof bedrockModels] @@ -390,16 +408,10 @@ function getSelectedModel({ } return { id, info } } - case providerIdentifiers.sambanova: { - const id = apiConfiguration.apiModelId ?? defaultModelId - const info = sambaNovaModels[id as keyof typeof sambaNovaModels] - return { id, info } - } - case providerIdentifiers.fireworks: { - const id = apiConfiguration.apiModelId ?? defaultModelId - const info = fireworksModels[id as keyof typeof fireworksModels] - return { id, info } - } + case providerIdentifiers.sambanova: + return getStaticCatalogModel(apiConfiguration.apiModelId, sambaNovaModels, defaultModelId) + case providerIdentifiers.fireworks: + return getStaticCatalogModel(apiConfiguration.apiModelId, fireworksModels, defaultModelId) case providerIdentifiers.friendli: { const id = apiConfiguration.apiModelId ?? defaultModelId const info = friendliModels[id as keyof typeof friendliModels]