Skip to content
63 changes: 63 additions & 0 deletions packages/types/src/__tests__/provider-default-model.test.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,63 @@
import type { ProviderName } from "../provider-settings.js"

vi.mock("../provider-identifiers.js", async (importOriginal) => {
const actual = await importOriginal<typeof import("../provider-identifiers.js")>()

return {
...actual,
providerIdentifiers: {
...actual.providerIdentifiers,
openrouter: "canonical-openrouter-test-value",
},
}
})

import { providerIdentifiers } from "../provider-identifiers.js"
import {
anthropicDefaultModelId,
getProviderDefaultModelId,
internationalZAiDefaultModelId,
kimiCodeDefaultModelId,
mainlandZAiDefaultModelId,
openRouterDefaultModelId,
vscodeLlmDefaultModelId,
zooGatewayDefaultModelId,
} from "../providers/index.js"

describe("getProviderDefaultModelId", () => {
it("selects a static default through the canonical provider identifier", () => {
expect(getProviderDefaultModelId(providerIdentifiers.openrouter as ProviderName)).toBe(openRouterDefaultModelId)
})

it("triangulates static selection with another provider category", () => {
expect(getProviderDefaultModelId(providerIdentifiers.vscodeLm)).toBe(vscodeLlmDefaultModelId)
})

it.each([
[providerIdentifiers.kimiCode, kimiCodeDefaultModelId],
[providerIdentifiers.zooGateway, zooGatewayDefaultModelId],
])("preserves the %s default added on main", (provider, expectedModelId) => {
expect(getProviderDefaultModelId(provider)).toBe(expectedModelId)
})

it("preserves region-dependent defaults", () => {
// These defaults currently share the same model ID, so the assertions document
// both branches but cannot detect swapped ternary arms until the IDs diverge.
expect(getProviderDefaultModelId(providerIdentifiers.zai, { isChina: true })).toBe(mainlandZAiDefaultModelId)
expect(getProviderDefaultModelId(providerIdentifiers.zai)).toBe(internationalZAiDefaultModelId)
})

it.each([providerIdentifiers.openai, providerIdentifiers.ollama, providerIdentifiers.lmstudio])(
"returns an empty default for custom or locally selected models from %s",
(provider) => {
expect(getProviderDefaultModelId(provider)).toBe("")
},
)

it.each([providerIdentifiers.anthropic, providerIdentifiers.geminiCli, providerIdentifiers.fakeAi])(
"preserves the Anthropic fallback for %s",
(provider) => {
expect(getProviderDefaultModelId(provider)).toBe(anthropicDefaultModelId)
},
)
})
1 change: 1 addition & 0 deletions packages/types/src/__tests__/provider-identifiers.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -106,6 +106,7 @@ describe("provider identifiers", () => {
providerIdentifiers.opencodeGo,
providerIdentifiers.kenari,
providerIdentifiers.kimiCode,
providerIdentifiers.friendli,
])
expect(localProviders).toEqual([providerIdentifiers.ollama, providerIdentifiers.lmstudio])
expect(internalProviders).toEqual([providerIdentifiers.vscodeLm])
Expand Down
1 change: 1 addition & 0 deletions packages/types/src/provider-settings.ts
Original file line number Diff line number Diff line change
Expand Up @@ -56,6 +56,7 @@ export const dynamicProviders = [
providerIdentifiers.opencodeGo,
providerIdentifiers.kenari,
providerIdentifiers.kimiCode,
providerIdentifiers.friendli,
] as const

export type DynamicProvider = (typeof dynamicProviders)[number]
Expand Down
12 changes: 10 additions & 2 deletions packages/types/src/providers/friendli.ts
Original file line number Diff line number Diff line change
Expand Up @@ -8,8 +8,12 @@ export type FriendliModelId =

export const friendliDefaultModelId: FriendliModelId = "zai-org/GLM-5.2"

