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
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);
|
|
}
|
|
|