Skip to content
Merged
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 @@ -466,6 +466,26 @@ describe('applyPreferredProvider', () => {
expect(request.body.provider).toEqual({ order: ['novita'] });
});

it.each(['moonshotai/kimi-k3', 'moonshotai/kimi-k3-fast', 'kimi-k3', 'moonshotai/kimi-k2.5'])(
'prefers Bedrock then Alibaba for Kimi model %s',
model => {
const request = makeRequest(model);

applyPreferredProvider(model, request.body);

expect(request.body.provider).toEqual({ order: ['amazon-bedrock', 'alibaba'] });
}
);

it('preserves explicit Kimi provider order and allowed providers', () => {
const request = makeRequest('moonshotai/kimi-k3');
request.body.provider = { only: ['alibaba'], order: ['alibaba'] };

applyPreferredProvider('moonshotai/kimi-k3', request.body);

expect(request.body.provider).toEqual({ only: ['alibaba'], order: ['alibaba'] });
});

it('prefers Friendli then Novita for GLM models', () => {
const request = makeRequest('z-ai/glm-5.2');

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -151,7 +151,10 @@ export function getPreferredProviderOrder(requestedModel: string): string[] {
return [OpenRouterInferenceProviderIdSchema.enum.mistral];
}
if (isKimiModel(requestedModel)) {
return [OpenRouterInferenceProviderIdSchema.enum.novita];
return [
OpenRouterInferenceProviderIdSchema.enum['amazon-bedrock'],
OpenRouterInferenceProviderIdSchema.enum.alibaba,
];
}
if (isStepModel(requestedModel)) {
return [OpenRouterInferenceProviderIdSchema.enum.stepfun];
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -111,10 +111,15 @@ describe('narrowProviderSlugsToVariant', () => {
});

describe('createModelsByProviderIndexLoader', () => {
function loader(models: Record<string, StoredModel> = storedModels) {
function loader(
models: Record<string, StoredModel> = storedModels,
vercelModels: Record<string, StoredModel> = {},
snapshot = makeSnapshot()
) {
return createModelsByProviderIndexLoader({
fetchSnapshot: async () => makeSnapshot(),
fetchStoredModels: async () => models,
fetchSnapshot: async () => snapshot,
fetchOpenRouterModels: async () => models,
fetchVercelModels: async () => vercelModels,
ttlMs: 60_000,
nowMs: () => 0,
});
Expand Down Expand Up @@ -149,4 +154,71 @@ describe('createModelsByProviderIndexLoader', () => {
await expect(getProviderSlugsForModel('unknown/model')).resolves.toEqual(new Set());
await expect(getProviderSlugsForModel('unknown/model:free')).resolves.toEqual(new Set());
});

it.each(['provider_name', 'tag'] as const)(
'retains Vercel-only Bedrock and Alibaba endpoints identified by %s',
async providerField => {
const modelId = 'moonshotai/kimi-k3';
const snapshot = makeSnapshot();
snapshot.providers = ['novita', 'amazon-bedrock', 'alibaba', 'deepinfra'].map(slug => ({
name: slug,
displayName: slug,
slug,
dataPolicy: { training: false, retainsPrompts: false, canPublish: false },
models: [snapshotModel(modelId, 'standard')],
}));
const { getProviderSlugsForModel } = loader(
{ [modelId]: storedModel(modelId, ['novita/bf16']) },
{
[modelId]: {
id: modelId,
name: modelId,
endpoints: ['bedrock', 'alibaba', 'fireworks'].map(provider => ({
[providerField]: provider,
})),
},
},
snapshot
);

await expect(getProviderSlugsForModel(modelId)).resolves.toEqual(
new Set(['novita', 'amazon-bedrock', 'alibaba'])
);
}
);

it('does not retain paid Vercel providers for a free variant', async () => {
const { getProviderSlugsForModel } = loader(storedModels, {
[MODEL]: { ...storedModel(MODEL, []), endpoints: [{ provider_name: 'deepinfra' }] },
});

await expect(getProviderSlugsForModel(FREE_MODEL)).resolves.toEqual(new Set(['nvidia']));
});

it('maps model and provider IDs when retaining Vercel endpoints', async () => {
const modelId = 'anthropic/claude-sonnet-4-6';
const vercelModelId = 'anthropic/claude-sonnet-4.6';
const snapshot = makeSnapshot();
snapshot.providers = ['anthropic', 'google-vertex'].map(slug => ({
name: slug,
displayName: slug,
slug,
dataPolicy: { training: false, retainsPrompts: false, canPublish: false },
models: [snapshotModel(modelId, 'standard')],
}));
const { getProviderSlugsForModel } = loader(
{ [modelId]: storedModel(modelId, ['anthropic']) },
{
[vercelModelId]: {
...storedModel(vercelModelId, []),
endpoints: [{ provider_name: 'vertexAnthropic' }],
},
},
snapshot
);

await expect(getProviderSlugsForModel(modelId)).resolves.toEqual(
new Set(['anthropic', 'google-vertex'])
);
});
});
Original file line number Diff line number Diff line change
Expand Up @@ -4,9 +4,15 @@ import { readDb } from '@/lib/drizzle';
import { normalizeModelId } from '@/lib/ai-gateway/model-utils';
import {
getOpenRouterModelsMetadataFromDatabase,
getVercelModelsMetadataFromDatabase,
type StoredModelMap,
} from '@/lib/ai-gateway/providers/gateway-models-cache';
import { normalizeInferenceProviderId } from '@/lib/ai-gateway/providers/openrouter/inference-provider-id';
import {
normalizeInferenceProviderId,
normalizeVercelInferenceProviderIdForRouting,
openRouterToVercelInferenceProviderId,
} from '@/lib/ai-gateway/providers/openrouter/inference-provider-id';
import { mapModelIdToVercel } from '@/lib/ai-gateway/providers/vercel/mapModelIdToVercel';
import type {
NormalizedOpenRouterResponse,
OpenRouterModel,
Expand All @@ -24,8 +30,8 @@ export type FetchModelsByProviderSnapshot = () => Promise<NormalizedOpenRouterRe

type ProviderIndexLoaderOptions = {
fetchSnapshot: FetchModelsByProviderSnapshot;
/** Gateway `/models/{id}/endpoints` metadata keyed by exact (variant-suffixed) model id. */
fetchStoredModels: () => Promise<StoredModelMap>;
fetchOpenRouterModels: () => Promise<StoredModelMap>;
fetchVercelModels: () => Promise<StoredModelMap>;
ttlMs: number;
nowMs: () => number;
};
Expand Down Expand Up @@ -139,8 +145,27 @@ export function createModelsByProviderIndexLoader(options: ProviderIndexLoaderOp
const index = await loadIndex();
const snapshotProviderSlugs = index.get(normalizeModelId(modelId));
if (!snapshotProviderSlugs) return new Set();
const storedModels = await options.fetchStoredModels();
return narrowProviderSlugsToVariant(snapshotProviderSlugs, storedModels[modelId]);
const [storedModels, vercelModels] = await Promise.all([
options.fetchOpenRouterModels(),
options.fetchVercelModels(),
]);
const openRouterProviderSlugs = narrowProviderSlugsToVariant(
snapshotProviderSlugs,
storedModels[modelId]
);
const vercelModel = vercelModels[mapModelIdToVercel(modelId)];
const vercelProviders = new Set(
vercelModel?.endpoints.map(endpoint =>
normalizeVercelInferenceProviderIdForRouting(endpoint.provider_name ?? endpoint.tag)
)
);
return new Set(
[...snapshotProviderSlugs].filter(
slug =>
openRouterProviderSlugs.has(slug) ||
vercelProviders.has(openRouterToVercelInferenceProviderId(slug))
)
);
}

return {
Expand All @@ -165,7 +190,8 @@ const DEFAULT_TTL_MS = 30_000;

const defaultLoader = createModelsByProviderIndexLoader({
fetchSnapshot: fetchLatestModelsByProviderSnapshotFromDb,
fetchStoredModels: getOpenRouterModelsMetadataFromDatabase,
fetchOpenRouterModels: getOpenRouterModelsMetadataFromDatabase,
fetchVercelModels: getVercelModelsMetadataFromDatabase,
ttlMs: DEFAULT_TTL_MS,
nowMs: () => Date.now(),
});
Expand Down