Files
Sybil-2/server/src/llm/provider-adapters.ts
T

285 lines
9.7 KiB
TypeScript
Raw Normal View History

2026-06-13 12:02:22 -07:00
import {
normalizeEnabledChatTools,
type ToolAwareCompletionParams,
type ToolAwareCompletionResult,
type ToolAwareStreamingEvent,
} from "./chat-tools.js";
import { completeWithChatCompletionsApi, streamWithChatCompletionsApi } from "./protocols/chat-completions-api.js";
2026-07-11 14:16:21 -07:00
import { completeWithGeminiApi, streamWithGeminiApi } from "./protocols/gemini-api.js";
2026-06-13 12:02:22 -07:00
import { completeWithMessagesApi, streamWithMessagesApi } from "./protocols/messages-api.js";
import { completeWithResponsesApi, streamWithResponsesApi } from "./protocols/responses-api.js";
import { env } from "../env.js";
2026-07-11 14:16:21 -07:00
import { anthropicClient, geminiClient, hermesAgentClient, isHermesAgentConfigured, openaiClient, xaiClient } from "./providers.js";
2026-06-13 12:02:22 -07:00
import type { ChatMessage, Provider } from "./types.js";
type ProviderAdapterParams = {
model: string;
messages: ChatMessage[];
enabledTools?: string[];
userLocation?: string;
temperature?: number;
maxTokens?: number;
logContext?: ToolAwareCompletionParams["logContext"];
};
export type ProviderChatAdapter = {
provider: Provider;
complete(params: ProviderAdapterParams): Promise<ToolAwareCompletionResult>;
stream(params: ProviderAdapterParams): AsyncGenerator<ToolAwareStreamingEvent>;
};
2026-07-11 14:16:21 -07:00
type ChatProtocolId = "chat-completions" | "gemini" | "messages" | "responses";
2026-06-13 12:02:22 -07:00
type ChatProtocol = {
id: ChatProtocolId;
complete(params: ToolAwareCompletionParams): Promise<ToolAwareCompletionResult>;
stream(params: ToolAwareCompletionParams): AsyncGenerator<ToolAwareStreamingEvent>;
};
type ModelCatalogSpec = {
enabled?: () => boolean;
fetchModels(client: any): Promise<string[]>;
fallbackModels?: () => string[];
2026-07-11 14:16:21 -07:00
sortModels?: (models: string[]) => string[];
2026-06-13 12:02:22 -07:00
};
type ProviderBackendSpec = {
createClient: () => any;
plainProtocol: ChatProtocol;
toolProtocol?: ChatProtocol;
managedTools?: boolean;
modelCatalog?: ModelCatalogSpec;
};
const chatCompletionsProtocol: ChatProtocol = {
id: "chat-completions",
complete: completeWithChatCompletionsApi,
stream: streamWithChatCompletionsApi,
};
const messagesProtocol: ChatProtocol = {
id: "messages",
complete: completeWithMessagesApi,
stream: streamWithMessagesApi,
};
2026-07-11 14:16:21 -07:00
const geminiProtocol: ChatProtocol = {
id: "gemini",
complete: completeWithGeminiApi,
stream: streamWithGeminiApi,
};
2026-06-13 12:02:22 -07:00
const responsesProtocol: ChatProtocol = {
id: "responses",
complete: completeWithResponsesApi,
stream: streamWithResponsesApi,
};
function uniqSorted(values: string[]) {
return [...new Set(values.map((value) => value.trim()).filter(Boolean))].sort((a, b) => a.localeCompare(b));
}
function modelIdsFromListResponse(page: any) {
return Array.isArray(page?.data)
? page.data.map((model: any) => model?.id).filter((id: unknown): id is string => typeof id === "string")
: [];
}
2026-07-11 14:16:21 -07:00
function stripModelResourcePrefix(model: string) {
return model.startsWith("models/") ? model.slice("models/".length) : model;
}
2026-06-13 12:02:22 -07:00
function isLikelyResponsesApiModel(model: string) {
const id = model.toLowerCase();
if (id.includes("embedding") || id.includes("moderation")) return false;
if (id.includes("audio") || id.includes("realtime") || id.includes("transcribe") || id.includes("tts")) return false;
if (id.includes("image") || id.includes("dall-e") || id.includes("sora")) return false;
if (id.includes("search") || id.includes("computer-use")) return false;
return id === "chat-latest" || /^(gpt-|o\d|chatgpt-)/.test(id);
2026-06-13 12:02:22 -07:00
}
2026-07-11 14:16:21 -07:00
function isLikelyGeminiChatModel(model: string) {
const id = model.toLowerCase();
if (!id.startsWith("gemini-")) return false;
if (id.includes("embedding") || id.includes("embed")) return false;
if (id.includes("image") || id.includes("imagen") || id.includes("veo")) return false;
if (id.includes("audio") || id.includes("tts") || id.includes("live")) return false;
if (id.includes("computer-use") || id.includes("robotics")) return false;
return true;
}
function preferGeminiModels(models: string[]) {
const preferred = [
"gemini-3.5-flash",
"gemini-flash-latest",
"gemini-3.1-flash-lite",
"gemini-3-flash-preview",
"gemini-pro-latest",
];
const modelSet = new Set(models);
return [...preferred.filter((model) => modelSet.delete(model)), ...[...modelSet].sort((a, b) => a.localeCompare(b))];
}
async function fetchJson(url: URL): Promise<any> {
const response = await fetch(url);
const body: any = await response.json().catch(() => null);
if (!response.ok) {
throw new Error(body?.error?.message ?? `Gemini model fetch failed with status ${response.status}.`);
}
return body;
}
2026-06-13 12:02:22 -07:00
function withClient(params: ProviderAdapterParams, client: any, enabledTools?: string[]): ToolAwareCompletionParams {
return {
client,
model: params.model,
messages: params.messages,
enabledTools,
userLocation: params.userLocation,
temperature: params.temperature,
maxTokens: params.maxTokens,
logContext: params.logContext,
};
}
function selectChatProtocol(spec: ProviderBackendSpec, params: Pick<ProviderAdapterParams, "enabledTools">) {
const enabledTools = normalizeEnabledChatTools(params.enabledTools);
const useManagedTools = spec.managedTools === true && spec.toolProtocol && enabledTools.length > 0;
return {
protocol: useManagedTools ? spec.toolProtocol! : spec.plainProtocol,
enabledTools: useManagedTools ? enabledTools : [],
managedTools: Boolean(useManagedTools),
};
}
function createProviderChatAdapter(provider: Provider, spec: ProviderBackendSpec): ProviderChatAdapter {
return {
provider,
complete(params) {
const selected = selectChatProtocol(spec, params);
return selected.protocol.complete(withClient(params, spec.createClient(), selected.enabledTools));
},
stream(params) {
const selected = selectChatProtocol(spec, params);
return selected.protocol.stream(withClient(params, spec.createClient(), selected.enabledTools));
},
};
}
const backendSpecs: Record<Provider, ProviderBackendSpec> = {
openai: {
createClient: openaiClient,
plainProtocol: chatCompletionsProtocol,
toolProtocol: responsesProtocol,
managedTools: true,
modelCatalog: {
async fetchModels(client) {
const page = await client.models.list();
return modelIdsFromListResponse(page).filter(isLikelyResponsesApiModel);
},
},
},
anthropic: {
createClient: anthropicClient,
plainProtocol: messagesProtocol,
toolProtocol: messagesProtocol,
managedTools: true,
modelCatalog: {
async fetchModels(client) {
const page = await client.models.list({ limit: 200 });
return modelIdsFromListResponse(page);
},
},
},
xai: {
createClient: xaiClient,
plainProtocol: chatCompletionsProtocol,
toolProtocol: chatCompletionsProtocol,
managedTools: true,
modelCatalog: {
async fetchModels(client) {
const page = await client.models.list();
return modelIdsFromListResponse(page);
},
},
},
2026-07-11 14:16:21 -07:00
gemini: {
createClient: geminiClient,
plainProtocol: geminiProtocol,
toolProtocol: geminiProtocol,
managedTools: true,
modelCatalog: {
async fetchModels(client) {
const url = new URL(`${client.baseURL.replace(/\/+$/, "")}/models`);
url.searchParams.set("key", client.apiKey);
url.searchParams.set("pageSize", "1000");
const page = await fetchJson(url);
return Array.isArray(page?.models)
? page.models
.filter((model: any) => Array.isArray(model?.supportedGenerationMethods) && model.supportedGenerationMethods.includes("generateContent"))
.map((model: any) => model?.name)
.filter((id: unknown): id is string => typeof id === "string")
.map(stripModelResourcePrefix)
.filter(isLikelyGeminiChatModel)
: [];
},
sortModels: preferGeminiModels,
},
},
2026-06-13 12:02:22 -07:00
"hermes-agent": {
createClient: hermesAgentClient,
plainProtocol: chatCompletionsProtocol,
managedTools: false,
modelCatalog: {
enabled: isHermesAgentConfigured,
async fetchModels(client) {
const page = await client.models.list();
const models = modelIdsFromListResponse(page);
if (env.HERMES_AGENT_MODEL) models.push(env.HERMES_AGENT_MODEL);
return models;
},
fallbackModels() {
return env.HERMES_AGENT_MODEL ? [env.HERMES_AGENT_MODEL] : [];
},
},
},
};
const providerChatAdapters: Record<Provider, ProviderChatAdapter> = Object.fromEntries(
Object.entries(backendSpecs).map(([provider, spec]) => [provider, createProviderChatAdapter(provider as Provider, spec)])
) as Record<Provider, ProviderChatAdapter>;
export function getProviderChatAdapter(provider: Provider) {
return providerChatAdapters[provider];
}
export function describeProviderChatBackend(provider: Provider, enabledTools?: string[]) {
const selected = selectChatProtocol(backendSpecs[provider], { enabledTools });
return {
provider,
protocol: selected.protocol.id,
managedTools: selected.managedTools,
enabledTools: selected.enabledTools,
};
}
export function listModelCatalogProviders(): Provider[] {
return (Object.entries(backendSpecs) as [Provider, ProviderBackendSpec][])
.filter(([, spec]) => {
const catalog = spec.modelCatalog;
return catalog !== undefined && catalog.enabled?.() !== false;
})
.map(([provider]) => provider);
}
export async function fetchProviderCatalogModels(provider: Provider) {
const spec = backendSpecs[provider].modelCatalog;
if (!spec) return [];
2026-07-11 14:16:21 -07:00
const models = uniqSorted(await spec.fetchModels(backendSpecs[provider].createClient()));
return spec.sortModels ? spec.sortModels(models) : models;
2026-06-13 12:02:22 -07:00
}
export function getProviderCatalogFallbackModels(provider: Provider) {
return uniqSorted(backendSpecs[provider].modelCatalog?.fallbackModels?.() ?? []);
}