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
7 changes: 7 additions & 0 deletions packages/types/src/__tests__/provider-identifiers.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@ import {
isProviderName,
isRetiredProvider,
localProviders,
MODELS_BY_PROVIDER,
providerIdentifiers,
providerNames,
providerNamesSchema,
Expand Down Expand Up @@ -113,6 +114,12 @@ describe("provider identifiers", () => {
expect(fauxProviders).toEqual([providerIdentifiers.fakeAi])
})

it("keeps model provider ids aligned with their keys", () => {
for (const [identifier, providerModels] of Object.entries(MODELS_BY_PROVIDER)) {
expect(providerModels.id).toBe(identifier)
}
})

it("preserves provider category type guards", () => {
for (const identifier of dynamicProviders) {
expect(isDynamicProvider(identifier)).toBe(true)
Expand Down
245 changes: 132 additions & 113 deletions packages/types/src/provider-settings.ts
Original file line number Diff line number Diff line change
Expand Up @@ -428,40 +428,40 @@ const defaultSchema = z.object({
})

export const providerSettingsSchemaDiscriminated = z.discriminatedUnion("apiProvider", [
anthropicSchema.merge(z.object({ apiProvider: z.literal("anthropic") })),
openRouterSchema.merge(z.object({ apiProvider: z.literal("openrouter") })),
bedrockSchema.merge(z.object({ apiProvider: z.literal("bedrock") })),
vertexSchema.merge(z.object({ apiProvider: z.literal("vertex") })),
openAiSchema.merge(z.object({ apiProvider: z.literal("openai") })),
ollamaSchema.merge(z.object({ apiProvider: z.literal("ollama") })),
vsCodeLmSchema.merge(z.object({ apiProvider: z.literal("vscode-lm") })),
lmStudioSchema.merge(z.object({ apiProvider: z.literal("lmstudio") })),
geminiSchema.merge(z.object({ apiProvider: z.literal("gemini") })),
geminiCliSchema.merge(z.object({ apiProvider: z.literal("gemini-cli") })),
openAiCodexSchema.merge(z.object({ apiProvider: z.literal("openai-codex") })),
openAiNativeSchema.merge(z.object({ apiProvider: z.literal("openai-native") })),
mistralSchema.merge(z.object({ apiProvider: z.literal("mistral") })),
deepSeekSchema.merge(z.object({ apiProvider: z.literal("deepseek") })),
poeSchema.merge(z.object({ apiProvider: z.literal("poe") })),
moonshotSchema.merge(z.object({ apiProvider: z.literal("moonshot") })),
kimiCodeSchema.merge(z.object({ apiProvider: z.literal("kimi-code") })),
minimaxSchema.merge(z.object({ apiProvider: z.literal("minimax") })),
mimoSchema.merge(z.object({ apiProvider: z.literal("mimo") })),
requestySchema.merge(z.object({ apiProvider: z.literal("requesty") })),
unboundSchema.merge(z.object({ apiProvider: z.literal("unbound") })),
fakeAiSchema.merge(z.object({ apiProvider: z.literal("fake-ai") })),
xaiSchema.merge(z.object({ apiProvider: z.literal("xai") })),
basetenSchema.merge(z.object({ apiProvider: z.literal("baseten") })),
litellmSchema.merge(z.object({ apiProvider: z.literal("litellm") })),
sambaNovaSchema.merge(z.object({ apiProvider: z.literal("sambanova") })),
zaiSchema.merge(z.object({ apiProvider: z.literal("zai") })),
fireworksSchema.merge(z.object({ apiProvider: z.literal("fireworks") })),
friendliSchema.merge(z.object({ apiProvider: z.literal("friendli") })),
qwenCodeSchema.merge(z.object({ apiProvider: z.literal("qwen-code") })),
vercelAiGatewaySchema.merge(z.object({ apiProvider: z.literal("vercel-ai-gateway") })),
opencodeGoSchema.merge(z.object({ apiProvider: z.literal("opencode-go") })),
kenariSchema.merge(z.object({ apiProvider: z.literal("kenari") })),
zooGatewaySchema.merge(z.object({ apiProvider: z.literal("zoo-gateway") })),
anthropicSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.anthropic) })),
openRouterSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.openrouter) })),
bedrockSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.bedrock) })),
vertexSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.vertex) })),
openAiSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.openai) })),
ollamaSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.ollama) })),
vsCodeLmSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.vscodeLm) })),
lmStudioSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.lmstudio) })),
geminiSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.gemini) })),
geminiCliSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.geminiCli) })),
openAiCodexSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.openaiCodex) })),
openAiNativeSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.openaiNative) })),
mistralSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.mistral) })),
deepSeekSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.deepseek) })),
poeSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.poe) })),
moonshotSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.moonshot) })),
kimiCodeSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.kimiCode) })),
minimaxSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.minimax) })),
mimoSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.mimo) })),
requestySchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.requesty) })),
unboundSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.unbound) })),
fakeAiSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.fakeAi) })),
xaiSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.xai) })),
basetenSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.baseten) })),
litellmSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.litellm) })),
sambaNovaSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.sambanova) })),
zaiSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.zai) })),
fireworksSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.fireworks) })),
friendliSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.friendli) })),
qwenCodeSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.qwenCode) })),
vercelAiGatewaySchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.vercelAiGateway) })),
opencodeGoSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.opencodeGo) })),
kenariSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.kenari) })),
zooGatewaySchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.zooGateway) })),
defaultSchema,
])

