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
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,12 @@ import {
minimaxModels,
friendliDefaultModelId,
friendliModels,
fireworksDefaultModelId,
fireworksModels,
basetenDefaultModelId,
basetenModels,
sambaNovaDefaultModelId,
sambaNovaModels,
deepSeekDefaultModelId,
deepSeekModels,
openRouterDefaultModelId,
Expand Down Expand Up @@ -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<string, ModelInfo>,
knownId: "accounts/fireworks/models/kimi-k2-instruct",
},
{
apiProvider: providerIdentifiers.baseten,
defaultId: basetenDefaultModelId,
models: basetenModels as Record<string, ModelInfo>,
knownId: Object.keys(basetenModels).find((id) => id !== basetenDefaultModelId)!,
},
{
apiProvider: providerIdentifiers.sambanova,
defaultId: sambaNovaDefaultModelId,
models: sambaNovaModels as Record<string, ModelInfo>,
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)
},
)
Comment thread
coderabbitai[bot] marked this conversation as resolved.
})

describe("Z AI provider", () => {
it("uses the International Coding catalog when no API line is configured", () => {
const apiConfiguration: ProviderSettings = {
Expand Down
42 changes: 27 additions & 15 deletions webview-ui/src/components/ui/hooks/useSelectedModel.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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<string, ModelInfo>,
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 } }

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Can we require the backend custom-model fix before this ships so the displayed model matches the one sent to the provider?

Comment thread
p12tic marked this conversation as resolved.
}
return { id: configuredId, info: models[configuredId] }
}

/**
* Resolves the model currently selected for the active API provider.
*
Expand Down Expand Up @@ -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]
Expand Down Expand Up @@ -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)
Comment thread
coderabbitai[bot] marked this conversation as resolved.
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]
Expand Down
Loading