You can not select more than 25 topics Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
 
 
 
 

92 lines
2.7 KiB

import { createOpenAI } from "@ai-sdk/openai";
import { createOpenAICompatible } from "@ai-sdk/openai-compatible";
import { type LanguageModel } from "ai";
import { getModelWithProviderById, getSystemModelWithProviderById, getModelWithProviderByIdAny } from "#server/service/llm";
import log4js from "logger";
const logger = log4js.getLogger("APP");
export interface ResolvedProvider {
name: string;
apiKey: string | null;
baseUrl: string | null;
parseMode: string;
}
export interface ResolvedModel {
dbId: number;
modelId: string;
provider: ResolvedProvider;
}
export function resolveLanguageModel(provider: ResolvedProvider, modelId: string): LanguageModel {
const baseUrl = provider.baseUrl?.replace(/\/+$/, "") || undefined;
if (provider.parseMode === "openai") {
const openai = createOpenAI({
apiKey: provider.apiKey || undefined,
baseURL: baseUrl,
});
return openai(modelId) as unknown as LanguageModel;
}
const openaiCompatible = createOpenAICompatible({
name: provider.name,
apiKey: provider.apiKey || undefined,
baseURL: baseUrl || "https://api.openai.com/v1",
});
return openaiCompatible(modelId) as unknown as LanguageModel;
}
export async function resolveModelForUser(
modelId: number,
userId: number | null,
): Promise<ResolvedModel | null> {
const modelRow = userId
? await getModelWithProviderById(modelId, userId)
: await getSystemModelWithProviderById(modelId);
if (!modelRow) {
logger.warn("[MODEL-RESOLVER] modelId=%d not found for userId=%s", modelId, userId);
return null;
}
const { model, provider } = modelRow;
if (!provider.apiKey) {
logger.warn("[MODEL-RESOLVER] provider has no apiKey, skip");
return null;
}
return {
dbId: model.id,
modelId: model.modelId,
provider: {
name: provider.name,
apiKey: provider.apiKey,
baseUrl: provider.baseUrl,
parseMode: provider.parseMode,
},
};
}
export async function resolveModelAny(modelId: number): Promise<ResolvedModel | null> {
const modelRow = await getModelWithProviderByIdAny(modelId);
if (!modelRow) {
logger.warn("[MODEL-RESOLVER] modelId=%d not found (any)", modelId);
return null;
}
const { model, provider } = modelRow;
if (!provider.apiKey) {
logger.warn("[MODEL-RESOLVER] provider has no apiKey, skip");
return null;
}
return {
dbId: model.id,
modelId: model.modelId,
provider: {
name: provider.name,
apiKey: provider.apiKey,
baseUrl: provider.baseUrl,
parseMode: provider.parseMode,
},
};
}
export function toLanguageModel(resolved: ResolvedModel): LanguageModel {
return resolveLanguageModel(resolved.provider, resolved.modelId);
}