Files
veridian/server/utils/ai-provider.ts
T

108 lines
3.9 KiB
TypeScript

import { createCerebras, type CerebrasProvider } from "@ai-sdk/cerebras";
import { createGoogleGenerativeAI, type GoogleGenerativeAIProvider } from "@ai-sdk/google";
import { createOpenRouter, type OpenRouterProvider } from "@openrouter/ai-sdk-provider";
import { createOllama, type OllamaProvider } from "ai-sdk-ollama";
import { createOpenAICompatible, type OpenAICompatibleProvider } from '@ai-sdk/openai-compatible';
import { createCohere, type CohereProvider } from '@ai-sdk/cohere';
import { type Entity } from "@triplit/client";
import { schema } from "~~/triplit/schema";
import { type StreamTextTransform } from "ai";
import { transformCerebrasReasoningStream } from "./cerebras";
import { providerBaseUrls } from "~/types/model";
// import { createLongcatTransformer } from "./longcat";
export type ModelGateway = OpenRouterProvider | OllamaProvider | CerebrasProvider | GoogleGenerativeAIProvider | OpenAICompatibleProvider | CohereProvider;
export interface Gateway {
gateway: ModelGateway;
streamTransformer: StreamTextTransform<{}> | StreamTextTransform<{}>[] | undefined;
textTransformer: ((text: string) => string) | ((text: string) => string)[] | undefined;
}
export async function getGateway(provider: Entity<typeof schema, 'providers'>, model: Entity<typeof schema, 'models'>, providerApiKey?: string): Promise<Gateway> {
let gateway: ModelGateway;
let streamTransformer = undefined;
let textTransformer = undefined;
let baseURL = undefined;
if (provider.config.apiProxyUrl && provider.config.apiProxyUrl.trim() !== '') {
baseURL = provider.config.apiProxyUrl;
}
switch (provider.type) {
case 'openrouter': {
if (providerApiKey === undefined) {
throw createError({
statusCode: 400,
message: 'OpenRouter provider requires an API key',
});
}
gateway = createOpenRouter({
apiKey: providerApiKey,
headers: {
'HTTP-Referer': 'https://localhost:3000',
'X-Title': 'Veridian',
},
});
break;
}
case 'ollama': {
if (baseURL === undefined) {
throw createError({
statusCode: 400,
message: 'Ollama provider requires an API proxy URL',
});
}
const innerGateway = createOllama({
apiKey: providerApiKey,
baseURL,
})
gateway = ((modelId: string) => innerGateway(modelId, { think: [...model.attributes.capabilities].includes('reasoning') })) as OllamaProvider;
break;
}
case 'cerebras': {
gateway = createCerebras({
apiKey: providerApiKey,
baseURL,
})
streamTransformer = transformCerebrasReasoningStream() as StreamTextTransform<{}>;
textTransformer = (text: string) => {
return text.split('</think>').at(-1)!.trim()
};
break;
}
case 'google': {
gateway = createGoogleGenerativeAI({
apiKey: providerApiKey,
baseURL,
})
break;
}
case 'longcat': {
gateway = createOpenAICompatible({
name: 'LongCat',
apiKey: providerApiKey,
baseURL: baseURL ?? providerBaseUrls[provider.type],
includeUsage: true,
})
// streamTransformer = createLongcatTransformer() as StreamTextTransform<{}>;
} break;
case 'cohere': {
gateway = createCohere({
apiKey: providerApiKey,
baseURL,
})
}
}
return {
gateway,
streamTransformer,
textTransformer,
};
}