// Static fallback for the Friendli provider. Used as a fallback when dynamic
// models cannot be fetched (cold start, network errors, API lag), in tests,
// and in the webview's MODELS_BY_PROVIDER fallback. The provider itself fetches
// the live list from https://api.friendli.ai/serverless/v1/models at runtime.
// Pricing sourced from https://friendli.ai/api/public/model-apis (per 1M tokens).
export const friendliModels = {
export const friendliModels: Record<string, ModelInfo> = {
"zai-org/GLM-5.2": {
maxTokens: 131_072,
contextWindow: 1_000_000,
Expand All @@ -20,6 +24,8 @@ export const friendliModels = {
outputPrice: 4.4,
cacheWritesPrice: 0,
cacheReadsPrice: 0.26,
supportsReasoningEffort: ["minimal", "low", "medium", "high", "xhigh", "max"],
reasoningEffort: "high",
description:
"GLM-5.2 is Zhipu's flagship model with a 1M context window and 128k max output, served via Friendli Model APIs. It delivers top-tier long-context reasoning, coding, and agentic performance for extended engineering sessions.",
},
Expand All @@ -33,6 +39,8 @@ export const friendliModels = {
outputPrice: 4.4,
cacheWritesPrice: 0,
cacheReadsPrice: 0.26,
supportsReasoningEffort: ["minimal", "low", "medium", "high", "xhigh", "max"],
reasoningEffort: "high",
description:
"GLM-5.1 is Zhipu's most capable model with a 200k context window and 128k max output, served via Friendli Model APIs. It delivers top-tier reasoning, coding, and agentic performance.",
},
Expand Down Expand Up @@ -60,4 +68,4 @@ export const friendliModels = {
description:
"MiniMax M2.5 is a high-performance language model with a 204.8K context window, optimized for long-context understanding and generation tasks, served via Friendli Model APIs.",
},
} as const satisfies Record<string, ModelInfo>
}
76 changes: 39 additions & 37 deletions packages/types/src/providers/index.ts
Original file line number Diff line number Diff line change
Expand Up @@ -62,6 +62,9 @@ import { zooGatewayDefaultModelId } from "./zoo-gateway.js"

// Import the ProviderName type from provider-settings to avoid duplication
import type { ProviderName } from "../provider-settings.js"
import { providerIdentifiers } from "../provider-identifiers.js"

const NO_DEFAULT_MODEL_ID = ""

/**
* Get the default model ID for a given provider.
Expand All @@ -73,71 +76,70 @@ export function getProviderDefaultModelId(
options: { isChina?: boolean } = { isChina: false },
): string {
switch (provider) {
case "openrouter":
case providerIdentifiers.openrouter:
return openRouterDefaultModelId
case "requesty":
case providerIdentifiers.requesty:
return requestyDefaultModelId
case "litellm":
case providerIdentifiers.litellm:
return litellmDefaultModelId
case "xai":
case providerIdentifiers.xai:
return xaiDefaultModelId
case "baseten":
case providerIdentifiers.baseten:
return basetenDefaultModelId
case "bedrock":
case providerIdentifiers.bedrock:
return bedrockDefaultModelId
case "vertex":
case providerIdentifiers.vertex:
return vertexDefaultModelId
case "gemini":
case providerIdentifiers.gemini:
return geminiDefaultModelId
case "deepseek":
case providerIdentifiers.deepseek:
return deepSeekDefaultModelId
case "moonshot":
case providerIdentifiers.moonshot:
return moonshotDefaultModelId
case "minimax":
case providerIdentifiers.minimax:
return minimaxDefaultModelId
case "mimo":
case providerIdentifiers.mimo:
return mimoDefaultModelId
case "zai":
case providerIdentifiers.zai:
return options?.isChina ? mainlandZAiDefaultModelId : internationalZAiDefaultModelId
case "openai-native":
case providerIdentifiers.openaiNative:
// TODO(#992): Replace this stale fallback with openAiNativeDefaultModelId.
return "gpt-4o" // Based on openai-native patterns
case "openai-codex":
case providerIdentifiers.openaiCodex:
return openAiCodexDefaultModelId
case "mistral":
case providerIdentifiers.mistral:
return mistralDefaultModelId
case "openai":
return "" // OpenAI provider uses custom model configuration
case "ollama":
return "" // Ollama uses dynamic model selection
case "lmstudio":
return "" // LMStudio uses dynamic model selection
case "vscode-lm":
case providerIdentifiers.openai:
case providerIdentifiers.ollama:
case providerIdentifiers.lmstudio:
return NO_DEFAULT_MODEL_ID
case providerIdentifiers.vscodeLm:
return vscodeLlmDefaultModelId
case "sambanova":
case providerIdentifiers.sambanova:
return sambaNovaDefaultModelId
case "fireworks":
case providerIdentifiers.fireworks:
return fireworksDefaultModelId
case "friendli":
case providerIdentifiers.friendli:
return friendliDefaultModelId
case "qwen-code":
case providerIdentifiers.qwenCode:
return qwenCodeDefaultModelId
case "poe":
case providerIdentifiers.poe:
return poeDefaultModelId
case "unbound":
case providerIdentifiers.unbound:
return unboundDefaultModelId
case "vercel-ai-gateway":
case providerIdentifiers.vercelAiGateway:
return vercelAiGatewayDefaultModelId
case "opencode-go":
case providerIdentifiers.opencodeGo:
return opencodeGoDefaultModelId
case "kenari":
case providerIdentifiers.kenari:
return kenariDefaultModelId
case "kimi-code":
case providerIdentifiers.kimiCode:
return kimiCodeDefaultModelId
case "zoo-gateway":
case providerIdentifiers.zooGateway:
return zooGatewayDefaultModelId
case "anthropic":
case "gemini-cli":
case "fake-ai":
case providerIdentifiers.anthropic:
case providerIdentifiers.geminiCli:
case providerIdentifiers.fakeAi:
default:
return anthropicDefaultModelId
}
Expand Down
146 changes: 136 additions & 10 deletions src/api/__tests__/index.spec.ts
Original file line number Diff line number Diff line change
Expand Up @@ -9,21 +9,147 @@ vitest.mock("vscode", () => ({
},
}))

import type { ProviderSettings } from "@roo-code/types"
// Handler constructors can require credentials or initialize SDK clients. Replace them
// with inert classes so these tests exercise only the factory's routing behavior.
vitest.mock("../providers", async () => {
const providers = await vitest.importActual<Record<string, unknown>>("../providers")

return Object.fromEntries(Object.keys(providers).map((name) => [name, class {}]))
})

vitest.mock("../providers/native-ollama", () => ({
NativeOllamaHandler: class {},
}))

import {
providerIdentifiers,
retiredProviderIdentifiers,
type ProviderName,
type ProviderNameWithRetired,
} from "@roo-code/types"

import { buildApiHandler } from "../index"
import { KenariHandler } from "../providers/kenari"
import {
AnthropicHandler,
AnthropicVertexHandler,
AwsBedrockHandler,
BasetenHandler,
DeepSeekHandler,
FakeAIHandler,
FireworksHandler,
FriendliHandler,
GeminiHandler,
KenariHandler,
KimiCodeHandler,
LiteLLMHandler,
LmStudioHandler,
MiniMaxHandler,
MimoHandler,
MistralHandler,
MoonshotHandler,
OpenAiCodexHandler,
OpenAiHandler,
OpenAiNativeHandler,
OpencodeGoHandler,
OpenRouterHandler,
PoeHandler,
QwenCodeHandler,
RequestyHandler,
SambaNovaHandler,
UnboundHandler,
VercelAiGatewayHandler,
VertexHandler,
VsCodeLmHandler,
XAIHandler,
ZAiHandler,
ZooGatewayHandler,
} from "../providers"
import { NativeOllamaHandler } from "../providers/native-ollama"

type HandlerConstructor = new (...args: never[]) => object

const expectedHandlers = {
[providerIdentifiers.anthropic]: AnthropicHandler,
[providerIdentifiers.openrouter]: OpenRouterHandler,
[providerIdentifiers.bedrock]: AwsBedrockHandler,
[providerIdentifiers.openai]: OpenAiHandler,
[providerIdentifiers.ollama]: NativeOllamaHandler,
[providerIdentifiers.lmstudio]: LmStudioHandler,
[providerIdentifiers.gemini]: GeminiHandler,
// Gemini CLI currently relies on the factory's default Anthropic handler.
[providerIdentifiers.geminiCli]: AnthropicHandler,
[providerIdentifiers.openaiCodex]: OpenAiCodexHandler,
[providerIdentifiers.openaiNative]: OpenAiNativeHandler,
[providerIdentifiers.deepseek]: DeepSeekHandler,
[providerIdentifiers.qwenCode]: QwenCodeHandler,
[providerIdentifiers.moonshot]: MoonshotHandler,
[providerIdentifiers.kimiCode]: KimiCodeHandler,
[providerIdentifiers.vscodeLm]: VsCodeLmHandler,
[providerIdentifiers.mistral]: MistralHandler,
[providerIdentifiers.requesty]: RequestyHandler,
[providerIdentifiers.unbound]: UnboundHandler,
[providerIdentifiers.fakeAi]: FakeAIHandler,
[providerIdentifiers.xai]: XAIHandler,
[providerIdentifiers.litellm]: LiteLLMHandler,
[providerIdentifiers.sambanova]: SambaNovaHandler,
[providerIdentifiers.mimo]: MimoHandler,
[providerIdentifiers.zai]: ZAiHandler,
[providerIdentifiers.fireworks]: FireworksHandler,
[providerIdentifiers.friendli]: FriendliHandler,
[providerIdentifiers.vercelAiGateway]: VercelAiGatewayHandler,
[providerIdentifiers.opencodeGo]: OpencodeGoHandler,
[providerIdentifiers.kenari]: KenariHandler,
[providerIdentifiers.zooGateway]: ZooGatewayHandler,
[providerIdentifiers.minimax]: MiniMaxHandler,
[providerIdentifiers.baseten]: BasetenHandler,
[providerIdentifiers.poe]: PoeHandler,
} satisfies Record<Exclude<ProviderName, typeof providerIdentifiers.vertex>, HandlerConstructor>

const expectedHandlerEntries = Object.entries(expectedHandlers) as Array<
[Exclude<ProviderName, typeof providerIdentifiers.vertex>, HandlerConstructor]
>

describe("buildApiHandler", () => {
it("returns a KenariHandler for the kenari provider", () => {
const configuration: ProviderSettings = {
apiProvider: "kenari",
kenariApiKey: "test-key",
kenariModelId: "glm-5-2",
}
it.each(expectedHandlerEntries)("returns the expected handler for %s", (apiProvider, Handler) => {
const handler = buildApiHandler({ apiProvider })

expect(handler).toBeInstanceOf(Handler)
})

it.each([
["an unspecified model", undefined, VertexHandler],
["non-Claude models", "non-claude-test-model", VertexHandler],
["Claude models", "claude-test-model", AnthropicVertexHandler],
] as const)("returns the expected Vertex handler for %s", (_description, apiModelId, Handler) => {
const handler = buildApiHandler({
apiProvider: providerIdentifiers.vertex,
apiModelId,
})

expect(handler).toBeInstanceOf(Handler)
})

it("preserves the dedicated removal error for the retired Roo provider", () => {
expect(() =>
buildApiHandler({
apiProvider: retiredProviderIdentifiers.roo,
}),
).toThrow("Roo Code Router has been removed")
})

it("rejects other retired providers", () => {
expect(() =>
buildApiHandler({
apiProvider: retiredProviderIdentifiers.cerebras,
}),
).toThrow("this provider is no longer supported")
})

const handler = buildApiHandler(configuration)
it("falls back to Anthropic for an unsupported provider value", () => {
const handler = buildApiHandler({
apiProvider: "unsupported-provider" as ProviderNameWithRetired,
})

expect(handler).toBeInstanceOf(KenariHandler)
expect(handler).toBeInstanceOf(AnthropicHandler)
})
})
Loading