Files
xk-ai-agent/internal/agent/runner.go
2026-08-14 21:50:48 +08:00

464 lines
15 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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]
}