Expand Down Expand Up @@ -553,37 +553,37 @@ export const isTypicalProvider = (key: unknown): key is TypicalProvider =>
isProviderName(key) && !isInternalProvider(key) && !isCustomProvider(key) && !isFauxProvider(key)

export const modelIdKeysByProvider: Record<TypicalProvider, ModelIdKey> = {
anthropic: "apiModelId",
openrouter: "openRouterModelId",
bedrock: "apiModelId",
vertex: "apiModelId",
"openai-codex": "apiModelId",
"openai-native": "openAiModelId",
ollama: "ollamaModelId",
lmstudio: "lmStudioModelId",
gemini: "apiModelId",
"gemini-cli": "apiModelId",
mistral: "apiModelId",
moonshot: "apiModelId",
"kimi-code": "apiModelId",
minimax: "apiModelId",
mimo: "apiModelId",
deepseek: "apiModelId",
poe: "apiModelId",
"qwen-code": "apiModelId",
requesty: "requestyModelId",
unbound: "unboundModelId",
xai: "apiModelId",
baseten: "apiModelId",
litellm: "litellmModelId",
sambanova: "apiModelId",
zai: "apiModelId",
fireworks: "apiModelId",
friendli: "apiModelId",
"vercel-ai-gateway": "vercelAiGatewayModelId",
"opencode-go": "opencodeGoModelId",
kenari: "kenariModelId",
"zoo-gateway": "zooGatewayModelId",
[providerIdentifiers.anthropic]: "apiModelId",
[providerIdentifiers.openrouter]: "openRouterModelId",
[providerIdentifiers.bedrock]: "apiModelId",
[providerIdentifiers.vertex]: "apiModelId",
[providerIdentifiers.openaiCodex]: "apiModelId",
[providerIdentifiers.openaiNative]: "openAiModelId",
[providerIdentifiers.ollama]: "ollamaModelId",
[providerIdentifiers.lmstudio]: "lmStudioModelId",
[providerIdentifiers.gemini]: "apiModelId",
[providerIdentifiers.geminiCli]: "apiModelId",
[providerIdentifiers.mistral]: "apiModelId",
[providerIdentifiers.moonshot]: "apiModelId",
[providerIdentifiers.kimiCode]: "apiModelId",
[providerIdentifiers.minimax]: "apiModelId",
[providerIdentifiers.mimo]: "apiModelId",
[providerIdentifiers.deepseek]: "apiModelId",
[providerIdentifiers.poe]: "apiModelId",
[providerIdentifiers.qwenCode]: "apiModelId",
[providerIdentifiers.requesty]: "requestyModelId",
[providerIdentifiers.unbound]: "unboundModelId",
[providerIdentifiers.xai]: "apiModelId",
[providerIdentifiers.baseten]: "apiModelId",
[providerIdentifiers.litellm]: "litellmModelId",
[providerIdentifiers.sambanova]: "apiModelId",
[providerIdentifiers.zai]: "apiModelId",
[providerIdentifiers.fireworks]: "apiModelId",
[providerIdentifiers.friendli]: "apiModelId",
[providerIdentifiers.vercelAiGateway]: "vercelAiGatewayModelId",
[providerIdentifiers.opencodeGo]: "opencodeGoModelId",
[providerIdentifiers.kenari]: "kenariModelId",
[providerIdentifiers.zooGateway]: "zooGatewayModelId",
}

