464 lines
15 KiB
Go
464 lines
15 KiB
Go
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]
|
||
}
|