738 lines
27 KiB
Go
738 lines
27 KiB
Go
package llm
|
||
|
||
import (
|
||
"context"
|
||
"fmt"
|
||
"log"
|
||
"os"
|
||
"strings"
|
||
"sync"
|
||
|
||
"tcm-agent/internal/config"
|
||
"tcm-agent/internal/dao"
|
||
"tcm-agent/internal/types"
|
||
)
|
||
|
||
// ========================================================================
|
||
// LLM 模型工厂 —— 设计思路说明
|
||
// ========================================================================
|
||
// 为什么要工厂模式?
|
||
// AI Agent 平台不可能只绑定一家模型厂商。今天用 DeepSeek,明天可能要
|
||
// 接入 GPT-4o、通义千问、本地 Ollama,甚至要同时用多个模型(比如:
|
||
// 病历生成用 DeepSeek,处方校验用 GPT-4,Embedding 用本地模型)。
|
||
//
|
||
// 工厂模式解决的问题:
|
||
// 1. 统一接口:所有模型实现同一个 LLMClient 接口
|
||
// 2. 按名取用:通过 Provider 名称("deepseek"/"openai")动态创建
|
||
// 3. 配置驱动:新增模型只需改 config.yaml,不改一行业务代码
|
||
// 4. 多模型共存:不同 Agent 可以用不同模型,互不干扰
|
||
// 5. 降级兜底:主模型挂了可以自动切到备用模型
|
||
//
|
||
// 架构层次:
|
||
// ┌─────────────────────────────────────────────────────────┐
|
||
// │ Agent Runner (调度层) │
|
||
// │ 病历Agent → 用 ModelRouter.Get("emr-generator") │
|
||
// │ 处方Agent → 用 ModelRouter.Get("prescription-maker") │
|
||
// ├─────────────────────────────────────────────────────────┤
|
||
// │ ModelRouter (路由层) │
|
||
// │ 按场景名 → 映射到具体 Provider → 返回 LLMClient │
|
||
// ├─────────────────────────────────────────────────────────┤
|
||
// │ LLMFactory (工厂层) │
|
||
// │ "deepseek" → DeepSeekClient │
|
||
// │ "openai" → OpenAIClient │
|
||
// │ "azure" → AzureOpenAIClient │
|
||
// │ "ollama" → OllamaClient (本地模型) │
|
||
// │ "qwen" → QwenClient (通义千问) │
|
||
// │ "spark" → SparkClient (讯飞星火 OpenAPI) │
|
||
// ├─────────────────────────────────────────────────────────┤
|
||
// │ LLMClient Interface (统一接口) │
|
||
// │ Chat() / Embed() / StreamChat() / Name() / Supports()│
|
||
// └─────────────────────────────────────────────────────────┘
|
||
// ========================================================================
|
||
|
||
// ========================================================================
|
||
// 统一接口定义
|
||
// ========================================================================
|
||
|
||
// LLMClient 大语言模型客户端统一接口
|
||
//
|
||
// 所有模型供应商(DeepSeek/OpenAI/Azure/Ollama/通义千问)都必须实现此接口。
|
||
// Agent Runner 只依赖这个接口,不关心背后是哪家的模型。
|
||
//
|
||
// 设计原则:
|
||
// - 输入统一:messages + tools,所有模型都一样
|
||
// - 输出统一:返回 *types.Message,Agent 不感知模型差异
|
||
// - 能力声明:Supports() 让调用方知道这个模型支持哪些特性
|
||
type LLMClient interface {
|
||
// Chat 发起一次对话请求(非流式)
|
||
// ctx - 上下文,支持超时取消
|
||
// messages - 对话历史(包含 system/user/assistant/tool 角色)
|
||
// tools - 可供模型调用的工具列表(Function Calling)
|
||
// 返回 - 模型生成的消息(可能包含工具调用指令)
|
||
Chat(ctx context.Context, messages []types.Message, tools []types.Tool) (*types.Message, error)
|
||
|
||
// StreamChat 发起流式对话(可选实现,用于实时输出)
|
||
// 返回 chan 逐块输出文本,调用方按需消费
|
||
StreamChat(ctx context.Context, messages []types.Message, tools []types.Tool) (<-chan string, error)
|
||
|
||
// Embed 生成文本向量(用于知识库检索、语义记忆)
|
||
// 输入文本列表,返回对应的向量数组
|
||
Embed(ctx context.Context, texts []string) ([][]float32, error)
|
||
|
||
// Name 返回模型标识(如 "deepseek-chat" / "gpt-4o")
|
||
Name() string
|
||
|
||
// Provider 返回供应商名称(如 "deepseek" / "openai")
|
||
Provider() string
|
||
|
||
// Supports 查询该模型是否支持某项能力
|
||
// 常用能力: "function_calling" / "vision" / "streaming" / "json_mode"
|
||
Supports(capability string) bool
|
||
|
||
// Close 释放资源(关闭连接池等)
|
||
Close() error
|
||
}
|
||
|
||
// ========================================================================
|
||
// 能力常量定义
|
||
// ========================================================================
|
||
|
||
const (
|
||
CapFunctionCalling = "function_calling" // 函数调用(工具使用)
|
||
CapVision = "vision" // 多模态视觉理解
|
||
CapStreaming = "streaming" // 流式输出
|
||
CapJSONMode = "json_mode" // JSON 模式输出
|
||
CapEmbedding = "embedding" // 文本向量化
|
||
CapLongContext = "long_context" // 长上下文(>32K)
|
||
)
|
||
|
||
// ------------------------------------------------------------------
|
||
// 可选能力接口(用于类型断言取 token/finish_reason)
|
||
// ------------------------------------------------------------------
|
||
|
||
// TokenAwareClient 可选接口:实现此接口的 provider 可以返回 token 用量
|
||
//
|
||
// 为什么不直接改 LLMClient.Chat 返回类型:
|
||
// Chat 返回 *types.Message 是为了"对话内容"语义清晰,且接口已被所有 provider
|
||
// 实现;强行加 token 字段会污染所有调用方。改用可选接口让需要 token 统计的
|
||
// 调用方(EnhancerService、ReactLoop)通过类型断言获取,老 provider 不实现
|
||
// 时自动回落到 0,零破坏。
|
||
//
|
||
// 实现方:DeepSeekClient / SparkClient / OpenAIClient
|
||
type TokenAwareClient interface {
|
||
// LastChatResult 返回最近一次 Chat 调用的 token 用量与 finish_reason
|
||
//
|
||
// 注意:本方法是"线程不安全"的——它返回的是客户端实例最近一次调用的快照。
|
||
// 业务上 LLMClient 通常是单例,多协程并发调用同一 client 时本方法返回值
|
||
// 不可靠。EnhancerService 调用模式是"一次请求串行一次 Chat",不存在并发,
|
||
// 所以可用。如果将来加并发,需要换成 Chat 直接返回 *ChatResult。
|
||
LastChatResult() *types.ChatResult
|
||
}
|
||
|
||
// ChatOpts 调用 LLM 时的可配置参数(覆盖客户端默认值)
|
||
//
|
||
// 用途:Agent 高级能力(ReAct / Token 预算)需要按场景动态调整 max_tokens、
|
||
// temperature 等参数。但 LLMClient.Chat 接口签名固定为 (ctx, messages, tools),
|
||
// 不能扩。所以本结构配合 OptAwareClient 可选接口使用:实现此接口的客户端
|
||
// 可以接受运行时参数。
|
||
type ChatOpts struct {
|
||
MaxTokens int // 单次响应最大 token 数(透传厂商 max_tokens)
|
||
Temperature float64 // 温度(0-2)
|
||
TopP float64 // nucleus sampling
|
||
}
|
||
|
||
// OptAwareClient 可选接口:支持运行时传 ChatOpts 的客户端
|
||
//
|
||
// 实现方:DeepSeekClient / SparkClient / OpenAIClient
|
||
// 调用方:ReactLoop、EnhancerService(通过类型断言判断是否支持)
|
||
type OptAwareClient interface {
|
||
// ChatWithOpts 带 opts 调用 LLM
|
||
//
|
||
// opts 字段为 0 值时客户端使用自身默认值
|
||
ChatWithOpts(ctx context.Context, messages []types.Message, tools []types.Tool, opts ChatOpts) (*types.Message, error)
|
||
}
|
||
|
||
// ========================================================================
|
||
// 工厂实现
|
||
// ========================================================================
|
||
|
||
// ProviderFactory 模型工厂
|
||
//
|
||
// 负责根据 Provider 名称创建对应的 LLMClient 实例。
|
||
// 采用"注册模式 + 工厂方法"的组合:
|
||
// - 内置常见供应商的创建逻辑
|
||
// - 支持外部注册自定义供应商
|
||
type ProviderFactory struct {
|
||
mu sync.RWMutex
|
||
creators map[string]CreatorFunc // 已注册的创建函数
|
||
configs map[string]*config.LLMConfigEx
|
||
}
|
||
|
||
// CreatorFunc 创建 LLMClient 的函数签名
|
||
// 外部扩展时实现此函数并注册到工厂
|
||
type CreatorFunc func(cfg *config.LLMConfigEx) (LLMClient, error)
|
||
|
||
// NewProviderFactory 创建模型工厂(注册所有内置供应商)
|
||
func NewProviderFactory(cfgs map[string]*config.LLMConfigEx) *ProviderFactory {
|
||
f := &ProviderFactory{
|
||
creators: make(map[string]CreatorFunc),
|
||
configs: cfgs,
|
||
}
|
||
|
||
// 注册内置供应商
|
||
f.Register("deepseek", createDeepSeekClient)
|
||
f.Register("openai", createOpenAIClient)
|
||
f.Register("azure", createAzureClient)
|
||
f.Register("ollama", createOllamaClient)
|
||
f.Register("qwen", createQwenClient)
|
||
f.Register("spark", createSparkClient) // 讯飞星火(OpenAPI 协议)
|
||
f.Register("mock", createMockClient) // 测试用
|
||
|
||
return f
|
||
}
|
||
|
||
// Register 注册自定义模型供应商
|
||
//
|
||
// 使用示例:
|
||
// factory.Register("my-custom-llm", func(cfg *config.LLMConfigEx) (LLMClient, error) {
|
||
// return &MyCustomClient{...}, nil
|
||
// })
|
||
func (f *ProviderFactory) Register(provider string, creator CreatorFunc) {
|
||
f.mu.Lock()
|
||
defer f.mu.Unlock()
|
||
f.creators[provider] = creator
|
||
log.Printf("[LLM工厂] 注册模型供应商: %s", provider)
|
||
}
|
||
|
||
// Create 根据供应商名称创建 LLMClient
|
||
//
|
||
// 参数:
|
||
// provider - 供应商名称("deepseek"/"openai" 等)
|
||
//
|
||
// 返回:
|
||
// 对应的 LLMClient 实例
|
||
//
|
||
// 错误:
|
||
// 供应商未注册 / 配置缺失 / 初始化失败
|
||
func (f *ProviderFactory) Create(provider string) (LLMClient, error) {
|
||
f.mu.RLock()
|
||
creator, ok := f.creators[provider]
|
||
cfg, hasCfg := f.configs[provider]
|
||
f.mu.RUnlock()
|
||
|
||
if !ok {
|
||
return nil, fmt.Errorf("[LLM工厂] 未注册的模型供应商: %s(已注册: %v)",
|
||
provider, f.ListProviders())
|
||
}
|
||
if !hasCfg {
|
||
return nil, fmt.Errorf("[LLM工厂] 供应商 %s 缺少配置信息", provider)
|
||
}
|
||
|
||
client, err := creator(cfg)
|
||
if err != nil {
|
||
return nil, fmt.Errorf("[LLM工厂] 创建 %s 客户端失败: %w", provider, err)
|
||
}
|
||
|
||
log.Printf("[LLM工厂] ✅ 成功创建模型: %s (模型名: %s, 地址: %s)",
|
||
provider, cfg.Model, cfg.BaseURL)
|
||
return client, nil
|
||
}
|
||
|
||
// ListProviders 列出所有已注册的供应商
|
||
func (f *ProviderFactory) ListProviders() []string {
|
||
f.mu.RLock()
|
||
defer f.mu.RUnlock()
|
||
names := make([]string, 0, len(f.creators))
|
||
for name := range f.creators {
|
||
names = append(names, name)
|
||
}
|
||
return names
|
||
}
|
||
|
||
// GetConfig 获取指定供应商的配置
|
||
func (f *ProviderFactory) GetConfig(provider string) (*config.LLMConfigEx, bool) {
|
||
f.mu.RLock()
|
||
defer f.mu.RUnlock()
|
||
cfg, ok := f.configs[provider]
|
||
return cfg, ok
|
||
}
|
||
|
||
// ========================================================================
|
||
// 模型路由 —— 按场景自动选择最合适的模型
|
||
// ========================================================================
|
||
|
||
// ModelRouter 模型路由器
|
||
//
|
||
// 核心职责:将"业务场景"映射到"具体模型",实现"不同任务用不同模型"。
|
||
//
|
||
// 为什么要路由?
|
||
// 不同的 Agent 任务对模型的要求不同:
|
||
// - 病历生成:需要强推理 + 医学知识 → 用 DeepSeek V3 / GPT-4o
|
||
// - 处方校验:需要严谨规则判断 → 用 GPT-4o / Claude
|
||
// - 简单问答:轻量快速即可 → 用 DeepSeek V2-Lite / 本地 Ollama
|
||
// - Embedding:纯向量化 → 用专门的 Embedding 模型
|
||
//
|
||
// 通过路由配置,可以灵活组合,且随时切换。
|
||
type ModelRouter struct {
|
||
factory *ProviderFactory
|
||
routes map[string]string // 场景名 → 供应商名
|
||
clients map[string]LLMClient // 已创建的客户端缓存(单例复用,按 provider 名)
|
||
configClients map[string]*clientCacheEntry // 按"运行时完整配置"缓存(支持配置变更自动重建)
|
||
mu sync.RWMutex
|
||
defaultProvider string
|
||
}
|
||
|
||
// NewModelRouter 创建模型路由器
|
||
//
|
||
// 参数:
|
||
// factory - 模型工厂实例
|
||
// routes - 场景到供应商的映射(如 "emr" → "deepseek")
|
||
// defaultProvider - 默认供应商(找不到路由时使用)
|
||
func NewModelRouter(factory *ProviderFactory, routes map[string]string, defaultProvider string) *ModelRouter {
|
||
return &ModelRouter{
|
||
factory: factory,
|
||
routes: routes,
|
||
clients: make(map[string]LLMClient),
|
||
configClients: make(map[string]*clientCacheEntry),
|
||
defaultProvider: defaultProvider,
|
||
}
|
||
}
|
||
|
||
// Get 根据场景名获取对应的 LLMClient
|
||
//
|
||
// 这是 Agent 代码中最常调用的方法:
|
||
// llm := router.Get("emr-generator")
|
||
// resp, err := llm.Chat(ctx, messages, tools)
|
||
func (r *ModelRouter) Get(scene string) (LLMClient, error) {
|
||
// 1. 查路由表,找到对应的供应商
|
||
provider, ok := r.routes[scene]
|
||
if !ok {
|
||
provider = r.defaultProvider
|
||
log.Printf("[模型路由] 场景 %s 未配置路由,使用默认: %s", scene, provider)
|
||
}
|
||
|
||
// 2. 检查缓存(已创建的客户端直接复用)
|
||
r.mu.RLock()
|
||
if client, ok := r.clients[provider]; ok {
|
||
r.mu.RUnlock()
|
||
return client, nil
|
||
}
|
||
r.mu.RUnlock()
|
||
|
||
// 3. 未缓存则通过工厂创建
|
||
client, err := r.factory.Create(provider)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
|
||
// 4. 写入缓存
|
||
r.mu.Lock()
|
||
r.clients[provider] = client
|
||
r.mu.Unlock()
|
||
|
||
return client, nil
|
||
}
|
||
|
||
// GetByProvider 直接按供应商名获取(绕过路由)
|
||
func (r *ModelRouter) GetByProvider(provider string) (LLMClient, error) {
|
||
r.mu.RLock()
|
||
if client, ok := r.clients[provider]; ok {
|
||
r.mu.RUnlock()
|
||
return client, nil
|
||
}
|
||
r.mu.RUnlock()
|
||
|
||
client, err := r.factory.Create(provider)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
r.mu.Lock()
|
||
r.clients[provider] = client
|
||
r.mu.Unlock()
|
||
return client, nil
|
||
}
|
||
|
||
// 内部结构:client 缓存里同时保存"配置指纹",用于检测 DB 配置变更
|
||
//
|
||
// 为什么需要指纹:DB 是配置单一可信源,后台改完 1 分钟内要生效;
|
||
// 没有指纹就只能用 client 创建时间,无法判断"是否需要重建"。
|
||
// cfgFingerprint 用 model + api_key_id + base_url + api_key 后 4 位组成(api_key 不全打出来避免泄露)
|
||
type clientCacheEntry struct {
|
||
client LLMClient
|
||
fingerprint string // 配置指纹,变化时需要重建 client
|
||
}
|
||
|
||
// computeFingerprint 计算配置指纹
|
||
//
|
||
// 用于判断"DB 里的配置是否变了"——
|
||
// 只要 model / base_url / api_key_id / api_key 拼接出的字符串变了,
|
||
// 指纹就变,从而触发 client 重建。
|
||
func computeFingerprint(cfg *config.LLMConfigEx) string {
|
||
// api_key 只取后 4 位避免日志/内存中泄露完整密钥
|
||
keyTail := ""
|
||
if len(cfg.APIKey) > 4 {
|
||
keyTail = cfg.APIKey[len(cfg.APIKey)-4:]
|
||
} else {
|
||
keyTail = cfg.APIKey
|
||
}
|
||
return fmt.Sprintf("%s|%s|%s|%s", cfg.Model, cfg.BaseURL, keyTail, cfg.Provider)
|
||
}
|
||
|
||
// GetByConfig 按完整配置获取 client,配置变更时自动重建
|
||
//
|
||
// 这是支持"DB 配置即时生效"的核心方法——
|
||
// 调用方每次都传最新的 cfg(来自 dao.LoadActiveLLMConfig 的结果),
|
||
// ModelRouter 通过指纹比对决定是复用现有 client 还是重建。
|
||
//
|
||
// 为什么不复用 GetByProvider:GetByProvider 只按 provider 名缓存,
|
||
// 无法感知配置变更(model/api_key 变了它不知道,会用旧 client)。
|
||
//
|
||
// 入参 cacheKey:缓存键,一般传 provider 名(同一 provider 多次调用复用同一条目)
|
||
// 入参 cfg:完整配置(含 Provider/APIKey/BaseURL/Model/Timeout 等)
|
||
//
|
||
// 返回:可用的 LLMClient(可能是新建的也可能是复用的)
|
||
func (r *ModelRouter) GetByConfig(cacheKey string, cfg *config.LLMConfigEx) (LLMClient, error) {
|
||
if cfg == nil {
|
||
return nil, fmt.Errorf("[模型路由] cfg 为空")
|
||
}
|
||
|
||
newFingerprint := computeFingerprint(cfg)
|
||
|
||
r.mu.Lock()
|
||
defer r.mu.Unlock()
|
||
|
||
// 同一 cacheKey 已有 client 且指纹一致 → 直接复用
|
||
if entry, ok := r.configClients[cacheKey]; ok && entry.fingerprint == newFingerprint {
|
||
return entry.client, nil
|
||
}
|
||
|
||
// 指纹变了或首次创建:通过工厂创建新 client
|
||
// 工厂需要 cfg 在 configs map 里能查到,这里临时写一份进去
|
||
r.factory.configs[cfg.Provider] = cfg
|
||
client, err := r.factory.Create(cfg.Provider)
|
||
if err != nil {
|
||
return nil, fmt.Errorf("[模型路由] 按 config 创建 %s 失败: %w", cfg.Provider, err)
|
||
}
|
||
|
||
// 关闭旧 client(如果有)释放 HTTP 连接池
|
||
if old, ok := r.configClients[cacheKey]; ok {
|
||
old.client.Close()
|
||
}
|
||
|
||
r.configClients[cacheKey] = &clientCacheEntry{
|
||
client: client,
|
||
fingerprint: newFingerprint,
|
||
}
|
||
|
||
reason := "首次创建"
|
||
if _, existed := r.clients[cacheKey]; existed {
|
||
reason = "配置变更重建"
|
||
}
|
||
log.Printf("[模型路由] %s: cacheKey=%s provider=%s model=%s fp=%s",
|
||
reason, cacheKey, cfg.Provider, cfg.Model, newFingerprint)
|
||
return client, nil
|
||
}
|
||
|
||
// RegisterRoute 动态注册/修改路由规则
|
||
func (r *ModelRouter) RegisterRoute(scene, provider string) {
|
||
r.mu.Lock()
|
||
defer r.mu.Unlock()
|
||
r.routes[scene] = provider
|
||
log.Printf("[模型路由] 注册路由: %s → %s", scene, provider)
|
||
}
|
||
|
||
// ListRoutes 列出所有路由规则
|
||
func (r *ModelRouter) ListRoutes() map[string]string {
|
||
r.mu.RLock()
|
||
defer r.mu.RUnlock()
|
||
result := make(map[string]string, len(r.routes))
|
||
for k, v := range r.routes {
|
||
result[k] = v
|
||
}
|
||
return result
|
||
}
|
||
|
||
// ResolveProvider 反查场景对应的 provider 名(不创建客户端,仅返回字符串)
|
||
//
|
||
// 用于上层记录"实际使用了哪个 provider"到审计日志/DB 子表
|
||
// 若场景未注册返回空串
|
||
func (r *ModelRouter) ResolveProvider(scene string) string {
|
||
r.mu.RLock()
|
||
defer r.mu.RUnlock()
|
||
return r.routes[scene]
|
||
}
|
||
|
||
// Close 关闭所有客户端连接
|
||
func (r *ModelRouter) Close() {
|
||
r.mu.Lock()
|
||
defer r.mu.Unlock()
|
||
for name, client := range r.clients {
|
||
client.Close()
|
||
log.Printf("[模型路由] 关闭模型连接: %s", name)
|
||
}
|
||
r.clients = make(map[string]LLMClient)
|
||
// 同时关闭 configClients(DB 驱动创建的 client 不在 clients map 里)
|
||
for key, entry := range r.configClients {
|
||
entry.client.Close()
|
||
log.Printf("[模型路由] 关闭 configClient: %s", key)
|
||
}
|
||
r.configClients = make(map[string]*clientCacheEntry)
|
||
}
|
||
|
||
// invalidateConfigClients 清空 configClients 缓存(不关闭 client,避免正在调用中的请求中断)
|
||
//
|
||
// 设计权衡:
|
||
// - 配置变更后旧 client 立即不可用 → 但已 in-flight 的 HTTP 请求通常几百毫秒就完成
|
||
// - 关闭旧 client 会触发底层连接池中断,可能导致刚发起的请求失败
|
||
// - 折中:清空缓存 map(让下次 GetByConfig 重建),但不主动 Close 旧 client,
|
||
// 让 Go GC 在所有引用消失后自动回收(HTTP 连接池的 idle conns 会在 KeepAlive 超时后自然关闭)
|
||
func (r *ModelRouter) invalidateConfigClients() {
|
||
r.mu.Lock()
|
||
count := len(r.configClients)
|
||
r.configClients = make(map[string]*clientCacheEntry)
|
||
r.mu.Unlock()
|
||
log.Printf("[模型路由] 已清空 configClients 缓存(%d 条),下次请求将按最新配置重建", count)
|
||
}
|
||
|
||
// ========================================================================
|
||
// 降级策略 —— 主模型挂了自动切换备用
|
||
// ========================================================================
|
||
|
||
// FallbackChain 降级链
|
||
//
|
||
// 当主模型 API 超时或报错时,按预设顺序尝试备用模型。
|
||
// 确保 Agent 服务的高可用。
|
||
type FallbackChain struct {
|
||
router *ModelRouter
|
||
chains map[string][]string // 场景 → 降级顺序列表
|
||
mu sync.RWMutex
|
||
}
|
||
|
||
// NewFallbackChain 创建降级链
|
||
//
|
||
// chains 示例:
|
||
// {
|
||
// "emr-generator": ["deepseek", "qwen", "ollama"], // 主→备1→备2
|
||
// "prescription": ["openai", "deepseek"],
|
||
// }
|
||
func NewFallbackChain(router *ModelRouter, chains map[string][]string) *FallbackChain {
|
||
return &FallbackChain{
|
||
router: router,
|
||
chains: chains,
|
||
}
|
||
}
|
||
|
||
// ChatWithFallback 带降级的对话调用
|
||
//
|
||
// 依次尝试链中的模型,第一个成功的即返回。
|
||
// 全部失败则返回最后一个错误。
|
||
func (fc *FallbackChain) ChatWithFallback(
|
||
ctx context.Context,
|
||
scene string,
|
||
messages []types.Message,
|
||
tools []types.Tool,
|
||
) (*types.Message, error) {
|
||
fc.mu.RLock()
|
||
chain, ok := fc.chains[scene]
|
||
fc.mu.RUnlock()
|
||
|
||
if !ok || len(chain) == 0 {
|
||
// 无降级配置,直接走正常路由
|
||
client, err := fc.router.Get(scene)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
return client.Chat(ctx, messages, tools)
|
||
}
|
||
|
||
// 按降级链依次尝试
|
||
var lastErr error
|
||
for i, provider := range chain {
|
||
client, err := fc.router.GetByProvider(provider)
|
||
if err != nil {
|
||
log.Printf("[降级] 获取 %s 失败: %v,尝试下一个", provider, err)
|
||
lastErr = err
|
||
continue
|
||
}
|
||
|
||
resp, err := client.Chat(ctx, messages, tools)
|
||
if err == nil {
|
||
if i > 0 {
|
||
log.Printf("[降级] ✅ 使用备用模型 %s 成功(主模型不可用)", provider)
|
||
}
|
||
return resp, nil
|
||
}
|
||
log.Printf("[降级] %s 调用失败: %v,尝试下一个", provider, err)
|
||
lastErr = err
|
||
}
|
||
|
||
return nil, fmt.Errorf("[降级] 所有模型均不可用,最后错误: %w", lastErr)
|
||
}
|
||
|
||
// RegisterChain 注册降级链
|
||
func (fc *FallbackChain) RegisterChain(scene string, providers []string) {
|
||
fc.mu.Lock()
|
||
defer fc.mu.Unlock()
|
||
fc.chains[scene] = providers
|
||
log.Printf("[降级] 注册降级链: %s → %v", scene, providers)
|
||
}// GetChain 读取某个 scene 的降级链(不创建客户端,仅返回字符串列表)
|
||
//
|
||
// 用于 enhancer.go 在主模型失败时手动遍历重试:
|
||
// - 拿到链后跳过已经失败的 provider,逐个 GetByProvider + callLLMWithMeta
|
||
// - 比 ChatWithFallback 更灵活,能保留 token 统计到 EnhanceStep
|
||
//
|
||
// 不存在返回 (nil, false)
|
||
func (fc *FallbackChain) GetChain(scene string) ([]string, bool) {
|
||
fc.mu.RLock()
|
||
defer fc.mu.RUnlock()
|
||
chain, ok := fc.chains[scene]
|
||
if !ok {
|
||
return nil, false
|
||
}
|
||
// 拷贝一份,避免调用方误改内部状态
|
||
out := make([]string, len(chain))
|
||
copy(out, chain)
|
||
return out, true
|
||
}
|
||
|
||
// ========================================================================
|
||
// 便捷初始化函数
|
||
// ========================================================================
|
||
|
||
// 包级全局变量:当前 ModelRouter 单例
|
||
//
|
||
// 用途:InvalidateConfigClientCache 这种"运维接口"在 handler 里无法直接拿到 router 引用,
|
||
// 通过包级变量让外部能调到 router 内部的清理逻辑。
|
||
//
|
||
// 启动时由 InitLLM 写入;进程生命周期内不变。
|
||
var globalRouter *ModelRouter
|
||
|
||
// InitLLM 一键初始化整个LLM层(工厂+路由+降级)
|
||
//
|
||
// 这是 main.go 中调用的入口函数,根据配置文件自动装配所有模型。
|
||
//
|
||
// 参数:
|
||
// cfg - 全局配置(包含 LLM 多模型配置)
|
||
//
|
||
// 返回:
|
||
// router - 模型路由器(Agent 日常使用)
|
||
// fallback - 降级链(带高可用保障)
|
||
// factory - 工厂(需要动态创建时用)
|
||
func InitLLM(cfg *config.Config) (*ModelRouter, *FallbackChain, *ProviderFactory) {
|
||
// 1. 构建配置映射
|
||
configs := make(map[string]*config.LLMConfigEx)
|
||
for name, llmCfg := range cfg.LLM.Models {
|
||
configs[name] = &llmCfg
|
||
}
|
||
|
||
// 2. 创建工厂
|
||
factory := NewProviderFactory(configs)
|
||
|
||
// 3. 构建路由表
|
||
// 从配置中读取场景→模型的映射,没有则用默认值
|
||
routes := cfg.LLM.Routes
|
||
if routes == nil {
|
||
// 默认路由:所有场景走默认模型
|
||
routes = make(map[string]string)
|
||
for scene := range defaultSceneRoutes() {
|
||
routes[scene] = cfg.LLM.DefaultProvider
|
||
}
|
||
}
|
||
|
||
router := NewModelRouter(factory, routes, cfg.LLM.DefaultProvider)
|
||
|
||
// 4. 构建降级链
|
||
fallback := NewFallbackChain(router, cfg.LLM.FallbackChains)
|
||
|
||
// ★ 把 router 存到包级变量,让运维接口能调到
|
||
globalRouter = router
|
||
|
||
log.Printf("[LLM初始化] ✅ 完成 | 已注册模型: %v | 默认: %s",
|
||
factory.ListProviders(), cfg.LLM.DefaultProvider)
|
||
|
||
return router, fallback, factory
|
||
}
|
||
|
||
// InvalidateConfigClientCache 清空所有"按 config 创建"的 client 缓存 + dao 的 active 配置缓存
|
||
//
|
||
// 使用场景:PHP 后台改完 xk_system_config 的 ai_active_* 后,
|
||
// 调 POST /api/v1/models/invalidate-cache,Go 端立刻:
|
||
// 1. 清 dao.LoadActiveLLMConfig 的 60s 缓存 → 下次读 DB 拿到新 provider/model
|
||
// 2. 清 ModelRouter.configClients 的指纹缓存 → 下次按新 config 重建 client
|
||
//
|
||
// 注意:不会清 clients map(yaml 路由的兜底 client 保留,避免 DB 抖动时全断)
|
||
func InvalidateConfigClientCache() {
|
||
// 先失效 dao 缓存,让下次 LoadActiveLLMConfig 重新查 DB
|
||
dao.InvalidateActiveLLMCache()
|
||
// 再失效 router 的 configClients 缓存
|
||
if globalRouter != nil {
|
||
globalRouter.invalidateConfigClients()
|
||
}
|
||
}
|
||
|
||
// PeekActiveLLMConfig 探测当前生效的 LLM 配置(绕过 dao 缓存)
|
||
//
|
||
// 用途:GET /api/v1/models/active-config 排查接口调用,
|
||
// 不走缓存直读 DB,让运维看到的就是当前 DB 里的真实值
|
||
func PeekActiveLLMConfig(fallbackProvider string) (*dao.ResolvedLLMConfig, error) {
|
||
// 先失效 dao 缓存,保证读到的是 DB 最新值
|
||
dao.InvalidateActiveLLMCache()
|
||
return dao.LoadActiveLLMConfig(fallbackProvider)
|
||
}
|
||
|
||
// defaultSceneRoutes 返回推荐的场景→模型路由(用户未配置时使用)
|
||
//
|
||
// 注意:scene 名必须与 PHP TcmAgentClient 透传的 scene 字段对齐,
|
||
// 否则 PHP 调过来路由表查不到,会走默认 provider,日志会出现
|
||
// "[模型路由] 场景 medical_record 未配置路由,使用默认"
|
||
// 并导致 enhancer 的 provider 字段为空串
|
||
func defaultSceneRoutes() map[string]string {
|
||
return map[string]string{
|
||
"medical_record": "emr-model", // 病历生成(PHP 业务侧命名,TcmAgentClient 透传)
|
||
"prescription": "rx-model", // 处方生成
|
||
"emr-generator": "emr-model", // 病历生成(Go 内部命名兼容)
|
||
"knowledge-qa": "qa-model", // 知识问答
|
||
"embedding": "embedding-model", // 向量化
|
||
"fallback": "fallback-model", // 降级兜底
|
||
}
|
||
}
|
||
|
||
// ========================================================================
|
||
// 环境变量覆盖(方便容器化部署)
|
||
// ========================================================================
|
||
|
||
// applyEnvOverrides 用环境变量覆盖配置中的敏感信息
|
||
//
|
||
// 优先级:环境变量 > 配置文件 > 默认值
|
||
// 这样在 K8s/Docker 中部署时,API Key 通过 Secret 注入,不写进配置文件
|
||
func applyEnvOverrides(configs map[string]*config.LLMConfigEx) {
|
||
// 通用 API Key 覆盖
|
||
if key := os.Getenv("LLM_API_KEY"); key != "" {
|
||
for _, cfg := range configs {
|
||
cfg.APIKey = key
|
||
}
|
||
}
|
||
|
||
// 按供应商分别覆盖
|
||
envMappings := map[string]string{
|
||
"DEEPSEEK_API_KEY": "deepseek",
|
||
"OPENAI_API_KEY": "openai",
|
||
"AZURE_API_KEY": "azure",
|
||
"QWEN_API_KEY": "qwen",
|
||
"OLLAMA_URL": "ollama",
|
||
"SPARK_API_KEY": "spark",
|
||
}
|
||
|
||
for envKey, provider := range envMappings {
|
||
if val := os.Getenv(envKey); val != "" {
|
||
if cfg, ok := configs[provider]; ok {
|
||
if strings.Contains(envKey, "URL") {
|
||
cfg.BaseURL = val
|
||
} else {
|
||
cfg.APIKey = val
|
||
}
|
||
}
|
||
}
|
||
}
|
||
}
|