diff --git a/packages/types/src/providers/openai.ts b/packages/types/src/providers/openai.ts index b5db12bb70..41d00c7a04 100644 --- a/packages/types/src/providers/openai.ts +++ b/packages/types/src/providers/openai.ts @@ -930,10 +930,11 @@ export const openAiNativeModels = { }, } as const satisfies Record +// `maxTokens` is intentionally omitted so that no max output token limit is sent to the +// provider by default. export const openAiModelInfoSaneDefaults: ModelInfo = { - maxTokens: -1, contextWindow: 128_000, - supportsImages: true, + supportsImages: false, supportsPromptCache: false, inputPrice: 0, outputPrice: 0, 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..b251168461 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,77 @@ 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) + expect(model.info).not.toHaveProperty("maxTokens") + }) + + 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("omits max_tokens for a custom model", async () => { + mockCreate.mockImplementationOnce(() => asyncStreamFrom([])) + + const customHandler = new TestOpenAiCompatibleProvider("test-api-key", { + apiModelId: "some/custom-model-not-in-list", + }) + + await collectStream(customHandler.createMessage("system prompt", [])) + + expect(mockCreate.mock.calls[0][0].max_tokens).toBeUndefined() + }) + + 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/__tests__/fireworks.spec.ts b/src/api/providers/__tests__/fireworks.spec.ts index de56c5a803..480f965894 100644 --- a/src/api/providers/__tests__/fireworks.spec.ts +++ b/src/api/providers/__tests__/fireworks.spec.ts @@ -98,6 +98,29 @@ describe("FireworksHandler", () => { expect(model.info).toEqual(expect.objectContaining(fireworksModels[testModelId])) }) + it("should omit max_tokens for a custom model", async () => { + const handlerWithCustomModel = new FireworksHandler({ + apiModelId: "accounts/fireworks/models/deepseek-v4p1-flash", + fireworksApiKey: "test-fireworks-api-key", + }) + + await collectStream(handlerWithCustomModel.createMessage("system prompt", [])) + + expect(mockCreate.mock.calls[0][0].model).toBe("accounts/fireworks/models/deepseek-v4p1-flash") + expect(mockCreate.mock.calls[0][0].max_tokens).toBeUndefined() + }) + + it("should still send a positive max_tokens for a known model", async () => { + const handlerWithModel = new FireworksHandler({ + apiModelId: "accounts/fireworks/models/kimi-k2-instruct", + fireworksApiKey: "test-fireworks-api-key", + }) + + await collectStream(handlerWithModel.createMessage("system prompt", [])) + + expect(mockCreate.mock.calls[0][0].max_tokens).toBe(16384) + }) + it.each([ { modelId: "accounts/fireworks/models/glm-5p1" as const, diff --git a/src/api/providers/__tests__/lmstudio.spec.ts b/src/api/providers/__tests__/lmstudio.spec.ts index 7ab674a0a9..ae773afb40 100644 --- a/src/api/providers/__tests__/lmstudio.spec.ts +++ b/src/api/providers/__tests__/lmstudio.spec.ts @@ -239,7 +239,7 @@ describe("LmStudioHandler", () => { const modelInfo = handler.getModel() expect(modelInfo.id).toBe(mockOptions.lmStudioModelId) expect(modelInfo.info).toBeDefined() - expect(modelInfo.info.maxTokens).toBe(-1) + expect(modelInfo.info.maxTokens).toBeUndefined() expect(modelInfo.info.contextWindow).toBe(128_000) }) }) diff --git a/src/api/providers/__tests__/openai.spec.ts b/src/api/providers/__tests__/openai.spec.ts index 754d57a6cd..08d4139ba6 100644 --- a/src/api/providers/__tests__/openai.spec.ts +++ b/src/api/providers/__tests__/openai.spec.ts @@ -1069,7 +1069,7 @@ describe("OpenAiHandler", () => { expect(model.id).toBe(mockOptions.openAiModelId) expect(model.info).toBeDefined() expect(model.info.contextWindow).toBe(128_000) - expect(model.info.supportsImages).toBe(true) + expect(model.info.supportsImages).toBe(false) }) it("should handle undefined model ID", () => { 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