From 46e906837f5061b428a50e2d93c4bf0001f98a3e Mon Sep 17 00:00:00 2001 From: Povilas Kanapickas Date: Mon, 28 Sep 2026 22:12:44 +0000 Subject: [PATCH] fix(api): honor custom model IDs for OpenAI-compatible providers Previously, getModel() silently fell back to the provider default when the configured apiModelId wasn't found in the static providerModels list. This sent the wrong model to the API for user-supplied custom models (e.g. newly released Fireworks models), surfacing as confusing "model not found" errors. Now getModel() honors the exact user-configured ID and supplies openAiModelInfoSaneDefaults so the rest of the pipeline works, only falling back to the provider default when no model is configured. - Return known model metadata when the ID is in providerModels - Return the custom ID with sane defaults when unknown but configured - Fall back to the default only when no model ID is set --- .../base-openai-compatible-provider.spec.ts | 81 ++++++++++++++++--- .../base-openai-compatible-provider.ts | 28 +++++-- .../hooks/__tests__/useSelectedModel.spec.ts | 41 +++++++++- .../components/ui/hooks/useSelectedModel.ts | 8 +- 4 files changed, 138 insertions(+), 20 deletions(-) diff --git a/src/api/providers/__tests__/base-openai-compatible-provider.spec.ts b/src/api/providers/__tests__/base-openai-compatible-provider.spec.ts index 60d4cae831..e761a380ea 100644 --- a/src/api/providers/__tests__/base-openai-compatible-provider.spec.ts +++ b/src/api/providers/__tests__/base-openai-compatible-provider.spec.ts @@ -3,7 +3,7 @@ import { Anthropic } from "@anthropic-ai/sdk" import OpenAI from "openai" -import type { ModelInfo } from "@roo-code/types" +import { type ModelInfo, openAiModelInfoSaneDefaults } from "@roo-code/types" import { BaseOpenAiCompatibleProvider } from "../base-openai-compatible-provider" import { asyncStreamFrom, collectStream } from "../../../test-utils/stream" @@ -25,18 +25,20 @@ vi.mock("openai", () => ({ }), })) +const TEST_MODEL_INFO: ModelInfo = { + maxTokens: 4096, + contextWindow: 128000, + supportsImages: false, + supportsPromptCache: false, + inputPrice: 0.5, + outputPrice: 1.5, +} + // Create a concrete test implementation of the abstract base class class TestOpenAiCompatibleProvider extends BaseOpenAiCompatibleProvider<"test-model"> { - constructor(apiKey: string) { + constructor(apiKey: string, options?: { apiModelId?: string; modelMaxTokens?: number }) { const testModels: Record<"test-model", ModelInfo> = { - "test-model": { - maxTokens: 4096, - contextWindow: 128000, - supportsImages: false, - supportsPromptCache: false, - inputPrice: 0.5, - outputPrice: 1.5, - }, + "test-model": TEST_MODEL_INFO, } super({ @@ -45,6 +47,7 @@ class TestOpenAiCompatibleProvider extends BaseOpenAiCompatibleProvider<"test-mo defaultProviderModelId: "test-model", providerModels: testModels, apiKey, + ...options, }) } } @@ -297,6 +300,64 @@ describe("BaseOpenAiCompatibleProvider", () => { }) }) + describe("getModel", () => { + it("returns the provider default when no model id is configured", () => { + const model = handler.getModel() + expect(model.id).toBe("test-model") + expect(model.info).toEqual(TEST_MODEL_INFO) + }) + + it("returns the predefined metadata for a known model id", () => { + const knownHandler = new TestOpenAiCompatibleProvider("test-api-key", { apiModelId: "test-model" }) + const model = knownHandler.getModel() + expect(model.id).toBe("test-model") + expect(model.info).toEqual(TEST_MODEL_INFO) + }) + + it("honors a custom model id that is not in the predefined list", () => { + // Regression: a user-supplied custom id must not be silently swapped for the + // provider default (which previously caused wrong-model requests / 404s). + const customHandler = new TestOpenAiCompatibleProvider("test-api-key", { + apiModelId: "some/custom-model-not-in-list", + }) + const model = customHandler.getModel() + expect(model.id).toBe("some/custom-model-not-in-list") + expect(model.id).not.toBe("test-model") + // Falls back to sane default metadata so the rest of the pipeline works. + expect(model.info).toEqual(openAiModelInfoSaneDefaults) + }) + + it("sends the custom model id verbatim to the API", async () => { + mockCreate.mockImplementationOnce(() => asyncStreamFrom([])) + + const customHandler = new TestOpenAiCompatibleProvider("test-api-key", { + apiModelId: "some/custom-model-not-in-list", + }) + + const stream = customHandler.createMessage("system prompt", []) + await collectStream(stream) + + expect(mockCreate).toHaveBeenCalledWith( + expect.objectContaining({ model: "some/custom-model-not-in-list" }), + undefined, + ) + }) + + it("completePrompt sends the custom model id verbatim to the API", async () => { + mockCreate.mockResolvedValueOnce({ choices: [{ message: { content: "ok" } }] }) + + const customHandler = new TestOpenAiCompatibleProvider("test-api-key", { + apiModelId: "some/custom-model-not-in-list", + }) + + await customHandler.completePrompt("hello") + + expect(mockCreate).toHaveBeenCalledWith( + expect.objectContaining({ model: "some/custom-model-not-in-list" }), + ) + }) + }) + describe("Tool call handling", () => { it("should yield tool_call_end events when finish_reason is tool_calls", async () => { mockCreate.mockImplementationOnce(() => diff --git a/src/api/providers/base-openai-compatible-provider.ts b/src/api/providers/base-openai-compatible-provider.ts index 1cf80784d0..63cf7ccb29 100644 --- a/src/api/providers/base-openai-compatible-provider.ts +++ b/src/api/providers/base-openai-compatible-provider.ts @@ -1,7 +1,7 @@ import { Anthropic } from "@anthropic-ai/sdk" import OpenAI from "openai" -import type { ModelInfo } from "@roo-code/types" +import { type ModelInfo, openAiModelInfoSaneDefaults } from "@roo-code/types" import { type ApiHandlerOptions, getModelMaxOutputTokens } from "../../shared/api" import { TagMatcher } from "../../utils/tag-matcher" @@ -243,11 +243,27 @@ export abstract class BaseOpenAiCompatibleProvider } override getModel() { - const id = - this.options.apiModelId && this.options.apiModelId in this.providerModels - ? (this.options.apiModelId as ModelName) - : this.defaultProviderModelId + const requestedId = this.options.apiModelId - return { id, info: this.providerModels[id] } + // A known model: use its predefined metadata. + if (requestedId && requestedId in this.providerModels) { + const id = requestedId as ModelName + return { id, info: this.providerModels[id] } + } + + // A user-supplied custom model that isn't in our static list (e.g. a newly + // released Fireworks model). Honor the exact id the user configured instead + // of silently falling back to the provider default, which would send the + // wrong model to the API and can surface as a confusing "model not found" + // error. Provide sane default metadata so the rest of the pipeline works. + if (requestedId) { + return { + id: requestedId as ModelName, + info: { ...openAiModelInfoSaneDefaults }, + } + } + + // No model configured: use the provider default. + return { id: this.defaultProviderModelId, info: this.providerModels[this.defaultProviderModelId] } } } 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..1b01719dc9 100644 --- a/webview-ui/src/components/ui/hooks/__tests__/useSelectedModel.spec.ts +++ b/webview-ui/src/components/ui/hooks/__tests__/useSelectedModel.spec.ts @@ -1413,7 +1413,9 @@ describe("useSelectedModel", () => { expect(result.current.info).toEqual(getZAiModels("international_api")["glm-5.3"]) }) - it("falls back when GLM-5.3 is unavailable on the China API", () => { + it("displays the configured model ID with sane defaults when it is not in the API line catalog", () => { + expect(getZAiModels("china_api")["glm-5.3"]).toBeUndefined() + const apiConfiguration: ProviderSettings = { apiProvider: providerIdentifiers.zai, apiModelId: "glm-5.3", @@ -1422,6 +1424,43 @@ describe("useSelectedModel", () => { const { result } = renderHook(() => useSelectedModel(apiConfiguration), { wrapper: createWrapper() }) + expect(result.current.id).toBe("glm-5.3") + expect(result.current.info).toEqual(openAiModelInfoSaneDefaults) + }) + + it("displays a custom model ID that is in no Z.ai catalog", () => { + const apiConfiguration: ProviderSettings = { + apiProvider: providerIdentifiers.zai, + apiModelId: "glm-custom-preview", + } + + const { result } = renderHook(() => useSelectedModel(apiConfiguration), { wrapper: createWrapper() }) + + expect(result.current.id).toBe("glm-custom-preview") + expect(result.current.info).toEqual(openAiModelInfoSaneDefaults) + }) + + it("falls back to the API line default when no model is configured", () => { + const apiConfiguration: ProviderSettings = { + apiProvider: providerIdentifiers.zai, + zaiApiLine: "china_api", + } + + const { result } = renderHook(() => useSelectedModel(apiConfiguration), { wrapper: createWrapper() }) + + expect(result.current.id).toBe(mainlandZAiDefaultModelId) + expect(result.current.info).toEqual(getZAiModels("china_api")[mainlandZAiDefaultModelId]) + }) + + it("falls back to the API line default when the configured model ID is empty", () => { + const apiConfiguration: ProviderSettings = { + apiProvider: providerIdentifiers.zai, + apiModelId: "", + zaiApiLine: "china_api", + } + + const { result } = renderHook(() => useSelectedModel(apiConfiguration), { wrapper: createWrapper() }) + expect(result.current.id).toBe(mainlandZAiDefaultModelId) expect(result.current.info).toEqual(getZAiModels("china_api")[mainlandZAiDefaultModelId]) }) diff --git a/webview-ui/src/components/ui/hooks/useSelectedModel.ts b/webview-ui/src/components/ui/hooks/useSelectedModel.ts index 8d4b70ad4a..e4a863088d 100644 --- a/webview-ui/src/components/ui/hooks/useSelectedModel.ts +++ b/webview-ui/src/components/ui/hooks/useSelectedModel.ts @@ -326,9 +326,11 @@ function getSelectedModel({ const isChina = zaiApiLineConfigs[apiLine].isChina const models = getZAiModels(apiLine) const defaultModelId = getProviderDefaultModelId(provider, { isChina }) - const id = getValidatedModelId(apiConfiguration.apiModelId, models, defaultModelId) - const info = models[id] - return { id, info } + const configuredId = apiConfiguration.apiModelId + if (configuredId) { + return { id: configuredId, info: models[configuredId] ?? openAiModelInfoSaneDefaults } + } + return { id: defaultModelId, info: models[defaultModelId] } } case providerIdentifiers.openaiNative: { const id = apiConfiguration.apiModelId ?? defaultModelId