UNPKG

@cherrystudio/ai-core

Version:

Cherry Studio AI Core - Unified AI Provider Interface Based on Vercel AI SDK

753 lines (747 loc) 25.1 kB
import { customProvider } from "ai"; import { createAnthropic } from "@ai-sdk/anthropic"; import { createAzure } from "@ai-sdk/azure"; import { createDeepSeek } from "@ai-sdk/deepseek"; import { createGoogleGenerativeAI } from "@ai-sdk/google"; import { createOpenAI } from "@ai-sdk/openai"; import { createOpenAICompatible } from "@ai-sdk/openai-compatible"; import { createXai } from "@ai-sdk/xai"; import { createCherryIn } from "@cherrystudio/ai-sdk-provider"; import { createOpenRouter } from "@openrouter/ai-sdk-provider"; import { LRUCache } from "lru-cache"; //#region src/core/providers/core/utils.ts /** * Provider 工具函数和错误类 * 合并自 utils.ts 和 errors.ts */ /** * 格式化私钥,确保它包含正确的PEM头部和尾部 */ function formatPrivateKey(privateKey) { if (!privateKey || typeof privateKey !== "string") throw new Error("Private key must be a non-empty string"); const key = privateKey.replace(/\\n/g, "\n"); const hasBeginMarker = key.includes("-----BEGIN PRIVATE KEY-----"); const hasEndMarker = key.includes("-----END PRIVATE KEY-----"); if (hasBeginMarker && hasEndMarker) return normalizePemFormat(key); return reconstructPemKey(key); } /** * 标准化 PEM 格式 */ function normalizePemFormat(pemKey) { const lines = pemKey.split("\n").map((line) => line.trim()).filter((line) => line.length > 0); let keyContent = ""; let foundBegin = false; let foundEnd = false; for (const line of lines) { if (line === "-----BEGIN PRIVATE KEY-----") { foundBegin = true; continue; } if (line === "-----END PRIVATE KEY-----") { foundEnd = true; break; } if (foundBegin && !foundEnd) keyContent += line; } if (!foundBegin || !foundEnd || !keyContent) throw new Error("Invalid PEM format: missing BEGIN/END markers or key content"); return `-----BEGIN PRIVATE KEY-----\n${keyContent.match(/.{1,64}/g)?.join("\n") || keyContent}\n-----END PRIVATE KEY-----`; } /** * 重新构建 PEM 私钥 */ function reconstructPemKey(key) { let cleanKey = key.replace(/\s+/g, ""); cleanKey = cleanKey.replace(/-----BEGIN[^-]*-----/g, ""); cleanKey = cleanKey.replace(/-----END[^-]*-----/g, ""); if (!cleanKey) throw new Error("Private key content is empty after cleaning"); if (!/^[A-Za-z0-9+/=]+$/.test(cleanKey)) throw new Error("Private key contains invalid characters (not valid Base64)"); return `-----BEGIN PRIVATE KEY-----\n${cleanKey.match(/.{1,64}/g)?.join("\n") || cleanKey}\n-----END PRIVATE KEY-----`; } /** * Provider 创建错误 * 当创建 provider 实例失败时抛出 */ var ProviderCreationError = class extends Error { constructor(message, providerId, cause) { super(message); this.providerId = providerId; this.cause = cause; this.name = "ProviderCreationError"; } }; //#endregion //#region src/core/providers/core/ExtensionRegistry.ts /** * Provider Extension 注册表 * * 职责: * - 注册和管理 Provider Extensions * - 根据 ID 查找对应的 Extension * - 创建并注册 provider 实例(包括变体) * * @example * ```typescript * import { extensionRegistry } from '@cherrystudio/ai-core/provider' * import { OpenAIExtension } from './extensions/openai' * * // 注册 extension * extensionRegistry.register(OpenAIExtension) * * // 批量注册 * extensionRegistry.registerAll([ * OpenAIExtension, * AzureExtension, * AnthropicExtension * ]) * * // 创建并注册 provider 实例 * await extensionRegistry.createAndRegisterProvider('openai', { * apiKey: 'sk-xxx' * }) * ``` */ var ExtensionRegistry = class { constructor() { this.extensions = /* @__PURE__ */ new Map(); this.aliasMap = /* @__PURE__ */ new Map(); } /** * 注册单个 Extension * 支持链式调用 */ register(extension) { const { name, aliases, variants } = extension.config; if (this.extensions.has(name)) return this; this.extensions.set(name, extension); if (aliases) for (const alias of aliases) { if (this.aliasMap.has(alias)) throw new Error(`Provider alias "${alias}" is already registered for "${this.aliasMap.get(alias)}"`); this.aliasMap.set(alias, name); } if (variants) for (const variant of variants) { const variantId = `${name}-${variant.suffix}`; if (this.aliasMap.has(variantId)) throw new Error(`Provider variant ID "${variantId}" is already registered for "${this.aliasMap.get(variantId)}"`); this.aliasMap.set(variantId, name); } return this; } /** * 批量注册 Extensions * 支持 readonly 数组(用于 as const 数组) */ registerAll(extensions) { for (const ext of extensions) this.register(ext); return this; } /** * 取消注册 Extension */ unregister(name) { const extension = this.extensions.get(name); if (!extension) return false; extension.clearCache(); this.extensions.delete(name); if (extension.config.aliases) for (const alias of extension.config.aliases) this.aliasMap.delete(alias); if (extension.config.variants) for (const variant of extension.config.variants) this.aliasMap.delete(`${name}-${variant.suffix}`); return true; } /** * 获取 Extension(支持别名) */ get(id) { if (this.extensions.has(id)) return this.extensions.get(id); const realName = this.aliasMap.get(id); if (realName) return this.extensions.get(realName); } /** * 获取 Extension * * @param id - Provider ID(必须是 RegisteredProviderId) * @returns Extension 或 undefined * * @example * ```typescript * const ext = extensionRegistry.getTyped('openai') * if (ext) { * const provider = await ext.createProvider({ * apiKey: 'sk-...' * }) * } * ``` */ getTyped(id) { return this.get(id); } /** * 检查 Extension 是否已注册 */ has(id) { return this.extensions.has(id) || this.aliasMap.has(id); } /** * 获取所有已注册的 Extension */ getAll() { return Array.from(this.extensions.values()); } /** * 获取所有已注册的 provider IDs(包含变体) * 返回类型安全的 RegisteredProviderId 数组,自动去重 */ getAllProviderIds() { const ids = /* @__PURE__ */ new Set(); for (const extension of this.extensions.values()) for (const id of extension.getProviderIds()) ids.add(id); return Array.from(ids); } /** * 根据 base ID + mode 解析到完整的 provider ID * * 支持别名:如果 baseId 是别名,会先解析到规范 ID * * @param baseId - 基础 provider ID(可以是别名) * @param mode - 模式(如 'chat', 'responses') * @returns 完整的 provider ID,如果无法解析则返回 null * * @example * ```typescript * resolveProviderIdWithMode('openai', 'chat') // → 'openai-chat' * resolveProviderIdWithMode('azure', 'responses') // → 'azure-responses' * resolveProviderIdWithMode('gemini', 'chat') // → null (google 没有 chat 变体) * resolveProviderIdWithMode('openai') // → 'openai' (没有 mode) * ``` */ resolveProviderIdWithMode(baseId, mode) { if (!mode) { const extension = this.get(baseId); return extension ? extension.config.name : null; } const extension = this.get(baseId); if (!extension) return null; if (!extension.config.variants) return null; const variant = extension.config.variants.find((v) => v.suffix === mode); if (!variant) return null; return `${extension.config.name}-${variant.suffix}`; } /** * 反向解析:从完整 ID 提取 base ID 和 mode * * 遍历所有 extensions 的变体,匹配 `${name}-${suffix}` 模式 * * @param providerId - 完整的 provider ID * @returns 解析结果,如果无法解析返回 null * * @example * ```typescript * parseProviderId('openai-chat') // → { baseId: 'openai', mode: 'chat', isVariant: true } * parseProviderId('azure-responses') // → { baseId: 'azure', mode: 'responses', isVariant: true } * parseProviderId('openai') // → { baseId: 'openai', isVariant: false } * parseProviderId('oai') // → { baseId: 'openai', isVariant: false } (别名) * parseProviderId('unknown') // → null * ``` */ parseProviderId(providerId) { for (const ext of this.extensions.values()) { if (!ext.config.variants) continue; for (const variant of ext.config.variants) if (`${ext.config.name}-${variant.suffix}` === providerId) return { baseId: ext.config.name, mode: variant.suffix, isVariant: true }; } const extension = this.get(providerId); if (extension) return { baseId: extension.config.name, isVariant: false }; return null; } /** * 检查是否为变体 ID * * @param id - Provider ID * @returns 如果是变体 ID 返回 true * * @example * ```typescript * isVariant('openai-chat') // → true * isVariant('azure-responses') // → true * isVariant('openai') // → false * isVariant('unknown') // → false * ``` */ isVariant(id) { return this.parseProviderId(id)?.isVariant ?? false; } /** * 获取基础 provider ID * * 对于变体ID,返回其基础provider ID; * 对于基础ID或别名,返回规范的provider ID; * 对于未知ID,返回null * * @param id - Provider ID(可以是基础ID、变体ID或别名) * @returns 基础 provider ID,如果无法解析则返回 null * * @example * ```typescript * getBaseProviderId('openai-chat') // → 'openai' (变体) * getBaseProviderId('azure-responses') // → 'azure' (变体) * getBaseProviderId('openai') // → 'openai' (基础ID) * getBaseProviderId('oai') // → 'openai' (别名) * getBaseProviderId('unknown') // → null * ``` */ getBaseProviderId(id) { return this.parseProviderId(id)?.baseId ?? null; } /** * 获取变体的模式/后缀 * * @param variantId - 变体 ID * @returns 模式/后缀,如果不是变体则返回 null * * @example * ```typescript * getVariantMode('openai-chat') // → 'chat' * getVariantMode('azure-responses') // → 'responses' * getVariantMode('openai') // → null (不是变体) * getVariantMode('unknown') // → null * ``` */ getVariantMode(variantId) { return this.parseProviderId(variantId)?.mode ?? null; } /** 获取 variant 的 resolveModel 函数(类型安全在 extension 声明处保证) */ getModelResolver(providerId) { const parsed = this.parseProviderId(providerId); if (!parsed) return void 0; const extension = this.get(parsed.baseId); if (!extension) return void 0; if (parsed.isVariant && parsed.mode) { const variant = extension.getVariant(parsed.mode); if (variant?.resolveModel) return variant.resolveModel; } } /** * 获取某个基础 provider 的所有变体 IDs * * @param baseId - 基础 provider ID(可以是别名) * @returns 变体 ID 数组,如果没有变体则返回空数组 * * @example * ```typescript * getVariants('openai') // → ['openai-chat'] * getVariants('azure') // → ['azure-responses'] * getVariants('google') // → ['google-chat'] * getVariants('xai') // → [] (没有变体) * getVariants('unknown') // → [] (未注册) * ``` */ getVariants(baseId) { const extension = this.get(baseId); if (!extension?.config.variants) return []; return extension.config.variants.map((v) => `${extension.config.name}-${v.suffix}`); } /** 获取指定 provider 的工具工厂(变体优先,回退到 base) */ getToolFactory(providerId, capability) { const parsed = this.parseProviderId(providerId); if (!parsed) return void 0; const { baseId, mode, isVariant } = parsed; const extension = this.get(baseId); if (!extension) return void 0; if (isVariant && mode) { const variant = extension.getVariant(mode); if (variant?.toolFactories?.[capability]) return variant.toolFactories[capability]; } return extension.config.toolFactories?.[capability]; } /** * 解析工具能力:返回 factory + provider 实例 * * 1. Direct — provider 自己有 toolFactories * 2. Aggregator fallback — 从 model.provider 段解析(如 "aihubmix.google" → google extension) */ async resolveToolCapability(providerId, capability, modelProvider) { const directFactory = this.getToolFactory(providerId, capability); if (directFactory) { const provider = await this.getToolProvider(providerId); if (provider) return { factory: directFactory, provider }; } if (typeof modelProvider === "string") { const segments = modelProvider.split("."); for (let i = segments.length - 1; i >= 0; i--) { const factory = this.getToolFactory(segments[i], capability); if (factory) { const provider = await this.getToolProvider(segments[i]); if (provider) return { factory, provider }; } } } } /** Get provider for .tools extraction (cached or dummy instance) */ async getToolProvider(providerId) { const parsed = this.parseProviderId(providerId); if (!parsed) return void 0; const extension = this.get(parsed.baseId); if (!extension) return void 0; try { return await extension.createProvider(extension.getCachedProvider() ? void 0 : { apiKey: "_tool_descriptor" }, parsed.isVariant ? parsed.mode : void 0); } catch { return; } } /** * 清空所有注册 */ clear() { this.extensions.clear(); this.aliasMap.clear(); } async createProvider(id, settings) { const parsed = this.parseProviderId(id); if (!parsed) throw new Error(`Provider extension "${id}" not found. Did you forget to register it?`); const { baseId, mode: variantSuffix } = parsed; const extension = this.get(baseId); if (!extension) throw new Error(`Provider extension "${baseId}" not found. Did you forget to register it?`); try { return await extension.createProvider(settings, variantSuffix); } catch (error) { throw new ProviderCreationError(`Failed to create provider "${id}"`, id, error instanceof Error ? error : new Error(String(error))); } } }; /** * 全局 Extension Registry 实例 * 单例模式,确保整个应用只有一个注册表 */ const extensionRegistry = new ExtensionRegistry(); //#endregion //#region src/core/utils/index.ts const isPlainObject = (value) => { return typeof value === "object" && value !== null && !Array.isArray(value); }; function deepMergeObjects(target, source) { const result = { ...target }; Object.entries(source).forEach(([key, value]) => { if (isPlainObject(value) && isPlainObject(result[key])) result[key] = deepMergeObjects(result[key], value); else result[key] = value; }); return result; } //#endregion //#region src/core/providers/core/ProviderExtension.ts /** * Provider Extension 类 * * @typeParam TSettings - Provider 配置类型 * @typeParam TProvider - 实际 provider 类型(用于 variants) * @typeParam TConfig - 配置对象类型(幻影类型参数,用于自动推导 Provider IDs) */ var ProviderExtension = class ProviderExtension { constructor(config) { this.config = config; this.pendingCreations = /* @__PURE__ */ new Map(); if (!config.name) throw new Error("ProviderExtension: name is required"); this.instances = new LRUCache({ max: 10, updateAgeOnGet: true }); } static create(config) { return new ProviderExtension(typeof config === "function" ? config() : config); } /** * Options getter - 只读配置 */ get options() { return Object.freeze({ ...this.config.defaultOptions }); } /** * 计算 settings 的稳定 hash */ computeHash(settings, variantSuffix) { const baseKey = (() => { if (settings === void 0 || settings === null) return "default"; const stableStringify = (obj) => { if (obj === null || obj === void 0) return "null"; if (typeof obj === "function") return "\"[function]\""; if (typeof obj !== "object") return JSON.stringify(obj); if (Array.isArray(obj)) return `[${obj.map(stableStringify).join(",")}]`; return `{${Object.keys(obj).sort().map((key) => `${JSON.stringify(key)}:${stableStringify(obj[key])}`).join(",")}}`; }; return stableStringify(settings); })(); return variantSuffix ? `${baseKey}:${variantSuffix}` : baseKey; } /** * 创建 Provider 实例 * 相同 settings 会复用实例,不同 settings 会创建新实例 */ async createProvider(settings, variantSuffix) { if (variantSuffix) { if (!this.getVariant(variantSuffix)) throw new Error(`ProviderExtension "${this.config.name}": variant "${variantSuffix}" not found. Available variants: ${this.config.variants?.map((v) => v.suffix).join(", ") || "none"}`); } const mergedSettings = deepMergeObjects(this.config.defaultOptions || {}, settings || {}); const hash = this.computeHash(mergedSettings, variantSuffix); const cachedInstance = this.instances.get(hash); if (cachedInstance) return cachedInstance; const pending = this.pendingCreations.get(hash); if (pending) return pending; const creationPromise = this._doCreateProvider(mergedSettings, variantSuffix, hash); this.pendingCreations.set(hash, creationPromise); try { return await creationPromise; } finally { this.pendingCreations.delete(hash); } } /** * 获取基础 provider 实例(无变体转换) * 用于访问 provider 的 .tools 属性 */ async getBaseProvider(settings) { return this.createProvider(settings); } async _doCreateProvider(mergedSettings, variantSuffix, hash) { let baseProvider; if (this.config.create) baseProvider = await Promise.resolve(this.config.create(mergedSettings)); else if (this.config.import && this.config.creatorFunctionName) { const creatorFn = (await this.config.import())[this.config.creatorFunctionName]; if (!creatorFn || typeof creatorFn !== "function") throw new Error(`ProviderExtension "${this.config.name}": creatorFunctionName "${this.config.creatorFunctionName}" not found in imported module`); baseProvider = await Promise.resolve(creatorFn(mergedSettings)); } else throw new Error(`ProviderExtension "${this.config.name}": cannot create provider, invalid configuration`); let finalProvider; if (variantSuffix) { const variant = this.getVariant(variantSuffix); if (variant.transform) { const baseHash = this.computeHash(mergedSettings); if (!this.instances.has(baseHash)) this.instances.set(baseHash, baseProvider); finalProvider = await Promise.resolve(variant.transform(baseProvider, mergedSettings)); } else finalProvider = baseProvider; } else finalProvider = baseProvider; this.instances.set(hash, finalProvider); return finalProvider; } /** * 配置 provider(链式调用) * 返回一个新的 Extension 实例,不修改原实例 */ configure(settings) { return new ProviderExtension({ ...this.config, defaultOptions: deepMergeObjects(this.config.defaultOptions || {}, settings) }); } /** * 获取所有 provider IDs(包含变体和别名) */ getProviderIds() { const ids = [this.config.name, ...this.config.aliases || []]; if (this.config.variants) for (const variant of this.config.variants) ids.push(`${this.config.name}-${variant.suffix}`); return ids; } /** * 检查给定 ID 是否属于此 Extension */ hasProviderId(id) { return this.getProviderIds().includes(id); } /** * 获取变体配置 */ getVariant(suffix) { return this.config.variants?.find((v) => v.suffix === suffix); } /** * 清除所有缓存的 Provider 实例 */ clearCache() { this.instances.clear(); this.pendingCreations.clear(); } /** * 获取已缓存的 provider 实例(如果存在) */ getCachedProvider() { for (const [key, value] of this.instances.entries()) if (!key.includes(":")) return value; for (const [, value] of this.instances.entries()) return value; } /** * 获取缓存统计信息 */ getCacheStats() { return { cachedInstances: this.instances.size }; } }; //#endregion //#region src/core/providers/core/initialization.ts const AnthropicExtension = ProviderExtension.create({ name: "anthropic", aliases: ["claude"], supportsImageGeneration: false, create: createAnthropic, toolFactories: { webSearch: (provider) => (config) => ({ tools: { webSearch: provider.tools.webSearch_20260209(config) } }), urlContext: (provider) => (config) => ({ tools: { urlContext: provider.tools.webFetch_20260209(config) } }) } }); /** * Azure Extension */ const AzureExtension = ProviderExtension.create({ name: "azure", aliases: ["azure-openai"], supportsImageGeneration: true, create: (settings) => { const provider = createAzure(settings); return customProvider({ fallbackProvider: { ...provider, languageModel: (modelId) => provider.chat(modelId) } }); }, toolFactories: { webSearch: (provider) => (config) => ({ tools: { webSearch: provider.tools.webSearchPreview(config) } }) }, variants: [{ suffix: "responses", name: "Azure OpenAI Responses", transform: (_provider, settings) => createAzure(settings), toolFactories: { webSearch: (provider) => (config) => ({ tools: { webSearch: provider.tools.webSearchPreview(config) } }) } }, { suffix: "anthropic", name: "Azure Anthropic", transform: (_provider, settings) => createAnthropic({ baseURL: (settings?.baseURL ?? "") + "/anthropic/v1", apiKey: settings?.apiKey ?? "", headers: settings?.headers }), toolFactories: { webSearch: (provider) => (config) => ({ tools: { webSearch: provider.tools.webSearch_20260209(config) } }), urlContext: (provider) => (config) => ({ tools: { urlContext: provider.tools.webFetch_20260209(config) } }) } }] }); const CherryInExtension = ProviderExtension.create({ name: "cherryin", supportsImageGeneration: true, create: createCherryIn, variants: [{ suffix: "chat", name: "CherryIN Chat", transform: (provider) => customProvider({ fallbackProvider: { ...provider, languageModel: (modelId) => provider.chat(modelId) } }) }] }); const DeepSeekExtension = ProviderExtension.create({ name: "deepseek", supportsImageGeneration: false, create: createDeepSeek }); const GoogleExtension = ProviderExtension.create({ name: "google", aliases: [ "google-ai", "gemini", "google-gemini" ], supportsImageGeneration: true, create: createGoogleGenerativeAI, toolFactories: { webSearch: (provider) => (config) => ({ tools: { webSearch: provider.tools.googleSearch(config) } }), urlContext: (provider) => (config) => ({ tools: { urlContext: provider.tools.urlContext(config) } }) } }); const OpenAICompatibleExtension = ProviderExtension.create({ name: "openai-compatible", supportsImageGeneration: true, create: (settings) => { if (!settings) throw new Error("OpenAI Compatible provider requires settings"); return createOpenAICompatible(settings); } }); const OpenAIExtension = ProviderExtension.create({ name: "openai", aliases: ["openai-response"], supportsImageGeneration: true, create: createOpenAI, toolFactories: { webSearch: (provider) => (config) => ({ tools: { webSearch: provider.tools.webSearch(config) } }) }, variants: [{ suffix: "chat", name: "OpenAI Chat", resolveModel: (provider, modelId) => provider.chat(modelId), toolFactories: { webSearch: (provider) => (config) => ({ tools: { webSearch: provider.tools.webSearchPreview(config) } }) } }] }); const OpenRouterExtension = ProviderExtension.create({ name: "openrouter", aliases: ["tokenflux"], supportsImageGeneration: true, create: createOpenRouter, toolFactories: { webSearch: () => (config) => ({ providerOptions: { openrouter: config } }) } }); const XaiExtension = ProviderExtension.create({ name: "xai", aliases: ["grok"], supportsImageGeneration: true, create: createXai, variants: [{ suffix: "responses", name: "xAI Responses", resolveModel: (provider, modelId) => provider.responses(modelId), toolFactories: { webSearch: (provider) => (config) => ({ tools: { webSearch: provider.tools.webSearch(config?.webSearch ?? {}), xSearch: provider.tools.xSearch(config?.xSearch ?? {}) } }) } }] }); /** * 核心 provider extensions 列表 */ const coreExtensions = [ OpenAIExtension, AnthropicExtension, AzureExtension, GoogleExtension, XaiExtension, DeepSeekExtension, OpenRouterExtension, OpenAICompatibleExtension, CherryInExtension ]; const registeredProviderIds = (() => { const map = {}; coreExtensions.forEach((ext) => { const config = ext.config; const name = config.name; map[name] = name; if (config.aliases) config.aliases.forEach((alias) => { map[alias] = name; }); if (config.variants) config.variants.forEach((variant) => { map[`${name}-${variant.suffix}`] = name; }); }); return map; })(); /** * 注册所有通用 extensions 到全局 registry * 在模块加载时自动执行 * * 注意:只注册通用的 provider extensions(OpenAI, Anthropic, Google 等) * 项目特定的 extensions 应该在应用层单独注册 */ extensionRegistry.registerAll(coreExtensions); /** * 检查是否有对应的 Provider Extension */ function hasProviderConfig(providerId) { return extensionRegistry.has(providerId); } //#endregion export { extensionRegistry as a, ExtensionRegistry as i, hasProviderConfig as n, ProviderCreationError as o, ProviderExtension as r, formatPrivateKey as s, coreExtensions as t };