package agent import ( "context" "crypto/rand" "fmt" "log" "sync" "time" "tcm-agent/internal/config" "tcm-agent/internal/llm" "tcm-agent/internal/tool" "tcm-agent/internal/types" ) // ======================================================================== // Runner —— Agent 引擎核心调度器 // ======================================================================== // 职责: // - 管理会话生命周期(创建/获取/清理) // - 管理工具注册表 // - 执行 Agent 推理循环(感知→规划→检索→工具→反思→输出) // // 关键升级: // - 不再硬编码 DeepSeek,通过 llm.ModelRouter 动态获取模型 // - 不同 Agent 场景可以绑定不同模型(病历用 DeepSeek,处方用 GPT-4o) // - 支持降级链:主模型挂了自动切备用 // // 架构关系: // Handler → Agent(EMR/Prescription) → Runner → ModelRouter → LLMClient // ↓ // ProviderFactory // ↓ // DeepSeek/OpenAI/... // ======================================================================== // Runner Agent 引擎核心调度器 type Runner struct { router *llm.ModelRouter // 模型路由器(核心升级点) fallback *llm.FallbackChain // 降级链(高可用保障) tools map[string]Tool // 注册的工具集 sessions map[string]*Session // 会话缓存(短期记忆) mu sync.RWMutex // 保护 sessions 的并发锁 maxkb *tool.MaxKBClient // MaxKB 知识库客户端 cfg *RunnerConfig // Runner 自身配置 } // RunnerConfig Runner 运行参数 type RunnerConfig struct { MaxIterations int // Agent 最大推理轮数 Timeout time.Duration // 单次调用超时 DefaultScene string // 默认场景名(用于获取模型) } // Session 单次 Agent 会话的上下文 // // 一个 Session 代表一次完整的"患者就诊"过程: // - 从患者描述主诉开始 // - 到病历生成、处方开具 // - 全程保持对话上下文 type Session struct { ID string `json:"id"` // 会话唯一 ID UserID string `json:"user_id"` // 关联用户(医生/患者) Scene string `json:"scene"` // 当前场景(emr/prescription) History []Message `json:"history"` // 对话历史(短期记忆) State map[string]any `json:"state"` // 中间状态 CreatedAt time.Time `json:"created_at"` // 创建时间 UpdatedAt time.Time `json:"updated_at"` // 最后更新 Status string `json:"status"` // running/completed/failed } // Message 单条对话消息(type 别名,向后兼容旧调用方) // // 真正的定义在 internal/types 包,目的是打破 llm <-> agent 循环依赖 type Message = types.Message // Tool Agent 可调用的工具接口(type 别名) type Tool = types.Tool // ToolCallInfo 工具调用记录(type 别名) type ToolCallInfo = types.ToolCallInfo // ======================================================================== // 初始化 // ======================================================================== // InitRunner 初始化 Agent 引擎 // // 参数: // router - 模型路由器(由 llm.InitLLM 创建) // fallback - 降级链(可为 nil,表示不启用降级) // cfg - 全局配置(用于读取 MaxKB 配置和 Agent 参数) // // 返回: // 初始化完成的 Runner 实例 func InitRunner(router *llm.ModelRouter, fallback *llm.FallbackChain, cfg interface{ GetMaxKB() MaxKBConfigGetter; GetAgent() AgentConfigGetter }) *Runner { // 从配置中提取需要的信息 maxkbCfg := config.MaxKBConfig{} var agentCfg RunnerConfig if cfg != nil { if m := cfg.GetMaxKB(); m != nil { maxkbCfg = config.MaxKBConfig{ BaseURL: m.GetBaseURL(), APIKey: m.GetAPIKey(), AppID: m.GetAppID(), } } if a := cfg.GetAgent(); a != nil { agentCfg = RunnerConfig{ MaxIterations: a.GetMaxIterations(), Timeout: time.Duration(a.GetTimeout()) * time.Second, } } } // 设置默认值 if agentCfg.MaxIterations == 0 { agentCfg.MaxIterations = 10 } if agentCfg.Timeout == 0 { agentCfg.Timeout = 120 * time.Second } r := &Runner{ router: router, fallback: fallback, tools: make(map[string]Tool), sessions: make(map[string]*Session), maxkb: tool.NewMaxKBClient(maxkbCfg), cfg: &agentCfg, } // 注册默认工具集 r.registerDefaultTools() // 启动会话清理协程 go r.cleanupExpiredSessions() log.Printf("[Runner] ✅ 初始化完成 | 工具数: %d | 最大推理轮数: %d", len(r.tools), r.cfg.MaxIterations) return r } // registerDefaultTools 注册 Agent 默认工具集 func (r *Runner) registerDefaultTools() { r.RegisterTool(tool.NewMaxKBRetrieveTool(r.maxkb)) // 知识库检索 r.RegisterTool(tool.NewHISTool(nil)) // HIS 系统查询 r.RegisterTool(tool.NewPharmacopoeiaTool(r.maxkb)) // 药典查询 r.RegisterTool(tool.NewRuleCheckTool()) // 规则引擎校验 } // MaxKBConfigGetter MaxKB 配置读取接口(type 别名指向 config 包同名接口) // // 为什么要别名:InitRunner 形参类型必须和 config.Config 的 GetMaxKB() 返回类型一致, // 否则 main.go 里传 cfg 给 InitRunner 会报"接口未实现"。 type MaxKBConfigGetter = config.MaxKBConfigGetter // AgentConfigGetter Agent 配置读取接口(type 别名指向 config 包同名接口) type AgentConfigGetter = config.AgentConfigGetter // ======================================================================== // 会话管理 // ======================================================================== // CreateSession 创建新会话 // // 参数: // userID - 用户标识(医生 ID 或患者 ID) // scene - 场景名称(对应路由表中的 key,如 "emr-generator") func (r *Runner) CreateSession(userID, scene string) *Session { r.mu.Lock() defer r.mu.Unlock() session := &Session{ ID: generateSessionID(), UserID: userID, Scene: scene, History: make([]Message, 0), State: make(map[string]any), CreatedAt: time.Now(), UpdatedAt: time.Now(), Status: "running", } r.sessions[session.ID] = session log.Printf("[Runner] 创建会话: %s | 用户: %s | 场景: %s", session.ID, userID, scene) return session } // GetSession 获取会话 func (r *Runner) GetSession(id string) (*Session, bool) { r.mu.RLock() defer r.mu.RUnlock() s, ok := r.sessions[id] return s, ok } // DeleteSession 删除会话(释放资源) func (r *Runner) DeleteSession(id string) { r.mu.Lock() defer r.mu.Unlock() delete(r.sessions, id) log.Printf("[Runner] 删除会话: %s", id) } // ======================================================================== // 工具管理 // ======================================================================== // RegisterTool 注册工具 func (r *Runner) RegisterTool(t Tool) { r.mu.Lock() defer r.mu.Unlock() r.tools[t.Name()] = t log.Printf("[Runner] 注册工具: %s - %s", t.Name(), t.Description()) } // GetTools 获取所有已注册工具 func (r *Runner) GetTools() []Tool { r.mu.RLock() defer r.mu.RUnlock() list := make([]Tool, 0, len(r.tools)) for _, t := range r.tools { list = append(list, t) } return list } // ======================================================================== // 核心:Agent 推理循环 // ======================================================================== // Run 执行一次完整的 Agent 生命周期 // // 这是整个系统最核心的方法,实现了: // // 感知 → 规划 → 检索 → 工具调用 → 反思 → 输出 // // 参数: // ctx - 上下文(支持超时取消) // sessionID - 会话 ID // userInput - 用户输入 // // 返回: // Agent 最终输出文本 // // 流程详解: // // ┌─────────────────────────────────────────────────────────────┐ // │ ① 感知:接收用户输入,加载会话记忆 │ // │ ② 获取模型:通过 Router 拿到当前场景对应的 LLM │ // │ ③ 规划+执行循环(最多 MaxIterations 轮): │ // │ a. 调用 LLM(附带工具定义) │ // │ b. LLM 决定:直接回复 or 调用工具 │ // │ c. 若调用工具 → 执行 → 结果喂回 LLM → 继续循环 │ // │ d. 若直接回复 → 结束循环 │ // │ ④ 输出:返回最终回复,更新会话状态 │ // └─────────────────────────────────────────────────────────────┘ func (r *Runner) Run(ctx context.Context, sessionID string, userInput string) (string, error) { session, ok := r.GetSession(sessionID) if !ok { return "", fmt.Errorf("[Runner] 会话不存在: %s", sessionID) } // ===== ① 感知阶段:加载上下文 ===== session.History = append(session.History, Message{ Role: "user", Content: userInput, Timestamp: time.Now().Unix(), }) session.UpdatedAt = time.Now() // ===== ② 获取当前场景对应的模型 ===== scene := session.Scene if scene == "" { scene = "default" } // ===== ③ Agent 推理循环 ===== for i := 0; i < r.cfg.MaxIterations; i++ { // --- 调用 LLM(带降级) --- var resp *Message var err error if r.fallback != nil { // 使用降级链:主模型挂了自动切备用 llmResp, fbErr := r.fallback.ChatWithFallback(ctx, scene, session.History, r.GetTools()) if fbErr != nil { return "", fmt.Errorf("[Runner] 所有模型均不可用: %w", fbErr) } resp = llmResp err = nil _ = err } else { // 直连模式:通过 Router 获取模型 client, rtErr := r.router.Get(scene) if rtErr != nil { return "", fmt.Errorf("[Runner] 获取模型失败: %w", rtErr) } resp, err = client.Chat(ctx, session.History, r.GetTools()) if err != nil { return "", fmt.Errorf("[Runner] LLM 调用失败: %w", err) } } // --- 检查是否需要调用工具 --- if resp.ToolCall != nil { // 【协议修复】先把 assistant 这一帧(含 tool_calls)入栈 // OpenAI Function Calling 协议要求 messages 数组中: // ... → assistant(tool_calls) → tool(result) → assistant(...) // 之前只 append tool 帧不 append assistant 帧,部分模型会报 // "messages must alternate between user/assistant/tool" 错误 session.History = append(session.History, *resp) // 执行工具 toolResult, toolErr := r.executeTool(ctx, resp.ToolCall) // tool 帧必须带 ToolCallID 与 assistant 帧 tool_calls[].id 对应(OpenAI 协议) toolMsg := Message{ Role: "tool", Content: toolResult, ToolCallID: resp.ToolCall.ID, Timestamp: time.Now().Unix(), } if toolErr != nil { toolMsg.Content = fmt.Sprintf("工具执行错误: %v", toolErr) log.Printf("[Runner] 工具执行失败: %s → %v", resp.ToolCall.ToolName, toolErr) } else { log.Printf("[Runner] 工具执行成功: %s → %.80s...", resp.ToolCall.ToolName, toolResult) } session.History = append(session.History, toolMsg) continue // 带着工具结果进入下一轮推理 } // --- LLM 产出最终回复 --- session.History = append(session.History, *resp) session.Status = "completed" session.UpdatedAt = time.Now() log.Printf("[Runner] ✅ Agent 完成 | 会话: %s | 推理轮数: %d", sessionID, i+1) return resp.Content, nil } // 达到最大轮数仍未完成 session.Status = "failed" return "", fmt.Errorf("[Runner] Agent 达到最大推理轮数(%d)仍未完成", r.cfg.MaxIterations) } // executeTool 执行工具调用 func (r *Runner) executeTool(ctx context.Context, call *ToolCallInfo) (string, error) { r.mu.RLock() t, ok := r.tools[call.ToolName] r.mu.RUnlock() if !ok { return "", fmt.Errorf("[Runner] 工具不存在: %s", call.ToolName) } result, err := t.Execute(ctx, call.Params) return result, err } // ======================================================================== // 会话清理 // ======================================================================== // cleanupExpiredSessions 定期清理过期会话(超过1小时) func (r *Runner) cleanupExpiredSessions() { ticker := time.NewTicker(10 * time.Minute) defer ticker.Stop() for range ticker.C { r.mu.Lock() now := time.Now() expired := 0 for id, s := range r.sessions { if now.Sub(s.UpdatedAt) > time.Hour { delete(r.sessions, id) expired++ } } r.mu.Unlock() if expired > 0 { log.Printf("[Runner] 清理 %d 个过期会话", expired) } } } // ======================================================================== // 辅助方法 // ======================================================================== // MaxKB 获取 MaxKB 客户端(供 Handler 直接调用知识库) func (r *Runner) MaxKB() *tool.MaxKBClient { return r.maxkb } // GetAllSessions 获取所有会话(管理/调试用) func (r *Runner) GetAllSessions() map[string]*Session { r.mu.RLock() defer r.mu.RUnlock() sessions := make(map[string]*Session) for k, v := range r.sessions { sessions[k] = v } return sessions } // UpdateSessionStatus 更新会话状态 func (r *Runner) UpdateSessionStatus(id string, status string) error { r.mu.Lock() defer r.mu.Unlock() session, ok := r.sessions[id] if !ok { return fmt.Errorf("[Runner] 会话不存在: %s", id) } session.Status = status session.UpdatedAt = time.Now() return nil } // GetSessionHistory 获取会话对话历史 func (r *Runner) GetSessionHistory(id string) ([]Message, error) { r.mu.RLock() defer r.mu.RUnlock() session, ok := r.sessions[id] if !ok { return nil, fmt.Errorf("[Runner] 会话不存在: %s", id) } return session.History, nil } // ListTools 列出所有已注册工具名称 func (r *Runner) ListTools() []string { r.mu.RLock() defer r.mu.RUnlock() names := make([]string, 0, len(r.tools)) for name := range r.tools { names = append(names, name) } return names } // generateSessionID 生成唯一会话 ID func generateSessionID() string { return fmt.Sprintf("sess-%d-%s", time.Now().UnixNano(), randomSuffix(6)) } // randomSuffix 生成 n 位 hex 随机后缀(用于会话 ID 防碰撞) // // 用 crypto/rand 保证多并发场景下的唯一性,比 math/rand 更安全 func randomSuffix(n int) string { if n <= 0 { n = 6 } b := make([]byte, n/2+1) if _, err := rand.Read(b); err != nil { // 极端情况下 rand 失败,回落到时间戳末位,保证不阻塞业务 return fmt.Sprintf("%x", time.Now().UnixNano()%0xffffff)[:n] } return fmt.Sprintf("%x", b)[:n] }