Skip to content

Commit c52f118

Browse files
authored
refactor(types): use canonical identifiers for default models (#991)
* refactor(types): use canonical identifiers for default models * test(types): cover international ZAI default * refactor(types): name absent provider default * test(types): cover defaults preserved during rebase * test(types): document identical Zai defaults
1 parent 772e716 commit c52f118

2 files changed

Lines changed: 102 additions & 37 deletions

File tree

Lines changed: 63 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,63 @@
1+
import type { ProviderName } from "../provider-settings.js"
2+
3+
vi.mock("../provider-identifiers.js", async (importOriginal) => {
4+
const actual = await importOriginal<typeof import("../provider-identifiers.js")>()
5+
6+
return {
7+
...actual,
8+
providerIdentifiers: {
9+
...actual.providerIdentifiers,
10+
openrouter: "canonical-openrouter-test-value",
11+
},
12+
}
13+
})
14+
15+
import { providerIdentifiers } from "../provider-identifiers.js"
16+
import {
17+
anthropicDefaultModelId,
18+
getProviderDefaultModelId,
19+
internationalZAiDefaultModelId,
20+
kimiCodeDefaultModelId,
21+
mainlandZAiDefaultModelId,
22+
openRouterDefaultModelId,
23+
vscodeLlmDefaultModelId,
24+
zooGatewayDefaultModelId,
25+
} from "../providers/index.js"
26+
27+
describe("getProviderDefaultModelId", () => {
28+
it("selects a static default through the canonical provider identifier", () => {
29+
expect(getProviderDefaultModelId(providerIdentifiers.openrouter as ProviderName)).toBe(openRouterDefaultModelId)
30+
})
31+
32+
it("triangulates static selection with another provider category", () => {
33+
expect(getProviderDefaultModelId(providerIdentifiers.vscodeLm)).toBe(vscodeLlmDefaultModelId)
34+
})
35+
36+
it.each([
37+
[providerIdentifiers.kimiCode, kimiCodeDefaultModelId],
38+
[providerIdentifiers.zooGateway, zooGatewayDefaultModelId],
39+
])("preserves the %s default added on main", (provider, expectedModelId) => {
40+
expect(getProviderDefaultModelId(provider)).toBe(expectedModelId)
41+
})
42+
43+
it("preserves region-dependent defaults", () => {
44+
// These defaults currently share the same model ID, so the assertions document
45+
// both branches but cannot detect swapped ternary arms until the IDs diverge.
46+
expect(getProviderDefaultModelId(providerIdentifiers.zai, { isChina: true })).toBe(mainlandZAiDefaultModelId)
47+
expect(getProviderDefaultModelId(providerIdentifiers.zai)).toBe(internationalZAiDefaultModelId)
48+
})
49+
50+
it.each([providerIdentifiers.openai, providerIdentifiers.ollama, providerIdentifiers.lmstudio])(
51+
"returns an empty default for custom or locally selected models from %s",
52+
(provider) => {
53+
expect(getProviderDefaultModelId(provider)).toBe("")
54+
},
55+
)
56+
57+
it.each([providerIdentifiers.anthropic, providerIdentifiers.geminiCli, providerIdentifiers.fakeAi])(
58+
"preserves the Anthropic fallback for %s",
59+
(provider) => {
60+
expect(getProviderDefaultModelId(provider)).toBe(anthropicDefaultModelId)
61+
},
62+
)
63+
})

packages/types/src/providers/index.ts

Lines changed: 39 additions & 37 deletions
Original file line numberDiff line numberDiff line change
@@ -62,6 +62,9 @@ import { zooGatewayDefaultModelId } from "./zoo-gateway.js"
6262

6363
// Import the ProviderName type from provider-settings to avoid duplication
6464
import type { ProviderName } from "../provider-settings.js"
65+
import { providerIdentifiers } from "../provider-identifiers.js"
66+
67+
const NO_DEFAULT_MODEL_ID = ""
6568

6669
/**
6770
* Get the default model ID for a given provider.
@@ -73,71 +76,70 @@ export function getProviderDefaultModelId(
7376
options: { isChina?: boolean } = { isChina: false },
7477
): string {
7578
switch (provider) {
76-
case "openrouter":
79+
case providerIdentifiers.openrouter:
7780
return openRouterDefaultModelId
78-
case "requesty":
81+
case providerIdentifiers.requesty:
7982
return requestyDefaultModelId
80-
case "litellm":
83+
case providerIdentifiers.litellm:
8184
return litellmDefaultModelId
82-
case "xai":
85+
case providerIdentifiers.xai:
8386
return xaiDefaultModelId
84-
case "baseten":
87+
case providerIdentifiers.baseten:
8588
return basetenDefaultModelId
86-
case "bedrock":
89+
case providerIdentifiers.bedrock:
8790
return bedrockDefaultModelId
88-
case "vertex":
91+
case providerIdentifiers.vertex:
8992
return vertexDefaultModelId
90-
case "gemini":
93+
case providerIdentifiers.gemini:
9194
return geminiDefaultModelId
92-
case "deepseek":
95+
case providerIdentifiers.deepseek:
9396
return deepSeekDefaultModelId
94-
case "moonshot":
97+
case providerIdentifiers.moonshot:
9598
return moonshotDefaultModelId
96-
case "minimax":
99+
case providerIdentifiers.minimax:
97100
return minimaxDefaultModelId
98-
case "mimo":
101+
case providerIdentifiers.mimo:
99102
return mimoDefaultModelId
100-
case "zai":
103+
case providerIdentifiers.zai:
101104
return options?.isChina ? mainlandZAiDefaultModelId : internationalZAiDefaultModelId
102-
case "openai-native":
105+
case providerIdentifiers.openaiNative:
106+
// TODO(#992): Replace this stale fallback with openAiNativeDefaultModelId.
103107
return "gpt-4o" // Based on openai-native patterns
104-
case "openai-codex":
108+
case providerIdentifiers.openaiCodex:
105109
return openAiCodexDefaultModelId
106-
case "mistral":
110+
case providerIdentifiers.mistral:
107111
return mistralDefaultModelId
108-
case "openai":
109-
return "" // OpenAI provider uses custom model configuration
110-
case "ollama":
111-
return "" // Ollama uses dynamic model selection
112-
case "lmstudio":
113-
return "" // LMStudio uses dynamic model selection
114-
case "vscode-lm":
112+
case providerIdentifiers.openai:
113+
case providerIdentifiers.ollama:
114+
case providerIdentifiers.lmstudio:
115+
return NO_DEFAULT_MODEL_ID
116+
case providerIdentifiers.vscodeLm:
115117
return vscodeLlmDefaultModelId
116-
case "sambanova":
118+
case providerIdentifiers.sambanova:
117119
return sambaNovaDefaultModelId
118-
case "fireworks":
120+
case providerIdentifiers.fireworks:
119121
return fireworksDefaultModelId
120-
case "friendli":
122+
case providerIdentifiers.friendli:
121123
return friendliDefaultModelId
122-
case "qwen-code":
124+
case providerIdentifiers.qwenCode:
123125
return qwenCodeDefaultModelId
124-
case "poe":
126+
case providerIdentifiers.poe:
125127
return poeDefaultModelId
126-
case "unbound":
128+
case providerIdentifiers.unbound:
127129
return unboundDefaultModelId
128-
case "vercel-ai-gateway":
130+
case providerIdentifiers.vercelAiGateway:
129131
return vercelAiGatewayDefaultModelId
130-
case "opencode-go":
132+
case providerIdentifiers.opencodeGo:
131133
return opencodeGoDefaultModelId
132-
case "kenari":
134+
case providerIdentifiers.kenari:
133135
return kenariDefaultModelId
134-
case "kimi-code":
136+
case providerIdentifiers.kimiCode:
135137
return kimiCodeDefaultModelId
136-
case "zoo-gateway":
138+
case providerIdentifiers.zooGateway:
137139
return zooGatewayDefaultModelId
138-
case "anthropic":
139-
case "gemini-cli":
140-
case "fake-ai":
140+
case providerIdentifiers.anthropic:
141+
case providerIdentifiers.geminiCli:
142+
case providerIdentifiers.fakeAi:
141143
default:
142144
return anthropicDefaultModelId
143145
}

0 commit comments

Comments
 (0)