/**
Expand Down Expand Up @@ -653,106 +653,125 @@ export const getApiProtocol = (provider: ProviderName | undefined, modelId?: str
*/

export const MODELS_BY_PROVIDER: Record<
Exclude<ProviderName, "fake-ai" | "gemini-cli" | "openai">,
Exclude<
ProviderName,
typeof providerIdentifiers.fakeAi | typeof providerIdentifiers.geminiCli | typeof providerIdentifiers.openai
>,
{ id: ProviderName; label: string; models: string[] }
> = {
anthropic: {
id: "anthropic",
[providerIdentifiers.anthropic]: {
id: providerIdentifiers.anthropic,
label: "Anthropic",
models: Object.keys(anthropicModels),
},
bedrock: {
id: "bedrock",
[providerIdentifiers.bedrock]: {
id: providerIdentifiers.bedrock,
label: "Amazon Bedrock",
models: Object.keys(bedrockModels),
},
deepseek: {
id: "deepseek",
[providerIdentifiers.deepseek]: {
id: providerIdentifiers.deepseek,
label: "DeepSeek",
models: Object.keys(deepSeekModels),
},
fireworks: {
id: "fireworks",
[providerIdentifiers.fireworks]: {
id: providerIdentifiers.fireworks,
label: "Fireworks",
models: Object.keys(fireworksModels),
},
friendli: {
id: "friendli",
[providerIdentifiers.friendli]: {
id: providerIdentifiers.friendli,
label: "Friendli",
models: Object.keys(friendliModels),
},
gemini: {
id: "gemini",
[providerIdentifiers.gemini]: {
id: providerIdentifiers.gemini,
label: "Google Gemini",
models: Object.keys(geminiModels),
},
mistral: {
id: "mistral",
[providerIdentifiers.mistral]: {
id: providerIdentifiers.mistral,
label: "Mistral",
models: Object.keys(mistralModels),
},
moonshot: {
id: "moonshot",
[providerIdentifiers.moonshot]: {
id: providerIdentifiers.moonshot,
label: "Moonshot",
models: Object.keys(moonshotModels),
},
"kimi-code": {
id: "kimi-code",
[providerIdentifiers.kimiCode]: {
id: providerIdentifiers.kimiCode,
label: "Kimi Code",
models: [],
},
minimax: {
id: "minimax",
[providerIdentifiers.minimax]: {
id: providerIdentifiers.minimax,
label: "MiniMax",
models: Object.keys(minimaxModels),
},
mimo: {
id: "mimo",
[providerIdentifiers.mimo]: {
id: providerIdentifiers.mimo,
label: "Xiaomi MiMo",
models: Object.keys(mimoModels),
},
"openai-codex": {
id: "openai-codex",
[providerIdentifiers.openaiCodex]: {
id: providerIdentifiers.openaiCodex,
label: "OpenAI - ChatGPT Plus/Pro",
models: Object.keys(openAiCodexModels),
},
"openai-native": {
id: "openai-native",
[providerIdentifiers.openaiNative]: {
id: providerIdentifiers.openaiNative,
label: "OpenAI",
models: Object.keys(openAiNativeModels),
},
"qwen-code": { id: "qwen-code", label: "Qwen Code", models: Object.keys(qwenCodeModels) },
sambanova: {
id: "sambanova",
[providerIdentifiers.qwenCode]: {
id: providerIdentifiers.qwenCode,
label: "Qwen Code",
models: Object.keys(qwenCodeModels),
},
[providerIdentifiers.sambanova]: {
id: providerIdentifiers.sambanova,
label: "SambaNova",
models: Object.keys(sambaNovaModels),
},
vertex: {
id: "vertex",
[providerIdentifiers.vertex]: {
id: providerIdentifiers.vertex,
label: "GCP Vertex AI",
models: Object.keys(vertexModels),
},
"vscode-lm": {
id: "vscode-lm",
[providerIdentifiers.vscodeLm]: {
id: providerIdentifiers.vscodeLm,
label: "VS Code LM API",
models: Object.keys(vscodeLlmModels),
},
xai: { id: "xai", label: "xAI (Grok)", models: Object.keys(xaiModels) },
zai: { id: "zai", label: "Z.ai", models: Object.keys(internationalZAiModels) },
baseten: { id: "baseten", label: "Baseten", models: Object.keys(basetenModels) },
[providerIdentifiers.xai]: { id: providerIdentifiers.xai, label: "xAI (Grok)", models: Object.keys(xaiModels) },
[providerIdentifiers.zai]: {
id: providerIdentifiers.zai,
label: "Z.ai",
models: Object.keys(internationalZAiModels),
},
[providerIdentifiers.baseten]: {
id: providerIdentifiers.baseten,
label: "Baseten",
models: Object.keys(basetenModels),
},

// Dynamic providers; models pulled from remote APIs.
poe: { id: "poe", label: "Poe", models: [] },
litellm: { id: "litellm", label: "LiteLLM", models: [] },
openrouter: { id: "openrouter", label: "OpenRouter", models: [] },
requesty: { id: "requesty", label: "Requesty", models: [] },
unbound: { id: "unbound", label: "Unbound", models: [] },
"vercel-ai-gateway": { id: "vercel-ai-gateway", label: "Vercel AI Gateway", models: [] },
"opencode-go": { id: "opencode-go", label: "Opencode Go", models: [] },
kenari: { id: "kenari", label: "Kenari", models: [] },
"zoo-gateway": { id: "zoo-gateway", label: "Zoo Gateway", models: [] },
[providerIdentifiers.poe]: { id: providerIdentifiers.poe, label: "Poe", models: [] },
[providerIdentifiers.litellm]: { id: providerIdentifiers.litellm, label: "LiteLLM", models: [] },
[providerIdentifiers.openrouter]: { id: providerIdentifiers.openrouter, label: "OpenRouter", models: [] },
[providerIdentifiers.requesty]: { id: providerIdentifiers.requesty, label: "Requesty", models: [] },
[providerIdentifiers.unbound]: { id: providerIdentifiers.unbound, label: "Unbound", models: [] },
[providerIdentifiers.vercelAiGateway]: {
id: providerIdentifiers.vercelAiGateway,
label: "Vercel AI Gateway",
models: [],
},
[providerIdentifiers.opencodeGo]: { id: providerIdentifiers.opencodeGo, label: "Opencode Go", models: [] },
[providerIdentifiers.kenari]: { id: providerIdentifiers.kenari, label: "Kenari", models: [] },
[providerIdentifiers.zooGateway]: { id: providerIdentifiers.zooGateway, label: "Zoo Gateway", models: [] },

// Local providers; models discovered from localhost endpoints.
lmstudio: { id: "lmstudio", label: "LM Studio", models: [] },
ollama: { id: "ollama", label: "Ollama", models: [] },
[providerIdentifiers.lmstudio]: { id: providerIdentifiers.lmstudio, label: "LM Studio", models: [] },
[providerIdentifiers.ollama]: { id: providerIdentifiers.ollama, label: "Ollama", models: [] },
}
Loading