1215 lines
46 KiB
Go
1215 lines
46 KiB
Go
package service
|
||
|
||
import (
|
||
"context"
|
||
"encoding/json"
|
||
"fmt"
|
||
"log"
|
||
"sort"
|
||
"strconv"
|
||
"strings"
|
||
"sync"
|
||
"time"
|
||
|
||
"tcm-agent/internal/agent"
|
||
"tcm-agent/internal/agentcfg"
|
||
"tcm-agent/internal/config"
|
||
"tcm-agent/internal/dao"
|
||
"tcm-agent/internal/kb"
|
||
"tcm-agent/internal/llm"
|
||
"tcm-agent/internal/tool"
|
||
"tcm-agent/internal/types"
|
||
)
|
||
|
||
// ========================================================================
|
||
// KnowledgeEnhancer —— 知识增强服务
|
||
// ========================================================================
|
||
// 这是 Go Agent 给 PHP 提供的"知识库增强 + 模型调用"一体化入口。
|
||
//
|
||
// 业务定位:
|
||
// PHP 端依然负责"提示词拼装 + 业务编排",但有一类"固定查询"内容
|
||
// (委托调剂规则、煎法、中医字典、ICD-10、药品库等)已经录到 MaxKB,
|
||
// Go Agent 接管这部分:从知识库检索 → 追加 system 消息 → 调 LLM → 回原文。
|
||
//
|
||
// 数据流:
|
||
// PHP POST /api/v1/agent/enhance
|
||
// body: { scene, context, messages: [...], kb_enabled }
|
||
// ↓
|
||
// [1] 如果 kb_enabled=true:
|
||
// - 用 context 关键词调 MaxKB.Search 取 TopK 文档
|
||
// - 把检索结果作为 system 消息插到 messages 头部
|
||
// [2] 按 scene 路由到对应 LLM 调用 Chat
|
||
// [3] 记录每一步的 token 用量、耗时(返回给 PHP,PHP 写 xk_ai_generation_step)
|
||
// [4] 返回 { content, steps: [...], provider, model }
|
||
//
|
||
// 设计原则:
|
||
// - 不接管提示词拼装:messages 直接来自 PHP,Go 只做"增强 + 调用"
|
||
// - 失败可降级:知识库检索失败不阻断,只用原始 messages 调 LLM
|
||
// - 可观测:每一步耗时/token 都记录,PHP 据此算成本
|
||
// ========================================================================
|
||
|
||
// EnhancerService 知识增强服务
|
||
type EnhancerService struct {
|
||
maxkb *tool.MaxKBClient // MaxKB 客户端(kb_enabled=false 或 source=local 时可为 nil)
|
||
localSearch *kb.Searcher // 本地知识库检索器(V1 默认走这条路径)
|
||
llmRouter *llm.ModelRouter // 模型路由(按 scene 取 LLM)
|
||
llmFallback *llm.FallbackChain
|
||
cfg *config.Config
|
||
toolRegistry map[string]types.Tool // ReactLoop 用:工具注册表(name → tool);非空时启用 Function Calling
|
||
}
|
||
|
||
// resolveMeta resolveClient 的解析元信息(配置来源 + api_key_id)
|
||
//
|
||
// 为什么用返回值而不是存到 EnhancerService 字段:
|
||
// EnhancerService 是单例,/agent/enhance 限流器允许 16 并发——
|
||
// 共享字段会被并发请求互相覆盖(run 记录的 cfg_source/api_key_id 串号 + data race),
|
||
// 用返回值让每个请求协程持有自己的一份,天然并发安全。
|
||
type resolveMeta struct {
|
||
APIKeyID int // DB 解析到的 api_key_id(审计用,写入 EnhanceStep.APIKeyID)
|
||
CfgSource string // 配置来源:active / yaml_force / yaml_route / yaml_default
|
||
}
|
||
|
||
// NewEnhancerService 构造函数
|
||
//
|
||
// 参数:
|
||
// maxkb - MaxKB 客户端(kb 关闭或 source=local 时可传 nil)
|
||
// router - 模型路由器(必填)
|
||
// fallback - 降级链(可空)
|
||
// cfg - 全局配置
|
||
func NewEnhancerService(maxkb *tool.MaxKBClient, router *llm.ModelRouter, fallback *llm.FallbackChain, cfg *config.Config) *EnhancerService {
|
||
return &EnhancerService{
|
||
maxkb: maxkb,
|
||
llmRouter: router,
|
||
llmFallback: fallback,
|
||
cfg: cfg,
|
||
}
|
||
}
|
||
|
||
// WithLocalSearcher 注入本地知识库检索器
|
||
//
|
||
// 调用方:router.Setup 启动时根据 agentcfg.KB.EmbeddingProvider 构造 Searcher,
|
||
// 调用本方法注入。V1 默认走 NoopEmbedder(仅全文检索)。
|
||
func (s *EnhancerService) WithLocalSearcher(searcher *kb.Searcher) *EnhancerService {
|
||
s.localSearch = searcher
|
||
return s
|
||
}
|
||
|
||
// WithTools 注入工具注册表(让 ReactLoop 可触发 Function Calling)
|
||
//
|
||
// 调用方:main.go 启动时若 agent.Runner 已注册工具,把同一份 map 注入给 EnhancerService。
|
||
// 不调用本方法时 toolRegistry 为 nil,ReactLoop 会以纯 chat 模式运行(无工具调用)。
|
||
func (s *EnhancerService) WithTools(tools map[string]types.Tool) *EnhancerService {
|
||
s.toolRegistry = tools
|
||
return s
|
||
}
|
||
|
||
// ========================================================================
|
||
// 请求/响应结构(与 PHP TcmAgentClient 严格对齐)
|
||
// ========================================================================
|
||
|
||
// EnhanceRequest PHP 发来的增强请求
|
||
type EnhanceRequest struct {
|
||
Scene string `json:"scene"` // 场景:medical_record / prescription
|
||
Context string `json:"context"` // 检索关键词(如"痰湿中阻 煎法")
|
||
Messages []types.Message `json:"messages"` // PHP 拼好的完整消息列表
|
||
KBEnabled bool `json:"kb_enabled"` // 是否启用知识库检索
|
||
TopK int `json:"top_k,omitempty"` // 检索条数(默认 5)
|
||
Provider string `json:"provider,omitempty"` // 强制用某 provider(空则走路由表)
|
||
AgentConfig *AgentConfigOpts `json:"agent_config,omitempty"` // 可选:PHP 透传覆盖 DB 中的 Agent 配置(不传则 Go 直读 DB)
|
||
// 生成参数(P0 修复:此前 PHP 按场景调的参数在 via-agent 路径被丢弃)
|
||
// 0 表示"未显式指定",走客户端默认(deepseek/spark 默认 0.3 / 4096)
|
||
Temperature float64 `json:"temperature,omitempty"` // 生成温度(0-2)
|
||
MaxTokens int `json:"max_tokens,omitempty"` // 单次响应 token 上限
|
||
// Count 一次生成几份(默认 1,上限 5)。>1 时 Go 并行打多路 LLM,按槽位返回 Results;
|
||
// 单路失败不拖死整批,至少成功 1 路才算业务成功。
|
||
Count int `json:"count,omitempty"`
|
||
}
|
||
|
||
// AgentConfigOpts 可选的 Agent 配置覆盖项(PHP 透传)
|
||
//
|
||
// 设计:默认情况下 Go Agent 直读 xk_system_config(agentcfg 包);
|
||
// 但 PHP 端如果已经有完整配置视图,可以通过本字段强制覆盖某些项。
|
||
// 所有字段都是指针,nil 表示不覆盖、用 DB 配置。
|
||
type AgentConfigOpts struct {
|
||
ReactEnabled *bool `json:"react_enabled,omitempty"`
|
||
ReactMaxIterations *int `json:"react_max_iterations,omitempty"`
|
||
}
|
||
|
||
// EnhanceStep 一次增强调用的过程步骤(写入 xk_ai_generation_step)
|
||
type EnhanceStep struct {
|
||
StepType string `json:"step_type"` // kb_retrieval / llm_call
|
||
Provider string `json:"provider,omitempty"` // llm_call 才填
|
||
Model string `json:"model,omitempty"` // llm_call 才填
|
||
APIKeyID int `json:"api_key_id,omitempty"`
|
||
PromptTokens int `json:"prompt_tokens,omitempty"`
|
||
CompletionTokens int `json:"completion_tokens,omitempty"`
|
||
TotalTokens int `json:"total_tokens,omitempty"`
|
||
Usage map[string]any `json:"usage,omitempty"`
|
||
DurationMs int `json:"duration_ms"`
|
||
StartedAt int64 `json:"started_at"` // unix 秒
|
||
FinishedAt int64 `json:"finished_at"`
|
||
Detail string `json:"detail,omitempty"` // 备注(如命中文档数 / 失败原因)
|
||
Status int `json:"status"` // 0进行中 1成功 2失败(与 xk_ai_generation.status 一致)
|
||
lastMessage *types.Message `json:"-"` // ReactLoop 内部用:缓存 LLM 返回的 Message(不入 JSON)
|
||
}
|
||
|
||
// EnhanceSlotResult 多份生成时每个槽位的结果(成功填 Content,失败填 Error)
|
||
type EnhanceSlotResult struct {
|
||
Index int `json:"index"` // 0-based 槽位
|
||
OK bool `json:"ok"` // 本槽是否成功
|
||
Content string `json:"content,omitempty"` // 成功时的 LLM 文本
|
||
Error string `json:"error,omitempty"` // 失败原因
|
||
Steps []EnhanceStep `json:"steps,omitempty"` // 本槽独立步骤
|
||
Provider string `json:"provider,omitempty"` // 本槽实际 provider
|
||
Model string `json:"model,omitempty"` // 本槽实际 model
|
||
}
|
||
|
||
// EnhanceResponse 返回给 PHP 的结果
|
||
type EnhanceResponse struct {
|
||
Content string `json:"content"` // 第一份成功 content(兼容旧 PHP)
|
||
Contents []string `json:"contents,omitempty"` // 仅成功内容数组(可选兼容)
|
||
Results []EnhanceSlotResult `json:"results,omitempty"` // 多份时按槽位完整结果
|
||
Provider string `json:"provider"` // 实际使用的供应商(首份成功)
|
||
Model string `json:"model"` // 实际使用的模型名(首份成功)
|
||
Steps []EnhanceStep `json:"steps"` // 共享步骤 + 汇总(单份时完整)
|
||
TotalMs int `json:"total_ms"` // 总耗时(并行时约等于最慢一路)
|
||
// cfgSource 本次请求的配置来源(内部字段,不返回 PHP):
|
||
// doEnhance 写入,Enhance 包装层读出来写 RunLog——
|
||
// 用 resp 传递而不是 service 字段,保证并发请求间不串号
|
||
cfgSource string `json:"-"`
|
||
}
|
||
|
||
// ========================================================================
|
||
// 核心方法
|
||
// ========================================================================
|
||
|
||
// Enhance 执行一次知识增强 + 模型调用(对外入口)
|
||
//
|
||
// 这是一个包装层:真正的业务逻辑在 doEnhance,本层负责把每次运行
|
||
// (成功/失败/守卫拦截)统一记入 RunLog 环形缓冲,供 /agent/view 面板观测。
|
||
// 包装层方案的好处:doEnhance 内部有 5 个 return 出口,不用逐个插埋点代码。
|
||
func (s *EnhancerService) Enhance(ctx context.Context, req *EnhanceRequest) (*EnhanceResponse, error) {
|
||
startedAt := time.Now()
|
||
resp, err := s.doEnhance(ctx, req)
|
||
// cfgSource 由 doEnhance 写入 resp(解析失败/守卫拦截时为空)
|
||
cfgSource := ""
|
||
if resp != nil {
|
||
cfgSource = resp.cfgSource
|
||
}
|
||
recordAgentRun(req, resp, err, startedAt, cfgSource)
|
||
return resp, err
|
||
}
|
||
|
||
// doEnhance 知识增强主流程(内部实现)
|
||
func (s *EnhancerService) doEnhance(ctx context.Context, req *EnhanceRequest) (*EnhanceResponse, error) {
|
||
totalStart := time.Now()
|
||
resp := &EnhanceResponse{Steps: []EnhanceStep{}}
|
||
|
||
if len(req.Messages) == 0 {
|
||
return nil, fmt.Errorf("messages 不能为空")
|
||
}
|
||
|
||
// 复制一份 messages,避免污染入参(追加 KB 检索结果时使用)
|
||
messages := make([]types.Message, len(req.Messages))
|
||
copy(messages, req.Messages)
|
||
|
||
// ===== 步骤 0:医疗相关性前置守卫 =====
|
||
// 在做任何 KB 检索 / LLM 调用之前先检查请求是否医疗相关:
|
||
// - PHP 端 prompt 组装出错(或接口被滥用)时,轻量模型会跑题输出通用内容
|
||
// - 提前拦截可以省一次完整 ReactLoop 的 token 花费
|
||
// 守卫默认开启,可通过 xk_system_config 的 ai_agent_medical_guard 关闭
|
||
guardCfg := agentcfg.Get()
|
||
if guardCfg.MedicalGuard.Enabled {
|
||
guard := validateMedicalRelevance(messages)
|
||
if !guard.Passed {
|
||
log.Printf("[Enhancer] ⛔ 医疗守卫拦截 scene=%s 命中黑名单=%v 命中白名单=%v",
|
||
req.Scene, guard.HitBlacklist, guard.HitWhitelist)
|
||
guardStep := EnhanceStep{
|
||
StepType: "medical_guard",
|
||
StartedAt: time.Now().Unix(),
|
||
FinishedAt: time.Now().Unix(),
|
||
Status: 2,
|
||
Detail: guard.Reason,
|
||
}
|
||
resp.Steps = append(resp.Steps, guardStep)
|
||
resp.TotalMs = int(time.Since(totalStart).Milliseconds())
|
||
return resp, fmt.Errorf("医疗守卫拦截: %s", guard.Reason)
|
||
}
|
||
// 通过时打一行简短日志(白名单命中数),便于观察守卫是否在正常工作
|
||
log.Printf("[Enhancer] 医疗守卫通过 scene=%s 白名单命中=%d", req.Scene, len(guard.HitWhitelist))
|
||
}
|
||
|
||
// ===== 步骤 1:知识库检索(可选) =====
|
||
// 根据 agentcfg.KB.Source 分流到 MaxKB 或本地知识库;
|
||
// 处方场景额外检索金方库(参考方剂),与药材知识分块注入
|
||
if req.KBEnabled && strings.TrimSpace(req.Context) != "" {
|
||
step, kbDocs, kbFormulas, kerr := s.performKBRetrieval(ctx, req)
|
||
if step != nil {
|
||
resp.Steps = append(resp.Steps, *step)
|
||
// 检索成功时把结果合并进首条 system 消息
|
||
if kerr == nil && (len(kbDocs) > 0 || len(kbFormulas) > 0) {
|
||
messages = injectKBContext(messages, kbDocs, kbFormulas)
|
||
}
|
||
}
|
||
}
|
||
|
||
// ===== 步骤 2:解析 LLM 客户端 =====
|
||
client, provider, meta, err := s.resolveClient(req.Scene, req.Provider)
|
||
resp.cfgSource = meta.CfgSource
|
||
if err != nil {
|
||
errStep := EnhanceStep{
|
||
StepType: "llm_call",
|
||
StartedAt: time.Now().Unix(),
|
||
FinishedAt: time.Now().Unix(),
|
||
DurationMs: 0,
|
||
Status: 2,
|
||
Detail: "解析 LLM 失败: " + err.Error(),
|
||
}
|
||
resp.Steps = append(resp.Steps, errStep)
|
||
resp.TotalMs = int(time.Since(totalStart).Milliseconds())
|
||
return resp, fmt.Errorf("解析 LLM 失败: %w", err)
|
||
}
|
||
|
||
// ===== 步骤 3:根据 AgentConfig 决定走 ReactLoop 还是单次 Chat =====
|
||
agentCfg := agentcfg.Get()
|
||
// PHP 透传的覆盖项优先
|
||
reactOn := agentCfg.ReAct.Enabled
|
||
maxIter := agentCfg.ReAct.MaxIterations
|
||
if req.AgentConfig != nil {
|
||
if req.AgentConfig.ReactEnabled != nil {
|
||
reactOn = *req.AgentConfig.ReactEnabled
|
||
}
|
||
if req.AgentConfig.ReactMaxIterations != nil {
|
||
maxIter = *req.AgentConfig.ReactMaxIterations
|
||
}
|
||
}
|
||
|
||
// 多份并行:KB/守卫/客户端解析只做一次,LLM 按槽位 goroutine 打;
|
||
// 处方严格 JSON 场景 PHP 已关 React,多份路径统一走单次 Chat,避免 N×多轮成本爆炸
|
||
count := normalizeEnhanceCount(req.Count)
|
||
if count > 1 {
|
||
return s.doEnhanceMulti(ctx, req, messages, client, provider, meta, agentCfg, totalStart, resp, count)
|
||
}
|
||
|
||
if reactOn {
|
||
// ---- 路径 A:ReactLoop(多轮 + Planning + Reflection + JSON 修复)----
|
||
reactReq := &ReactLoopRequest{
|
||
Scene: req.Scene,
|
||
Messages: messages,
|
||
Client: client,
|
||
Provider: provider,
|
||
Temperature: req.Temperature, // PHP 场景化温度透传(0=用默认)
|
||
}
|
||
// 工具集:若 EnhancerService 注入了 toolRegistry,转换为 []types.Tool
|
||
if s.toolRegistry != nil {
|
||
tools := make([]types.Tool, 0, len(s.toolRegistry))
|
||
for _, t := range s.toolRegistry {
|
||
tools = append(tools, t)
|
||
}
|
||
reactReq.Tools = tools
|
||
}
|
||
// 应用 PHP 透传的 maxIter 覆盖
|
||
reactCfg := agentCfg.ReAct
|
||
if maxIter > 0 {
|
||
reactCfg.MaxIterations = maxIter
|
||
}
|
||
loopResult := s.runReactLoop(ctx, reactReq, reactCfg, agentCfg.TokenBudget)
|
||
|
||
// 把 ReactLoop 的步骤合并进 resp.Steps
|
||
resp.Steps = append(resp.Steps, loopResult.Steps...)
|
||
resp.Content = loopResult.Content
|
||
resp.Provider = provider
|
||
resp.Model = client.Name()
|
||
resp.TotalMs = int(time.Since(totalStart).Milliseconds())
|
||
|
||
// 给 ReactLoop 的所有 step 补 api_key_id(来自 DB 解析,client 内部拿不到)
|
||
if meta.APIKeyID > 0 {
|
||
for i := range resp.Steps {
|
||
if resp.Steps[i].APIKeyID == 0 {
|
||
resp.Steps[i].APIKeyID = meta.APIKeyID
|
||
}
|
||
}
|
||
}
|
||
|
||
logTag := "ReactLoop"
|
||
if loopResult.Aborted {
|
||
logTag = "ReactLoop(BudgetAborted)"
|
||
}
|
||
if raw, err := json.Marshal(resp.Steps); err == nil {
|
||
log.Printf("[Enhancer] %s scene=%s provider=%s source=%s total=%dms steps=%d budget=%s raw=%s",
|
||
logTag, req.Scene, provider, meta.CfgSource, resp.TotalMs, len(resp.Steps),
|
||
jsonSnapshot(loopResult.Budget), string(raw))
|
||
}
|
||
return resp, nil
|
||
}
|
||
|
||
// ---- 路径 B:单次 Chat(兼容原有行为,但补 token 统计 + max_tokens 透传)----
|
||
llmStep := EnhanceStep{
|
||
StepType: "llm_call",
|
||
StartedAt: time.Now().Unix(),
|
||
}
|
||
llmStart := time.Now()
|
||
|
||
// 按 agentcfg 决定 max_tokens(开启预算管理时透传 MaxTokensPerCall)
|
||
opts := llm.ChatOpts{}
|
||
if agentCfg.TokenBudget.Enabled {
|
||
opts.MaxTokens = agentCfg.TokenBudget.MaxTokensPerCall
|
||
}
|
||
// PHP 显式传入的生成参数优先(P0:处方/病历按场景调温不再被忽略)
|
||
if req.Temperature > 0 {
|
||
opts.Temperature = req.Temperature
|
||
}
|
||
if req.MaxTokens > 0 {
|
||
// 与预算上限取小:业务可以要求更小的输出,但不能借此绕开 token 预算
|
||
if opts.MaxTokens == 0 || req.MaxTokens < opts.MaxTokens {
|
||
opts.MaxTokens = req.MaxTokens
|
||
}
|
||
}
|
||
|
||
msg, chatResult := s.callLLMWithMeta(ctx, client, messages, nil, opts)
|
||
llmStep.DurationMs = int(time.Since(llmStart).Milliseconds())
|
||
llmStep.FinishedAt = time.Now().Unix()
|
||
llmStep.Provider = provider
|
||
llmStep.Model = client.Name()
|
||
|
||
// ★ api_key_id 优先用 DB 解析出来的真实 id(resolveClient 透传)
|
||
// client 内部的 chatResult.APIKeyID 通常为 0(厂商 API 不会回传 key_id),
|
||
// 只有 Go 端从 DB 取的 key 才能拿到准确 id 用于审计
|
||
if meta.APIKeyID > 0 {
|
||
llmStep.APIKeyID = meta.APIKeyID
|
||
}
|
||
if chatResult != nil {
|
||
llmStep.PromptTokens = chatResult.PromptTokens
|
||
llmStep.CompletionTokens = chatResult.CompletionTokens
|
||
llmStep.TotalTokens = chatResult.TotalTokens
|
||
llmStep.Usage = chatResult.Usage
|
||
// chatResult.APIKeyID 兜底(仅当 DB 没解析到时才有意义)
|
||
if llmStep.APIKeyID == 0 {
|
||
llmStep.APIKeyID = chatResult.APIKeyID
|
||
}
|
||
}
|
||
|
||
if msg == nil {
|
||
// 主模型调用失败,尝试从 fallback 链里取下一个 provider 重试一次
|
||
//
|
||
// 为什么不直接调 llmFallback.ChatWithFallback:
|
||
// 那个方法会按链从头开始尝试,主模型已经在上面调过失败,
|
||
// 这里只想跳过当前 provider 用下一个,避免重复打主模型
|
||
if s.llmFallback != nil && provider != "" {
|
||
sceneKey := req.Scene
|
||
if sceneKey == "" {
|
||
sceneKey = "emr-generator"
|
||
}
|
||
chain, ok := s.llmFallback.GetChain(sceneKey)
|
||
if ok {
|
||
for _, nextProvider := range chain {
|
||
if nextProvider == provider {
|
||
continue // 跳过已经失败的当前 provider
|
||
}
|
||
fbClient, fbErr := s.llmRouter.GetByProvider(nextProvider)
|
||
if fbErr != nil {
|
||
log.Printf("[Enhancer] fallback 取 %s 失败: %v", nextProvider, fbErr)
|
||
continue
|
||
}
|
||
log.Printf("[Enhancer] 主模型 %s 调用失败,尝试 fallback %s", provider, nextProvider)
|
||
fbMsg, fbChatResult := s.callLLMWithMeta(ctx, fbClient, messages, nil, opts)
|
||
if fbMsg != nil {
|
||
// fallback 成功:用新结果替换上下文中的 client/provider/chatResult
|
||
client = fbClient
|
||
provider = nextProvider
|
||
msg = fbMsg
|
||
chatResult = fbChatResult
|
||
llmStep.Provider = provider
|
||
llmStep.Model = client.Name()
|
||
if fbChatResult != nil {
|
||
llmStep.PromptTokens = fbChatResult.PromptTokens
|
||
llmStep.CompletionTokens = fbChatResult.CompletionTokens
|
||
llmStep.TotalTokens = fbChatResult.TotalTokens
|
||
llmStep.Usage = fbChatResult.Usage
|
||
llmStep.APIKeyID = fbChatResult.APIKeyID
|
||
llmStep.Detail = "成功(fallback)"
|
||
}
|
||
break
|
||
}
|
||
log.Printf("[Enhancer] fallback %s 也失败,继续尝试下一个", nextProvider)
|
||
}
|
||
}
|
||
}
|
||
}
|
||
|
||
if msg == nil {
|
||
llmStep.Status = 2
|
||
llmStep.Detail = "LLM 调用失败"
|
||
resp.Steps = append(resp.Steps, llmStep)
|
||
resp.TotalMs = int(time.Since(totalStart).Milliseconds())
|
||
return resp, fmt.Errorf("LLM 调用失败")
|
||
}
|
||
|
||
llmStep.Status = 1
|
||
if chatResult != nil && chatResult.FinishReason != "" {
|
||
llmStep.Detail = "成功 | finish_reason=" + chatResult.FinishReason
|
||
} else {
|
||
llmStep.Detail = "成功"
|
||
}
|
||
resp.Steps = append(resp.Steps, llmStep)
|
||
|
||
resp.Content = msg.Content
|
||
resp.Provider = provider
|
||
resp.Model = client.Name()
|
||
resp.TotalMs = int(time.Since(totalStart).Milliseconds())
|
||
|
||
// 把 step 详情打成日志,便于运维排查
|
||
if raw, err := json.Marshal(resp.Steps); err == nil {
|
||
log.Printf("[Enhancer] scene=%s provider=%s source=%s total=%dms steps=%s",
|
||
req.Scene, provider, meta.CfgSource, resp.TotalMs, string(raw))
|
||
}
|
||
|
||
return resp, nil
|
||
}
|
||
|
||
// normalizeEnhanceCount 钳制一次生成份数:默认 1,上限 5
|
||
func normalizeEnhanceCount(n int) int {
|
||
if n <= 1 {
|
||
return 1
|
||
}
|
||
if n > 5 {
|
||
return 5
|
||
}
|
||
return n
|
||
}
|
||
|
||
// doEnhanceMulti 多份并行 LLM:共享前置步骤,按槽位返回成功/失败
|
||
//
|
||
// 为什么并行:份数>1 时串行会把医生等待拉到 N× 单次;WaitGroup 并行后墙钟约等于最慢一路。
|
||
// 部分失败语义:某路超时/LLM 错误只标该槽 ok=false,其它成功槽照常返回;全部失败才 error。
|
||
func (s *EnhancerService) doEnhanceMulti(
|
||
ctx context.Context,
|
||
req *EnhanceRequest,
|
||
messages []types.Message,
|
||
client llm.LLMClient,
|
||
provider string,
|
||
meta resolveMeta,
|
||
agentCfg *agentcfg.AllConfig,
|
||
totalStart time.Time,
|
||
resp *EnhanceResponse,
|
||
count int,
|
||
) (*EnhanceResponse, error) {
|
||
log.Printf("[Enhancer] 多份并行开始 scene=%s count=%d provider=%s", req.Scene, count, provider)
|
||
|
||
slots := make([]EnhanceSlotResult, count)
|
||
var wg sync.WaitGroup
|
||
for i := 0; i < count; i++ {
|
||
wg.Add(1)
|
||
go func(idx int) {
|
||
defer wg.Done()
|
||
defer func() {
|
||
if r := recover(); r != nil {
|
||
slots[idx] = EnhanceSlotResult{
|
||
Index: idx,
|
||
OK: false,
|
||
Error: fmt.Sprintf("panic: %v", r),
|
||
}
|
||
}
|
||
}()
|
||
slotMsgs := cloneMessagesWithVariantHint(messages, idx, count)
|
||
content, steps, usedProvider, usedModel, err := s.runSingleChatSlot(
|
||
ctx, req, slotMsgs, client, provider, meta, agentCfg, idx,
|
||
)
|
||
if err != nil {
|
||
slots[idx] = EnhanceSlotResult{
|
||
Index: idx,
|
||
OK: false,
|
||
Error: err.Error(),
|
||
Steps: steps,
|
||
Provider: usedProvider,
|
||
Model: usedModel,
|
||
}
|
||
return
|
||
}
|
||
slots[idx] = EnhanceSlotResult{
|
||
Index: idx,
|
||
OK: true,
|
||
Content: content,
|
||
Steps: steps,
|
||
Provider: usedProvider,
|
||
Model: usedModel,
|
||
}
|
||
}(i)
|
||
}
|
||
wg.Wait()
|
||
|
||
resp.Results = slots
|
||
successContents := make([]string, 0, count)
|
||
var firstOK *EnhanceSlotResult
|
||
successN := 0
|
||
for i := range slots {
|
||
if slots[i].OK {
|
||
successN++
|
||
successContents = append(successContents, slots[i].Content)
|
||
if firstOK == nil {
|
||
cp := slots[i]
|
||
firstOK = &cp
|
||
}
|
||
}
|
||
// 汇总步骤时带上槽位前缀,便于 PHP 写子表排障
|
||
for _, st := range slots[i].Steps {
|
||
st.Detail = fmt.Sprintf("[slot%d] %s", i, st.Detail)
|
||
resp.Steps = append(resp.Steps, st)
|
||
}
|
||
}
|
||
resp.Contents = successContents
|
||
resp.TotalMs = int(time.Since(totalStart).Milliseconds())
|
||
|
||
if firstOK == nil {
|
||
log.Printf("[Enhancer] 多份并行全部失败 scene=%s count=%d total=%dms", req.Scene, count, resp.TotalMs)
|
||
return resp, fmt.Errorf("多份生成全部失败(%d 路)", count)
|
||
}
|
||
|
||
resp.Content = firstOK.Content
|
||
resp.Provider = firstOK.Provider
|
||
resp.Model = firstOK.Model
|
||
log.Printf("[Enhancer] 多份并行完成 scene=%s success=%d/%d total=%dms",
|
||
req.Scene, successN, count, resp.TotalMs)
|
||
return resp, nil
|
||
}
|
||
|
||
// cloneMessagesWithVariantHint 复制 messages,并在末条 user 追加「第 i 套方案」提示以拉开差异
|
||
func cloneMessagesWithVariantHint(src []types.Message, idx, total int) []types.Message {
|
||
out := make([]types.Message, len(src))
|
||
copy(out, src)
|
||
if total <= 1 || len(out) == 0 {
|
||
return out
|
||
}
|
||
hint := fmt.Sprintf(
|
||
"\n\n【多方案要求】请生成第 %d/%d 套组方方案,与其它套方案在用药组合或君臣佐使上尽量有可对比的差异;仍只输出约定 JSON,不要解释。",
|
||
idx+1, total,
|
||
)
|
||
// 从后往前找最后一条 user,找不到则追加
|
||
for i := len(out) - 1; i >= 0; i-- {
|
||
if strings.EqualFold(strings.TrimSpace(out[i].Role), "user") {
|
||
out[i].Content = out[i].Content + hint
|
||
return out
|
||
}
|
||
}
|
||
out = append(out, types.Message{Role: "user", Content: strings.TrimSpace(hint)})
|
||
return out
|
||
}
|
||
|
||
// runSingleChatSlot 单槽一次 Chat(含 fallback),供多份并行调用
|
||
//
|
||
// 每个 goroutine 独立温度微调(基温 + idx*0.05,上限 0.6),进一步提高方案差异。
|
||
func (s *EnhancerService) runSingleChatSlot(
|
||
ctx context.Context,
|
||
req *EnhanceRequest,
|
||
messages []types.Message,
|
||
client llm.LLMClient,
|
||
provider string,
|
||
meta resolveMeta,
|
||
agentCfg *agentcfg.AllConfig,
|
||
idx int,
|
||
) (content string, steps []EnhanceStep, usedProvider, usedModel string, err error) {
|
||
usedProvider = provider
|
||
usedModel = client.Name()
|
||
llmStep := EnhanceStep{
|
||
StepType: "llm_call",
|
||
StartedAt: time.Now().Unix(),
|
||
}
|
||
llmStart := time.Now()
|
||
|
||
opts := llm.ChatOpts{}
|
||
if agentCfg != nil && agentCfg.TokenBudget.Enabled {
|
||
opts.MaxTokens = agentCfg.TokenBudget.MaxTokensPerCall
|
||
}
|
||
baseTemp := req.Temperature
|
||
if baseTemp <= 0 {
|
||
baseTemp = 0.2
|
||
}
|
||
// 槽位温度微调:略抬高后续方案温度,拉开差异但仍偏稳
|
||
opts.Temperature = baseTemp + float64(idx)*0.05
|
||
if opts.Temperature > 0.6 {
|
||
opts.Temperature = 0.6
|
||
}
|
||
if req.MaxTokens > 0 {
|
||
if opts.MaxTokens == 0 || req.MaxTokens < opts.MaxTokens {
|
||
opts.MaxTokens = req.MaxTokens
|
||
}
|
||
}
|
||
|
||
msg, chatResult := s.callLLMWithMeta(ctx, client, messages, nil, opts)
|
||
llmStep.DurationMs = int(time.Since(llmStart).Milliseconds())
|
||
llmStep.FinishedAt = time.Now().Unix()
|
||
llmStep.Provider = usedProvider
|
||
llmStep.Model = usedModel
|
||
if meta.APIKeyID > 0 {
|
||
llmStep.APIKeyID = meta.APIKeyID
|
||
}
|
||
if chatResult != nil {
|
||
llmStep.PromptTokens = chatResult.PromptTokens
|
||
llmStep.CompletionTokens = chatResult.CompletionTokens
|
||
llmStep.TotalTokens = chatResult.TotalTokens
|
||
llmStep.Usage = chatResult.Usage
|
||
if llmStep.APIKeyID == 0 {
|
||
llmStep.APIKeyID = chatResult.APIKeyID
|
||
}
|
||
}
|
||
|
||
if msg == nil && s.llmFallback != nil && provider != "" {
|
||
sceneKey := req.Scene
|
||
if sceneKey == "" {
|
||
sceneKey = "emr-generator"
|
||
}
|
||
chain, ok := s.llmFallback.GetChain(sceneKey)
|
||
if ok {
|
||
for _, nextProvider := range chain {
|
||
if nextProvider == provider {
|
||
continue
|
||
}
|
||
fbClient, fbErr := s.llmRouter.GetByProvider(nextProvider)
|
||
if fbErr != nil {
|
||
continue
|
||
}
|
||
fbMsg, fbChatResult := s.callLLMWithMeta(ctx, fbClient, messages, nil, opts)
|
||
if fbMsg != nil {
|
||
client = fbClient
|
||
usedProvider = nextProvider
|
||
usedModel = client.Name()
|
||
msg = fbMsg
|
||
chatResult = fbChatResult
|
||
llmStep.Provider = usedProvider
|
||
llmStep.Model = usedModel
|
||
if fbChatResult != nil {
|
||
llmStep.PromptTokens = fbChatResult.PromptTokens
|
||
llmStep.CompletionTokens = fbChatResult.CompletionTokens
|
||
llmStep.TotalTokens = fbChatResult.TotalTokens
|
||
llmStep.Usage = fbChatResult.Usage
|
||
llmStep.APIKeyID = fbChatResult.APIKeyID
|
||
llmStep.Detail = "成功(fallback)"
|
||
}
|
||
break
|
||
}
|
||
}
|
||
}
|
||
}
|
||
|
||
if msg == nil {
|
||
llmStep.Status = 2
|
||
llmStep.Detail = "LLM 调用失败"
|
||
steps = append(steps, llmStep)
|
||
return "", steps, usedProvider, usedModel, fmt.Errorf("LLM 调用失败")
|
||
}
|
||
|
||
llmStep.Status = 1
|
||
if chatResult != nil && chatResult.FinishReason != "" {
|
||
llmStep.Detail = "成功 | finish_reason=" + chatResult.FinishReason
|
||
} else if llmStep.Detail == "" {
|
||
llmStep.Detail = "成功"
|
||
}
|
||
steps = append(steps, llmStep)
|
||
return msg.Content, steps, usedProvider, usedModel, nil
|
||
}
|
||
|
||
// jsonSnapshot 把 BudgetSnapshot 序列化成短字符串(仅用于日志)
|
||
func jsonSnapshot(s agent.BudgetSnapshot) string {
|
||
b, _ := json.Marshal(s)
|
||
return string(b)
|
||
}
|
||
|
||
// resolveClient 按场景 + 强制 provider 解析出 LLM 客户端
|
||
//
|
||
// ★★★ 核心升级:DB 优先 + 即时生效 ★★★
|
||
//
|
||
// 解析优先级:
|
||
// 1. req.Provider 非空 → 强制用此 provider(PHP 显式指定,覆盖一切)
|
||
// 但仍会从 DB 拉取该 provider 的最新 model/api_key(保证 key 轮换即时生效)
|
||
// 2. dao.LoadActiveLLMConfig 实时解析(带 60s 缓存):
|
||
// 读 xk_system_config 的 ai_active_provider / ai_active_api_key_id / ai_active_model,
|
||
// 按 PHP AiRuntimeConfigService::resolve() 同款规则合并到最终配置
|
||
// 这样后台切完 provider/model 1 分钟内全集群生效,无需重启 Go 进程
|
||
// 3. 若 DB 解析失败 → 回落到 yaml/env 配置(保证服务可用性)
|
||
//
|
||
// 返回值:
|
||
// - client:可用的 LLMClient(可能新建可能复用,由 ModelRouter.GetByConfig 决定)
|
||
// - provider:实际使用的 provider 名(用于审计/日志)
|
||
// - meta:解析元信息(api_key_id + 配置来源),调用方持有局部拷贝,并发安全
|
||
func (s *EnhancerService) resolveClient(scene, forceProvider string) (llm.LLMClient, string, resolveMeta, error) {
|
||
if s.llmRouter == nil {
|
||
return nil, "", resolveMeta{}, fmt.Errorf("llmRouter 未初始化")
|
||
}
|
||
|
||
// ===== 路径 1:DB 实时解析(带 60s 缓存) =====
|
||
//
|
||
// 这是默认路径——后台运维改完模型配置后最长 60s 内全集群生效。
|
||
// 失败时(如 DB 抖动)回落到老路径,保证服务可用。
|
||
if dao.DB != nil {
|
||
// 决定"DB 兜底 provider":forceProvider > yaml DefaultProvider
|
||
fallback := forceProvider
|
||
if fallback == "" {
|
||
fallback = s.cfg.LLM.DefaultProvider
|
||
}
|
||
|
||
resolved, err := dao.LoadActiveLLMConfig(fallback)
|
||
if err != nil {
|
||
// 仅打日志,继续走老路径(不要因 DB 抖动中断业务)
|
||
log.Printf("[Enhancer] ⚠️ DB 解析生效配置失败,回落 yaml 路由: %v", err)
|
||
} else {
|
||
// 如果 PHP 强制指定了 provider,但 DB 解析出的 provider 与之不同,
|
||
// 说明 PHP 显式覆盖——这种情况下以 PHP 指定为准,
|
||
// 但 model/api_key 仍用 DB 的(保证 key 轮换即时生效)
|
||
finalProvider := resolved.Provider
|
||
if forceProvider != "" && forceProvider != resolved.Provider {
|
||
log.Printf("[Enhancer] PHP 强制 provider=%s 覆盖 DB provider=%s",
|
||
forceProvider, resolved.Provider)
|
||
finalProvider = forceProvider
|
||
}
|
||
|
||
// 组装完整 LLMConfigEx
|
||
cfg := &config.LLMConfigEx{
|
||
Provider: finalProvider,
|
||
APIKey: resolved.APIKey,
|
||
BaseURL: resolved.APIURL,
|
||
Model: resolved.Model,
|
||
Timeout: s.getTimeoutForProvider(finalProvider),
|
||
}
|
||
|
||
// 按配置创建/复用 client(指纹变化时自动重建)
|
||
client, cErr := s.llmRouter.GetByConfig(finalProvider, cfg)
|
||
if cErr != nil {
|
||
log.Printf("[Enhancer] ⚠️ 按 config 创建 %s 失败,回落 yaml 路由: %v",
|
||
finalProvider, cErr)
|
||
} else {
|
||
// api_key_id / source 通过 meta 返回,供 Enhance 主流程写 step 用
|
||
return client, finalProvider, resolveMeta{APIKeyID: resolved.APIKeyID, CfgSource: resolved.Source}, nil
|
||
}
|
||
}
|
||
}
|
||
|
||
// ===== 路径 2:DB 不可用时的兜底(保留原有 yaml 路由逻辑) =====
|
||
//
|
||
// 仅在 DB 连接失败、LoadActiveLLMConfig 报错、GetByConfig 失败时进入。
|
||
// 保证 LLM 调用链不因 DB 抖动完全中断。
|
||
|
||
// 1) PHP 强制指定 provider
|
||
if forceProvider != "" {
|
||
c, err := s.llmRouter.GetByProvider(forceProvider)
|
||
if err != nil {
|
||
return nil, "", resolveMeta{}, fmt.Errorf("强制 provider=%s 取模型失败: %w", forceProvider, err)
|
||
}
|
||
return c, forceProvider, resolveMeta{CfgSource: "yaml_force"}, nil
|
||
}
|
||
|
||
// 2) 按 scene 路由(路由表里 scene 名约定如 "medical_record"/"prescription")
|
||
//
|
||
// 注意:ModelRouter.Get(scene) 在 scene 未注册时会**内部回落 defaultProvider**
|
||
// 并成功返回 client,但同行的 ResolveProvider(scene) 只查 routes[scene],
|
||
// 未注册时返回空串。如果不补这个空串判断,会让日志/响应里出现 provider=空。
|
||
if scene != "" {
|
||
if c, err := s.llmRouter.Get(scene); err == nil {
|
||
p := s.llmRouter.ResolveProvider(scene)
|
||
if p == "" {
|
||
// scene 未注册但 Get 内部已回落 default,这里对齐 provider 字符串
|
||
p = s.cfg.LLM.DefaultProvider
|
||
if p == "" {
|
||
p = "deepseek"
|
||
}
|
||
log.Printf("[Enhancer] scene=%s 未注册路由,已回落默认 provider=%s", scene, p)
|
||
}
|
||
return c, p, resolveMeta{CfgSource: "yaml_route"}, nil
|
||
}
|
||
}
|
||
|
||
// 3) 默认 provider
|
||
defProvider := s.cfg.LLM.DefaultProvider
|
||
if defProvider == "" {
|
||
defProvider = "deepseek"
|
||
}
|
||
c, err := s.llmRouter.GetByProvider(defProvider)
|
||
if err != nil {
|
||
return nil, "", resolveMeta{}, fmt.Errorf("默认 provider=%s 取模型失败: %w", defProvider, err)
|
||
}
|
||
return c, defProvider, resolveMeta{CfgSource: "yaml_default"}, nil
|
||
}
|
||
|
||
// getTimeoutForProvider 取指定 provider 的 HTTP 超时(秒)
|
||
//
|
||
// DB 表 xk_ai_platform 没有超时字段,从 yaml/env 配置中查对应 provider 的 timeout。
|
||
// 如果 yaml 里没有该 provider,回落到 120 秒(与 config.go 默认值一致)。
|
||
func (s *EnhancerService) getTimeoutForProvider(provider string) int {
|
||
if s.cfg == nil {
|
||
return 120
|
||
}
|
||
if cfg, ok := s.cfg.LLM.Models[provider]; ok && cfg.Timeout > 0 {
|
||
return cfg.Timeout
|
||
}
|
||
return 120
|
||
}
|
||
|
||
// kbRetrievalLimits 注入长度预算(rune 计)
|
||
//
|
||
// 为什么要限:TopK*整段 chunk + 方剂可能过万字,直接塞 system 会稀释业务指令
|
||
// 且逼近轻量模型上下文上限;预算按"药材知识为主、方剂为辅"分配
|
||
const (
|
||
kbDocMaxRunes = 500 // 单条药材知识注入上限
|
||
kbDocsTotalRunes = 2600 // 药材知识总预算
|
||
kbFormulaMaxRunes = 400 // 单条参考方剂注入上限
|
||
kbFormulaTopK = 3 // 方剂检索条数(方剂是"参考骨架",2~3 个足够化裁)
|
||
kbMinScoreRatio = 0.15 // 相对分数阈值:低于最高分 15% 的命中视为弱相关噪声
|
||
)
|
||
|
||
// performKBRetrieval 执行知识库检索(按 agentcfg.KB.Source 分流)
|
||
//
|
||
// 返回值:
|
||
// - step:检索过程步骤(写入 EnhanceResponse.Steps,PHP 据此写子表)
|
||
// - docs:药材/文档知识块(多库归并 + 阈值过滤 + title 加权后格式化)
|
||
// - formulas:参考方剂块(仅处方场景,金方库 xk_golden_formula 动态检索)
|
||
// - err:检索失败原因(仅用于调用方决定是否阻断;step 内已经记录)
|
||
//
|
||
// 分流策略:
|
||
// 1. agentcfg.KB.Source == "local" → 走 LocalKB(V1 默认,多库归并)
|
||
// 2. agentcfg.KB.Source == "maxkb" → 走 MaxKB(需 maxkb 客户端可用)
|
||
// 3. 选定路径失败时不阻断主流程,调用方仍用原始 messages 调 LLM
|
||
// 4. 金方检索独立于 KB source(直查 z_xk 库),失败只降级不影响文档检索
|
||
func (s *EnhancerService) performKBRetrieval(ctx context.Context, req *EnhanceRequest) (*EnhanceStep, []string, []string, error) {
|
||
kbCfg := agentcfg.GetKBConfig()
|
||
topK := req.TopK
|
||
if topK <= 0 {
|
||
topK = kbCfg.TopK
|
||
if topK <= 0 {
|
||
topK = 5
|
||
}
|
||
}
|
||
|
||
step := &EnhanceStep{
|
||
StepType: "kb_retrieval",
|
||
StartedAt: time.Now().Unix(),
|
||
}
|
||
start := time.Now()
|
||
|
||
// 别名归一扩展(P1):查询中的药材别名追加正名,提高召回
|
||
// 例:主诉里写"淮山",扩展后能命中"山药"的性味/配伍语料
|
||
query, aliasHits := kb.GetAliasIndex().ExpandQuery(req.Context)
|
||
detailParts := []string{}
|
||
if len(aliasHits) > 0 {
|
||
detailParts = append(detailParts, "别名扩展→"+strings.Join(aliasHits, "/"))
|
||
}
|
||
|
||
var docs []string
|
||
var docsDetail string
|
||
var err error
|
||
source := kbCfg.Source
|
||
|
||
switch source {
|
||
case "local":
|
||
docs, docsDetail, err = s.retrieveFromLocalMulti(ctx, query, topK, kbCfg.SearchMode)
|
||
case "maxkb":
|
||
docs, err = s.retrieveFromMaxKB(ctx, query, topK)
|
||
docsDetail = fmt.Sprintf("命中 %d 条", len(docs))
|
||
default:
|
||
// 未知 source 回落到本地
|
||
source = "local(fallback)"
|
||
docs, docsDetail, err = s.retrieveFromLocalMulti(ctx, query, topK, kbCfg.SearchMode)
|
||
}
|
||
|
||
// 金方库检索(P1):处方场景第二检索源,动态读取平台维护的方剂
|
||
// 与药材知识独立:金方失败不影响文档结果(只补一条 detail 说明)
|
||
var formulas []string
|
||
if req.Scene == "prescription" {
|
||
var ferr error
|
||
formulas, ferr = retrieveGoldenFormulas(query, kbFormulaTopK)
|
||
if ferr != nil {
|
||
log.Printf("[Enhancer] 金方检索失败(不阻断): %v", ferr)
|
||
detailParts = append(detailParts, "方剂检索失败:"+ferr.Error())
|
||
} else if len(formulas) > 0 {
|
||
detailParts = append(detailParts, fmt.Sprintf("参考方剂 %d 条", len(formulas)))
|
||
}
|
||
}
|
||
|
||
step.DurationMs = int(time.Since(start).Milliseconds())
|
||
step.FinishedAt = time.Now().Unix()
|
||
|
||
if err != nil {
|
||
step.Status = 2
|
||
step.Detail = fmt.Sprintf("[%s] 检索失败: %s", source, err.Error())
|
||
log.Printf("[Enhancer] [%s] 检索失败,降级用原始消息: %v", source, err)
|
||
// 文档检索失败但方剂命中时仍可注入方剂(调用方按 err==nil 判断,这里直接吞掉
|
||
// 文档错误会丢失告警信息——保持原有"失败降级"语义,方剂一并放弃)
|
||
return step, nil, nil, err
|
||
}
|
||
|
||
step.Status = 1
|
||
detail := fmt.Sprintf("[%s] %s", source, docsDetail)
|
||
if len(detailParts) > 0 {
|
||
detail += "|" + strings.Join(detailParts, "|")
|
||
}
|
||
step.Detail = detail
|
||
return step, docs, formulas, nil
|
||
}
|
||
|
||
// localKBHit 本地多库归并检索的中间结果(按 doc 去重、加权后排序用)
|
||
type localKBHit struct {
|
||
docID uint
|
||
title string
|
||
content string
|
||
score float64
|
||
libName string
|
||
}
|
||
|
||
// retrieveFromLocalMulti 本地知识库多库归并检索(P1 主路径)
|
||
//
|
||
// 与旧版(只查第一个库)的差异:
|
||
// 1. 遍历全部 status=1 的库分别检索,按分数归并
|
||
// 2. title 命中加权:查询词包含文档标题(如药名/方名直接出现在证候上下文里)
|
||
// 的命中 ×1.3——标题级命中的相关性远高于正文碰撞
|
||
// 3. 相对分数阈值:低于最高分 15% 的命中丢弃(过滤 ngram 2-gram 碰撞噪声)
|
||
// 4. 按 doc 去重:同一文档多个分段命中时只保留最高分那段,避免同一味药占满 TopK
|
||
//
|
||
// 返回:格式化好的知识块列表 + 命中摘要(step detail 用)
|
||
func (s *EnhancerService) retrieveFromLocalMulti(ctx context.Context, query string, topK int, mode string) ([]string, string, error) {
|
||
if s.localSearch == nil {
|
||
return nil, "", fmt.Errorf("本地知识库检索器未初始化")
|
||
}
|
||
libs, err := listActiveLibraries()
|
||
if err != nil {
|
||
return nil, "", err
|
||
}
|
||
if len(libs) == 0 {
|
||
return nil, "", fmt.Errorf("本地无可用知识库(请先在 /kb/view 导入文档)")
|
||
}
|
||
// 逐库检索(库数量少,串行足够;每库拿 topK 再全局归并)
|
||
all := make([]localKBHit, 0, topK*len(libs))
|
||
libHitDesc := make([]string, 0, len(libs))
|
||
for _, lib := range libs {
|
||
results, serr := s.localSearch.Search(ctx, kb.SearchOptions{
|
||
LibraryID: lib.ID,
|
||
Query: query,
|
||
TopK: topK,
|
||
Mode: mode,
|
||
})
|
||
if serr != nil {
|
||
// 单库失败不整体报错(可能只是某库空/索引缺失),记录后继续
|
||
log.Printf("[Enhancer] 库[%s]检索失败(跳过): %v", lib.Name, serr)
|
||
continue
|
||
}
|
||
if len(results) > 0 {
|
||
libHitDesc = append(libHitDesc, fmt.Sprintf("%s:%d", lib.Name, len(results)))
|
||
}
|
||
for _, r := range results {
|
||
score := r.Score
|
||
title := strings.TrimSpace(r.Title)
|
||
// title 加权:查询上下文直接点名该文档(药名/方名)→ 相关性强提升
|
||
if title != "" && strings.Contains(query, title) {
|
||
score *= 1.3
|
||
}
|
||
all = append(all, localKBHit{
|
||
docID: r.DocID,
|
||
title: title,
|
||
content: strings.TrimSpace(r.Content),
|
||
score: score,
|
||
libName: lib.Name,
|
||
})
|
||
}
|
||
}
|
||
if len(all) == 0 {
|
||
return []string{}, "命中 0 条", nil
|
||
}
|
||
// 按 doc 去重:保留每个文档的最高分分段
|
||
bestByDoc := make(map[uint]localKBHit, len(all))
|
||
for _, h := range all {
|
||
if cur, ok := bestByDoc[h.docID]; !ok || h.score > cur.score {
|
||
bestByDoc[h.docID] = h
|
||
}
|
||
}
|
||
merged := make([]localKBHit, 0, len(bestByDoc))
|
||
for _, h := range bestByDoc {
|
||
merged = append(merged, h)
|
||
}
|
||
sort.Slice(merged, func(i, j int) bool { return merged[i].score > merged[j].score })
|
||
// 相对分数阈值:过滤远低于最高分的弱相关噪声
|
||
minScore := merged[0].score * kbMinScoreRatio
|
||
filtered := merged[:0]
|
||
for _, h := range merged {
|
||
if h.score >= minScore {
|
||
filtered = append(filtered, h)
|
||
}
|
||
}
|
||
if len(filtered) > topK {
|
||
filtered = filtered[:topK]
|
||
}
|
||
// 格式化 + 总长度预算控制
|
||
out := make([]string, 0, len(filtered))
|
||
usedRunes := 0
|
||
for _, h := range filtered {
|
||
content := truncateRunesStr(h.content, kbDocMaxRunes)
|
||
var block string
|
||
if h.title == "" {
|
||
block = content
|
||
} else {
|
||
// 带来源标题与库名:LLM 引用时可指明出处,排障时可定位语料
|
||
block = fmt.Sprintf("【%s|来源:%s】\n%s", h.title, h.libName, content)
|
||
}
|
||
blockRunes := len([]rune(block))
|
||
if usedRunes+blockRunes > kbDocsTotalRunes && len(out) > 0 {
|
||
break // 预算用尽:保留已有条目(至少注入 1 条)
|
||
}
|
||
usedRunes += blockRunes
|
||
out = append(out, block)
|
||
}
|
||
desc := fmt.Sprintf("命中 %d 条(%s)", len(out), strings.Join(libHitDesc, ", "))
|
||
return out, desc, nil
|
||
}
|
||
|
||
// retrieveGoldenFormulas 金方库检索(处方场景专用)
|
||
//
|
||
// 两段式与 KB 检索一致:BOOLEAN 精确 → 零命中回退 NATURAL LANGUAGE。
|
||
// 输出格式固定三行:方名 / 组成(解析 drugs_json)/ 主治(优先译文),
|
||
// 提示词侧告知 LLM"可在参考方基础上化裁",不强制照搬
|
||
func retrieveGoldenFormulas(query string, topK int) ([]string, error) {
|
||
hits, err := dao.GoldenFormulaSearch(query, topK, false)
|
||
if (err != nil || len(hits) == 0) && len([]rune(query)) >= 3 {
|
||
hits, err = dao.GoldenFormulaSearch(query, topK, true)
|
||
}
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
out := make([]string, 0, len(hits))
|
||
for _, h := range hits {
|
||
drugs := formatFormulaDrugs(h.DrugsJSON)
|
||
if drugs == "" {
|
||
// drugs_json 解析不出组成时回退药材速览(herb_overview 是人读格式)
|
||
drugs = strings.TrimSpace(h.HerbOverview)
|
||
}
|
||
indication := strings.TrimSpace(h.IndicationTranslation)
|
||
if indication == "" {
|
||
indication = strings.TrimSpace(h.IndicationOriginal)
|
||
}
|
||
block := fmt.Sprintf("〔%s〕组成:%s。主治:%s", h.Name, drugs, indication)
|
||
out = append(out, truncateRunesStr(block, kbFormulaMaxRunes))
|
||
}
|
||
return out, nil
|
||
}
|
||
|
||
// formatFormulaDrugs 把金方 drugs_json 解析成"桂枝46.9g、芍药46.9g"形式
|
||
//
|
||
// drugs_json 元素结构(金方后台维护):{name, dose(克数), unit, usage, prep, ancient_dose}
|
||
// 解析失败返回空串,调用方回退 herb_overview
|
||
func formatFormulaDrugs(raw *string) string {
|
||
if raw == nil || *raw == "" {
|
||
return ""
|
||
}
|
||
var items []struct {
|
||
Name string `json:"name"`
|
||
Dose float64 `json:"dose"`
|
||
Unit string `json:"unit"`
|
||
Usage string `json:"usage"`
|
||
}
|
||
if err := json.Unmarshal([]byte(*raw), &items); err != nil {
|
||
return ""
|
||
}
|
||
parts := make([]string, 0, len(items))
|
||
for _, it := range items {
|
||
name := strings.TrimSpace(it.Name)
|
||
if name == "" {
|
||
continue
|
||
}
|
||
p := name
|
||
if it.Dose > 0 {
|
||
unit := it.Unit
|
||
if unit == "" {
|
||
unit = "g"
|
||
}
|
||
// 去掉多余小数(46.88 → 46.88,36 → 36)
|
||
p += strconv.FormatFloat(it.Dose, 'f', -1, 64) + unit
|
||
}
|
||
if it.Usage != "" {
|
||
p += "(" + it.Usage + ")"
|
||
}
|
||
parts = append(parts, p)
|
||
}
|
||
return strings.Join(parts, "、")
|
||
}
|
||
|
||
// truncateRunesStr 按 rune 截断(不破坏 UTF-8 边界),超长加省略号
|
||
func truncateRunesStr(s string, max int) string {
|
||
runes := []rune(s)
|
||
if max <= 0 || len(runes) <= max {
|
||
return s
|
||
}
|
||
return string(runes[:max]) + "…"
|
||
}
|
||
|
||
// retrieveFromMaxKB 从 MaxKB 检索(兼容旧路径,需要 MaxKB Pro)
|
||
func (s *EnhancerService) retrieveFromMaxKB(ctx context.Context, query string, topK int) ([]string, error) {
|
||
if s.maxkb == nil {
|
||
return nil, fmt.Errorf("MaxKB 客户端未初始化")
|
||
}
|
||
return s.maxkb.Search(ctx, query, topK)
|
||
}
|
||
|
||
// listActiveLibraries 取所有启用的本地知识库(按 id 升序)
|
||
//
|
||
// 单独抽出来避免在 retrieveFromLocal 里直接 import dao(service 包不直接依赖 dao)
|
||
// 但这里仍要调 dao.KBListLibraries —— service 包对 dao 的依赖在 enhancer 之外已有先例
|
||
// (agentcfg 也依赖 dao),保持一致
|
||
func listActiveLibraries() ([]libraryView, error) {
|
||
rows, err := listActiveLibrariesFromDAO()
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
out := make([]libraryView, 0, len(rows))
|
||
for _, r := range rows {
|
||
out = append(out, libraryView{ID: r.ID, Name: r.Name})
|
||
}
|
||
return out, nil
|
||
}
|
||
|
||
//
|
||
// 插入位置策略:
|
||
// - 若 messages[0] 是 system → 在其后插入(保持业务 system 优先级最高)
|
||
// - 若不是 → 在最前面插入
|
||
//
|
||
// 这样设计的理由:业务 system prompt 是"医生角色/输出格式"等核心指令,
|
||
// 不能被 KB 内容稀释;KB 内容是"参考资料",定位次于业务 system。
|
||
func injectKBContext(messages []types.Message, docs []string, formulas []string) []types.Message {
|
||
if len(docs) == 0 && len(formulas) == 0 {
|
||
return messages
|
||
}
|
||
|
||
// 分块注入:药材/文档知识块 + 参考方剂块(P1)
|
||
// 方剂块单独成段并明确"可化裁"——避免 LLM 把参考方当唯一答案照抄剂量
|
||
var sb strings.Builder
|
||
if len(docs) > 0 {
|
||
sb.WriteString("以下是知识库检索到的参考资料,请在生成时参考:\n\n")
|
||
for i, d := range docs {
|
||
sb.WriteString(fmt.Sprintf("【参考 %d】\n%s\n\n", i+1, strings.TrimSpace(d)))
|
||
}
|
||
}
|
||
if len(formulas) > 0 {
|
||
sb.WriteString("以下是金方库检索到的参考方剂(与患者证候相关的经典方)。")
|
||
sb.WriteString("可在参考方基础上按患者实际辨证化裁加减,不必照搬原方与剂量;")
|
||
sb.WriteString("若证不对应则不要用:\n\n")
|
||
for i, f := range formulas {
|
||
sb.WriteString(fmt.Sprintf("【参考方剂 %d】%s\n\n", i+1, strings.TrimSpace(f)))
|
||
}
|
||
}
|
||
sb.WriteString("(参考资料结束,请基于以上资料和你的医学知识完成后续任务)")
|
||
kbText := sb.String()
|
||
|
||
// ★ 注入方式必须是"合并进首条 system"而不是"追加第二条 system":
|
||
// 讯飞 Spark 等厂商强制 system 只能作为第一条消息,出现第二条 system
|
||
// 会直接报 HTTP 500 NotFirstSystemError(code=10049)(P0 打通 context 后实测踩坑)。
|
||
// 合并进首条对 deepseek/openai 也完全兼容,行为一致性最好
|
||
out := make([]types.Message, len(messages))
|
||
copy(out, messages)
|
||
if len(out) > 0 && out[0].Role == "system" {
|
||
// 业务 system 指令在前,KB 参考资料附在其后(保持业务指令优先级最高)
|
||
out[0].Content = out[0].Content + "\n\n" + kbText
|
||
return out
|
||
}
|
||
// 没有业务 system 时,KB 参考作为唯一 system 放最前
|
||
kbMsg := types.Message{
|
||
Role: "system",
|
||
Content: kbText,
|
||
Timestamp: time.Now().Unix(),
|
||
}
|
||
return append([]types.Message{kbMsg}, out...)
|
||
}
|