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 @@ -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"
Expand All @@ -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({
Expand All @@ -45,6 +47,7 @@ class TestOpenAiCompatibleProvider extends BaseOpenAiCompatibleProvider<"test-mo
defaultProviderModelId: "test-model",
providerModels: testModels,
apiKey,
...options,
})
}
}
Expand Down Expand Up @@ -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(() =>
Expand Down
28 changes: 22 additions & 6 deletions src/api/providers/base-openai-compatible-provider.ts
Original file line number Diff line number Diff line change
@@ -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"
Expand Down Expand Up @@ -243,11 +243,27 @@ export abstract class BaseOpenAiCompatibleProvider<ModelName extends string>
}

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 {
Comment thread
p12tic marked this conversation as resolved.
id: requestedId as ModelName,
info: { ...openAiModelInfoSaneDefaults },
Comment thread
p12tic marked this conversation as resolved.
}
}

// No model configured: use the provider default.
return { id: this.defaultProviderModelId, info: this.providerModels[this.defaultProviderModelId] }
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand All @@ -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])
})
Expand Down
8 changes: 5 additions & 3 deletions webview-ui/src/components/ui/hooks/useSelectedModel.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Loading