From 26cb3253dc48f4dfae1b4dc1a7ecc2a197a757d1 Mon Sep 17 00:00:00 2001 From: st <2663600842@qq.com> Date: Thu, 6 Aug 2026 18:03:43 +0800 Subject: [PATCH 01/10] =?UTF-8?q?feat(graph):=20Rewrite=20=E8=8A=82?= =?UTF-8?q?=E7=82=B9=E5=8D=87=E7=BA=A7=E4=B8=BA=20LLM=20=E6=9F=A5=E8=AF=A2?= =?UTF-8?q?=E6=94=B9=E5=86=99=20+=20=E6=84=8F=E5=9B=BE=E8=AF=86=E5=88=AB?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - quickRewriteFn 从占位符改为真正调 cm.Generate 做改写 - 新增 5 类意图: greeting/chitchat/question/identity/meta - greeting/chitchat 设置 SkipRetrieve=true 跳过知识库 - BuildMsgs 用 RewrittenQuery 替换用户消息内容 - eino_adapter Retrieve 空 query 快速返回 nil - 全链路降级: ChatModel 缺失/Generate 失败/JSON 解析失败 fallback 原始 query --- internal/rag/eino_adapter.go | 3 + internal/service/chat_service_graph_quick.go | 194 ++++++++++++++++++- 2 files changed, 187 insertions(+), 10 deletions(-) diff --git a/internal/rag/eino_adapter.go b/internal/rag/eino_adapter.go index dea2c1e..19c3ece 100644 --- a/internal/rag/eino_adapter.go +++ b/internal/rag/eino_adapter.go @@ -97,6 +97,9 @@ func (a *EinoRetrieverAdapter) Retrieve(ctx context.Context, query string, opts if a == nil || a.inner == nil { return nil, fmt.Errorf("eino retriever adapter: inner retriever is nil") } + if strings.TrimSpace(query) == "" { + return nil, nil + } defaultTopK := a.defaultTopK common := retriever.GetCommonOptions(&retriever.Options{ diff --git a/internal/service/chat_service_graph_quick.go b/internal/service/chat_service_graph_quick.go index 6e72cc9..9f1d319 100644 --- a/internal/service/chat_service_graph_quick.go +++ b/internal/service/chat_service_graph_quick.go @@ -2,6 +2,7 @@ package service import ( "context" + "encoding/json" "errors" "fmt" "io" @@ -47,10 +48,52 @@ type quickGraphInput struct { // quickGraphState Graph Local State,通过 ProcessState 读写。 type quickGraphState struct { Input *quickGraphInput - RewrittenQuery string + RewrittenQuery string // 改写后的查询,供 Retrieve / BuildMsgs 使用 + Intent string // greeting / chitchat / question / identity / meta + SkipRetrieve bool // Greeting/Chitchat 跳过知识库检索 + Keywords []string // 改写时提取的关键词,可用于日志/调试 RetrievedDocs []*schema.Document } +// 查询改写意图类型 +const ( + intentGreeting = "greeting" // 问候语 + intentChitchat = "chitchat" // 闲聊 + intentQuestion = "question" // 知识查询(默认,最常见) + intentIdentity = "identity" // 身份确认(你是谁、你能做什么) + intentMeta = "meta" // 元问题(我的历史记录、你刚才说了什么) +) + +// rewriteResult LLM 返回的 JSON 解析结果 +type rewriteResult struct { + Rewritten string `json:"rewritten"` + Intent string `json:"intent"` + Keywords []string `json:"keywords"` +} + +// rewriteMaxHistoryRounds 改写时拼入历史的最大轮数(每轮=user+assistant) +const rewriteMaxHistoryRounds = 3 + +// rewriteSystemPrompt 改写专用的 System Prompt +const rewriteSystemPrompt = `你是一个查询改写助手。根据用户的原始问题和对话历史,对问题进行改写并识别意图。 + +## 改写规则 +1. 消解指代:把"这个"、"那个方案"、"它"、"之前说的"等代词替换为对话历史中的具体名词 +2. 扩展关键词:补充与问题相关的同义词、上下位词,方便知识库检索 +3. 拆分复合问题:如果原问题包含多个子问题,改写为一个完整句子即可(不要拆成多行) +4. 保持原意:改写后的问题必须和原问题核心意图一致,不要引入新主题 + +## 意图识别 +- greeting: 问候语(你好、hi、在吗、早上好) +- chitchat: 闲聊(今天天气怎么样、讲个笑话、随便聊聊) +- question: 知识查询(业务问题、技术问题、需要从知识库找答案) +- identity: 身份确认(你是谁、你能做什么、介绍一下你自己) +- meta: 元问题(我的历史记录、你刚才说了什么、回顾对话) + +## 输出格式 +严格使用 JSON,不要输出任何多余文字或 Markdown 代码块: +{"rewritten": "改写后的完整问题", "intent": "question", "keywords": ["关键词1", "关键词2"]}` + const ( graphQuickNodeRewrite = "query_rewrite" graphQuickNodeRetrieve = "retrieve" @@ -94,7 +137,7 @@ func wrapGraphErr(stage string, err error) error { return apperrors.WrapDefault(apperrors.CodeInternalError, fmt.Errorf("%s: %w", stage, err)) } -// addQuickRewriteNode 节点 1:QueryRewrite(占位,返回原文)。 +// addQuickRewriteNode 节点 1:QueryRewrite,调 LLM 做查询改写 + 意图识别。 func addQuickRewriteNode(g *einoCompose.Graph[*quickGraphInput, *schema.StreamReader[*schema.Message]]) error { return g.AddLambdaNode(graphQuickNodeRewrite, einoCompose.InvokableLambda(quickRewriteFn), @@ -102,7 +145,9 @@ func addQuickRewriteNode(g *einoCompose.Graph[*quickGraphInput, *schema.StreamRe ) } -// quickRewriteFn 节点 1 实现:暂存输入到 State 并原样返回查询 +// quickRewriteFn 节点 1 实现:调 LLM 对用户问题做改写 + 意图识别。 +// 降级策略:LLM 改写失败 → fallback 原始 query,不阻塞主流程。 +// SkipRetrieve=true 时返回空串,Retriever 收到空串会快速返回空 docs。 func quickRewriteFn(ctx context.Context, input *quickGraphInput) (string, error) { if input == nil { return "", apperrors.NewDefault(apperrors.CodeInvalidParam) @@ -113,22 +158,135 @@ func quickRewriteFn(ctx context.Context, input *quickGraphInput) (string, error) }); err != nil { return "", err } + startAt := time.Now() - rewritten := input.OriginalQuery + rewritten, intent, keywords, skipRetrieve := doRewriteWithLLM(ctx, input) + _ = einoCompose.ProcessState(ctx, func(_ context.Context, state *quickGraphState) error { state.RewrittenQuery = rewritten + state.Intent = intent + state.Keywords = keywords + state.SkipRetrieve = skipRetrieve return nil }) + durMs := time.Since(startAt).Milliseconds() observability.SetSpanAttrs(ctx, observability.Attrs{ - "original_query": input.OriginalQuery, - "rewritten_query": rewritten, - "rewrite_ms": durMs, - "model_id": input.ModelName, + "original_query": input.OriginalQuery, + "rewritten_query": rewritten, + "intent": intent, + "skip_retrieve": fmt.Sprintf("%v", skipRetrieve), + "rewrite_ms": durMs, + "model_id": input.ModelName, }) return rewritten, nil } +// doRewriteWithLLM 调 LLM 做改写,失败时 fallback 原始 query。 +// 返回 (rewritten, intent, keywords, skipRetrieve) +func doRewriteWithLLM(ctx context.Context, input *quickGraphInput) (string, string, []string, bool) { + // 1. 从 context 拿 ChatModel + cm, ok := graphChatModelFromContext(ctx) + if !ok || cm == nil { + logger.Warnf("quickRewriteFn: context 中没有 ChatModel,跳过改写") + return input.OriginalQuery, intentQuestion, nil, false + } + + // 2. 从 InputMsgs 提取最近几轮用户-助手历史(排除 system 和当前问题) + historyStr := buildRewriteHistory(input.InputMsgs, input.UserQuestionIndex, rewriteMaxHistoryRounds) + + // 3. 构造改写请求消息 + var userContent strings.Builder + userContent.WriteString("原始问题:") + userContent.WriteString(input.OriginalQuery) + if historyStr != "" { + userContent.WriteString("\n\n对话历史:\n") + userContent.WriteString(historyStr) + } + + msgs := []*schema.Message{ + schema.SystemMessage(rewriteSystemPrompt), + schema.UserMessage(userContent.String()), + } + + // 4. 同步调 Generate(改写不需要流式) + msg, err := cm.Generate(ctx, msgs) + if err != nil || msg == nil || msg.Content == "" { + logger.Warnf("quickRewriteFn: LLM 改写失败,fallback 原始 query: err=%v", err) + return input.OriginalQuery, intentQuestion, nil, false + } + + // 5. 解析 JSON 返回 + var result rewriteResult + content := strings.TrimSpace(msg.Content) + content = strings.TrimPrefix(content, "```json") + content = strings.TrimPrefix(content, "```") + content = strings.TrimSuffix(content, "```") + content = strings.TrimSpace(content) + + if err := json.Unmarshal([]byte(content), &result); err != nil { + logger.Warnf("quickRewriteFn: LLM 改写返回 JSON 解析失败,fallback 原始 query: err=%v, content=%s", err, content) + return input.OriginalQuery, intentQuestion, nil, false + } + + // 6. 清洗 + 验证 + if strings.TrimSpace(result.Rewritten) == "" { + result.Rewritten = input.OriginalQuery + } + if !isValidIntent(result.Intent) { + result.Intent = intentQuestion + } + + // 7. 判定是否跳过检索(greeting/chitchat 不需要知识库) + skipRetrieve := result.Intent == intentGreeting || result.Intent == intentChitchat + + return result.Rewritten, result.Intent, result.Keywords, skipRetrieve +} + +// isValidIntent 检查 LLM 返回的意图是否在合法枚举内 +func isValidIntent(intent string) bool { + switch intent { + case intentGreeting, intentChitchat, intentQuestion, intentIdentity, intentMeta: + return true + } + return false +} + +// buildRewriteHistory 从 InputMsgs 提取最近 N 轮 user-assistant 对话(排除 system 和当前问题)。 +// maxRounds 控制最大轮数,避免改写 prompt 太长。 +func buildRewriteHistory(msgs []*schema.Message, currentUserMsgIdx, maxRounds int) string { + if len(msgs) == 0 { + return "" + } + // 从 currentUserMsgIdx 往前找,跳过 system,收集 user+assistant 对 + // 简化实现:找最近 maxRounds*2 条非 system 消息,倒序输出 + var pairs []string + for i := currentUserMsgIdx - 1; i >= 0 && len(pairs) < maxRounds*2; i-- { + m := msgs[i] + if m == nil || m.Content == "" { + continue + } + role := string(m.Role) + if role == "system" { + continue + } + // user 和 assistant 交替收集,用最近的优先 + roleLabel := "用户" + if role == "assistant" { + roleLabel = "助手" + } + pairs = append([]string{fmt.Sprintf("%s:%s", roleLabel, m.Content)}, pairs...) + } + if len(pairs) == 0 { + return "" + } + // 只取 maxRounds*2 条(即 maxRounds 轮) + if len(pairs) > maxRounds*2 { + pairs = pairs[len(pairs)-maxRounds*2:] + } + return strings.Join(pairs, "\n") +} + // addQuickRetrieveNode 节点 2:Retrieve,PostHandler 把结果写回 State。 func addQuickRetrieveNode(g *einoCompose.Graph[*quickGraphInput, *schema.StreamReader[*schema.Message]], einoRetriever *rag.EinoRetrieverAdapter) error { return g.AddRetrieverNode(graphQuickNodeRetrieve, einoRetriever, @@ -151,13 +309,24 @@ func addQuickBuildMsgsNode(g *einoCompose.Graph[*quickGraphInput, *schema.Stream // quickBuildMsgsFn 节点 3 实现:在用户问题前插入检索上下文块 func quickBuildMsgsFn(ctx context.Context, docs []*schema.Document) ([]*schema.Message, error) { - var input *quickGraphInput + var ( + input *quickGraphInput + rewrittenQuery string + ) if err := einoCompose.ProcessState(ctx, func(_ context.Context, state *quickGraphState) error { input = state.Input + rewrittenQuery = state.RewrittenQuery return nil }); err != nil || input == nil { return nil, apperrors.NewDefault(apperrors.CodeInternalError) } + + // 用 RewrittenQuery 替换用户问题(如果有改写结果) + questionContent := input.OriginalQuery + if rewrittenQuery != "" && rewrittenQuery != input.OriginalQuery { + questionContent = rewrittenQuery + } + msgs := make([]*schema.Message, 0, len(input.InputMsgs)+2) injected := false for i, m := range input.InputMsgs { @@ -166,7 +335,12 @@ func quickBuildMsgsFn(ctx context.Context, docs []*schema.Document) ([]*schema.M msgs = append(msgs, schema.UserMessage(block)) injected = true } - msgs = append(msgs, m) + // 替换 UserQuestion 位置的内容为改写后的 query + if i == input.UserQuestionIndex { + msgs = append(msgs, schema.UserMessage(questionContent)) + } else { + msgs = append(msgs, m) + } } if len(docs) > 0 && !injected { last := msgs[len(msgs)-1] From 4d58ef3336c9f1de10b4e0fa3ec6e6804becad98 Mon Sep 17 00:00:00 2001 From: st <2663600842@qq.com> Date: Thu, 6 Aug 2026 18:03:50 +0800 Subject: [PATCH 02/10] =?UTF-8?q?fix(prompt):=20=E6=B7=B1=E5=BA=A6?= =?UTF-8?q?=E6=A8=A1=E5=BC=8F=20UserCtx=20=E6=89=A9=E5=B1=95=2012=20?= =?UTF-8?q?=E5=AD=97=E6=AE=B5=20+=20=E5=88=A0=E9=99=A4=E6=AD=BB=E4=BB=A3?= =?UTF-8?q?=E7=A0=81?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - chat_prompt_builder UserCtx 从 4 字段扩展到 12 字段 - 新增 toAgentPromptUserContext 辅助函数做字段映射 - 删除 execute.go 中冗余的 buildEnhancedSystemPromptForAgent - PromptBuilder.BuildSystem 统一作为唯一 Prompt 构造入口 --- internal/agent/execute.go | 116 ++---------------------- internal/service/chat_prompt_builder.go | 33 +++++-- 2 files changed, 30 insertions(+), 119 deletions(-) diff --git a/internal/agent/execute.go b/internal/agent/execute.go index 4f7329a..617344c 100644 --- a/internal/agent/execute.go +++ b/internal/agent/execute.go @@ -126,9 +126,13 @@ func (e *Engine) runAgent(ctx context.Context, req Request, chatModel model.Tool } var systemPromptFinal string if req.SystemPrompt != "" { - systemPromptFinal = baseSystemPrompt + "\n\n" + req.SystemPrompt + // BuildSystem() 以空 baseSystem 构建时,结果会前导 "\n\n",去掉避免多余空行 + enhanced := strings.TrimLeft(req.SystemPrompt, "\n") + systemPromptFinal = baseSystemPrompt + "\n\n" + enhanced } else { - systemPromptFinal = buildEnhancedSystemPromptForAgent(baseSystemPrompt, req.Summary, req.Memories, req.UserCtx) + // 兜底:req.SystemPrompt 为空时只用 ReAct 规则(正常流程不会走到这里, + // PromptBuilder.BuildSystem() 总会产出摘要/记忆/用户信息之一) + systemPromptFinal = baseSystemPrompt } logger.Infof("[Agent] SystemPrompt (前400字符): %s", truncateStr(systemPromptFinal, 400)) inputMessages := buildInputMessages(req.Query, req.History) @@ -414,111 +418,3 @@ func isToolChoiceUnsupportedError(errMsg string) bool { } return false } - -// buildEnhancedSystemPromptForAgent 在 ReAct 系统提示词上注入:时间/用户信息 + 对话摘要 + 用户记忆 -// 与快速模式的 buildEnhancedSystemPrompt 逻辑保持一致,确保双模式行为统一 -func buildEnhancedSystemPromptForAgent(base string, summary *entity.ChatSummary, memories []entity.UserMemory, userCtx PromptUserContext) string { - var extras []string - - userInfo := "## 当前信息\n" - if userCtx.TimeStr != "" { - userInfo += "- 当前时间:" + userCtx.TimeStr + "\n" - } - if userCtx.Timezone != "" { - userInfo += "- 用户时区:" + userCtx.Timezone + "\n" - } - if userCtx.Username != "" { - userInfo += "- 用户:" + userCtx.Username + "\n" - } - if userCtx.Role != "" { - userInfo += "- 系统角色:" + userCtx.Role + "\n" - } - if userCtx.Department != "" { - userInfo += "- 部门:" + userCtx.Department + "\n" - } - if userCtx.Position != "" { - userInfo += "- 职位:" + userCtx.Position + "\n" - } - if userCtx.Expertise != "" { - userInfo += "- 擅长/关注:" + userCtx.Expertise + "\n" - } - if userCtx.Language != "" { - userInfo += "- 偏好语言:" + userCtx.Language + "\n" - } - if userInfo != "## 当前信息\n" { - extras = append(extras, userInfo) - } - - if userCtx.AnswerStyle != "" || userCtx.TableFirst || userCtx.CitationStyle != "" { - var p strings.Builder - p.WriteString("## 用户回答偏好\n") - switch userCtx.AnswerStyle { - case "concise": - p.WriteString("- 回答风格:简洁凝练,直击要点,3~5 句说完,不过度展开\n") - case "detailed": - p.WriteString("- 回答风格:详细展开,先结论再分点论述,必要时给例子和注意事项\n") - case "step_by_step": - p.WriteString("- 回答风格:分步讲解,用 1/2/3…编号或小标题组织步骤\n") - default: - p.WriteString("- 回答风格:平衡简洁与完整,先结论再展开\n") - } - if userCtx.TableFirst { - p.WriteString("- 结构化呈现:对比、列表、映射等数据优先用 Markdown 表格组织\n") - } - switch userCtx.CitationStyle { - case "none": - p.WriteString("- 引用格式:正文不标注引用,引用信息仅由消息底部来源区展示\n") - case "doc_title_only": - p.WriteString("- 引用格式:正文引用时只提「根据《文档名》」,不要章节\n") - default: - p.WriteString("- 引用格式:正文引用时以「根据《文档名》· 章节标题」形式说明来源\n") - } - extras = append(extras, p.String()) - } - - if strings.TrimSpace(userCtx.RoleTemplatePrompt) != "" { - extras = append(extras, "## 角色模板设定\n"+strings.TrimSpace(userCtx.RoleTemplatePrompt)) - } - - if userCtx.Language != "" { - langHint := "## 回答语言\n" - switch userCtx.Language { - case "en-US": - langHint += "- 请使用英文回答(美式英语)。\n" - case "ja-JP": - langHint += "- 请使用日语回答。\n" - case "ko-KR": - langHint += "- 请使用韩语回答。\n" - case "fr-FR": - langHint += "- 请使用法语回答。\n" - case "de-DE": - langHint += "- 请使用德语回答。\n" - case "es-ES": - langHint += "- 请使用西班牙语回答。\n" - default: - langHint += "- 请使用简体中文回答。\n" - } - extras = append(extras, langHint) - } - - if summary != nil && summary.Summary != "" { - extras = append(extras, "## 本次对话摘要\n"+summary.Summary) - } - - if len(memories) > 0 { - var memoryText strings.Builder - memoryText.WriteString("## 关于用户的已知信息\n") - for _, m := range memories { - memoryText.WriteString("- ") - memoryText.WriteString(m.Content) - memoryText.WriteString("\n") - } - extras = append(extras, memoryText.String()) - } - - if len(extras) == 0 { - return base - } - - return base + "\n\n" + strings.Join(extras, "\n\n") -} diff --git a/internal/service/chat_prompt_builder.go b/internal/service/chat_prompt_builder.go index 007d215..c4a1a46 100644 --- a/internal/service/chat_prompt_builder.go +++ b/internal/service/chat_prompt_builder.go @@ -254,8 +254,10 @@ func (b *PromptBuilder) BuildHistoryForAgent(history []entity.ChatMessage) []ent } // BuildAgentRequestFields 深度模式:把 builder 中的摘要 / 记忆 / 用户上下文填充到 agent.Request 对应字段 -// P1-⑦(PromptBuilder 单入口):System Prompt 由 PromptBuilder.BuildSystem() 统一产出后直接塞到 -// agent.Request.SystemPrompt,agent.runAgent 不再二次拼接,保证快速/深度两模式的摘要/记忆/偏好注入完全一致。 +// System Prompt 由 PromptBuilder.BuildSystem() 统一产出后塞到 agent.Request.SystemPrompt, +// agent.runAgent 只负责在前面拼接 ReAct 规则,保证快速/深度两模式的摘要/记忆/偏好注入完全一致。 +// 同时完整填充 UserCtx 作为兜底:如果将来 runAgent 因某种原因拿到空的 SystemPrompt, +// 还能从 UserCtx + Summary + Memories 重建增强提示词。 func (b *PromptBuilder) BuildAgentRequestFields(userID, query, modelID, modelType string, kbIDs []string, history []entity.ChatMessage) agentpkg.Request { return agentpkg.Request{ UserID: userID, @@ -266,13 +268,26 @@ func (b *PromptBuilder) BuildAgentRequestFields(userID, query, modelID, modelTyp ModelType: modelType, Summary: b.summary, Memories: b.memories, - UserCtx: agentpkg.PromptUserContext{ - ID: b.userCtx.ID, - Username: b.userCtx.Username, - Role: b.userCtx.Role, - TimeStr: b.userCtx.TimeStr, - }, - SystemPrompt: b.BuildSystem(), + UserCtx: b.toAgentPromptUserContext(), + SystemPrompt: b.BuildSystem(), + } +} + +// toAgentPromptUserContext 把内部 UserContext 完整转换为 agent.PromptUserContext +func (b *PromptBuilder) toAgentPromptUserContext() agentpkg.PromptUserContext { + return agentpkg.PromptUserContext{ + ID: b.userCtx.ID, + Username: b.userCtx.Username, + Role: b.userCtx.Role, + TimeStr: b.userCtx.TimeStr, + Department: b.userCtx.Department, + Position: b.userCtx.Position, + Expertise: b.userCtx.Expertise, + Language: b.userCtx.Language, + Timezone: b.userCtx.Timezone, + AnswerStyle: b.userCtx.AnswerStyle, + TableFirst: b.userCtx.TableFirst, + CitationStyle: b.userCtx.CitationStyle, } } From b0ce5bd7e63e6a243d5246cb2bb7c729046fc0fa Mon Sep 17 00:00:00 2001 From: st <2663600842@qq.com> Date: Fri, 7 Aug 2026 09:48:47 +0800 Subject: [PATCH 03/10] =?UTF-8?q?test(graph):=20QuickAnswerGraph=20Rewrite?= =?UTF-8?q?=20=E8=8A=82=E7=82=B9=E9=9B=86=E6=88=90=E6=B5=8B=E8=AF=95?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - mockChatModel 实现 eino BaseChatModel(Generate + Stream) - 6 组 doRewriteWithLLM 测试覆盖: - greeting 意图识别 + skipRetrieve=true + rewritten fallback 原始 query - question 指代消解(历史上下文拼入 + keywords 提取) - chitchat 意图识别 + skipRetrieve=true - 降级链路:无 ChatModel / Generate 报错 / JSON 解析失败 均 fallback 原始 query - init() 调用 logger.InitDefault() 避免 nil pointer - go build / go vet / go test ./... 全通过 --- .../service/chat_service_graph_quick_test.go | 342 ++++++++++++++++++ 1 file changed, 342 insertions(+) create mode 100644 internal/service/chat_service_graph_quick_test.go diff --git a/internal/service/chat_service_graph_quick_test.go b/internal/service/chat_service_graph_quick_test.go new file mode 100644 index 0000000..ea8708e --- /dev/null +++ b/internal/service/chat_service_graph_quick_test.go @@ -0,0 +1,342 @@ +package service + +import ( + "context" + "fmt" + "io" + "testing" + + einoModel "github.com/cloudwego/eino/components/model" + "github.com/cloudwego/eino/schema" + pkgLogger "solvify-agent/pkg/logger" +) + +// mockChatModel eino BaseChatModel 的最小实现,用于测试 +type mockChatModel struct { + generateFn func(ctx context.Context, input []*schema.Message, opts ...einoModel.Option) (*schema.Message, error) + streamFn func(ctx context.Context, input []*schema.Message, opts ...einoModel.Option) (*schema.StreamReader[*schema.Message], error) +} + +func (m *mockChatModel) Generate(ctx context.Context, input []*schema.Message, opts ...einoModel.Option) (*schema.Message, error) { + if m.generateFn != nil { + return m.generateFn(ctx, input, opts...) + } + return &schema.Message{Content: "mock response"}, nil +} + +func (m *mockChatModel) Stream(ctx context.Context, input []*schema.Message, opts ...einoModel.Option) (*schema.StreamReader[*schema.Message], error) { + if m.streamFn != nil { + return m.streamFn(ctx, input, opts...) + } + return nil, io.EOF +} + +func TestIsValidIntent(t *testing.T) { + cases := []struct { + name string + intent string + want bool + }{ + {"greeting 合法", intentGreeting, true}, + {"chitchat 合法", intentChitchat, true}, + {"question 合法", intentQuestion, true}, + {"identity 合法", intentIdentity, true}, + {"meta 合法", intentMeta, true}, + {"空字符串非法", "", false}, + {"未知值非法", "tool_call", false}, + {"拼写错误非法", "questions", false}, + } + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + if got := isValidIntent(c.intent); got != c.want { + t.Errorf("isValidIntent(%q) = %v, want %v", c.intent, got, c.want) + } + }) + } +} + +func TestBuildRewriteHistory(t *testing.T) { + // 构造 InputMsgs: [System, User1, Assistant1, User2, Assistant2, User3(current)] + msgs := []*schema.Message{ + schema.SystemMessage("你是知识助理"), + schema.UserMessage("Go 的接口怎么实现"), + schema.AssistantMessage("接口用 interface 关键字定义...", nil), + schema.UserMessage("错误处理有哪些方式"), + schema.AssistantMessage("Go 用 error 返回值...", nil), + schema.UserMessage("那 context 怎么用"), // UserQuestionIndex = 5 + } + + history := buildRewriteHistory(msgs, 5, 3) + + // 应排除 system 和当前 User3,只保留 User1/Assistant1/User2/Assistant2 + if history == "" { + t.Fatal("history 为空,应该有内容") + } + if !containsAll(history, "Go 的接口怎么实现", "接口用 interface 关键字", "错误处理有哪些方式", "Go 用 error 返回值") { + t.Errorf("history 未包含预期历史内容: %s", history) + } + if containsAll(history, "你是知识助理") { + t.Error("history 不应该包含 system message") + } + if containsAll(history, "那 context 怎么用") { + t.Error("history 不应该包含当前用户问题") + } +} + +func TestBuildRewriteHistory_EmptyMsgs(t *testing.T) { + history := buildRewriteHistory(nil, 0, 3) + if history != "" { + t.Errorf("nil msgs 应返回空,got: %q", history) + } +} + +func TestBuildRewriteHistory_NoPreviousMessages(t *testing.T) { + msgs := []*schema.Message{ + schema.SystemMessage("你好"), + schema.UserMessage("这是第一个问题"), + } + history := buildRewriteHistory(msgs, 1, 3) + if history != "" { + t.Errorf("没有历史消息应返回空,got: %q", history) + } +} + +func TestBuildRewriteHistory_MaxRoundsLimit(t *testing.T) { + // 构造 6 轮历史,maxRounds=2 → 应该只保留最后 2 轮 + msgs := []*schema.Message{ + schema.SystemMessage("sys"), + schema.UserMessage("q1"), schema.AssistantMessage("a1", nil), + schema.UserMessage("q2"), schema.AssistantMessage("a2", nil), + schema.UserMessage("q3"), schema.AssistantMessage("a3", nil), + schema.UserMessage("q4"), schema.AssistantMessage("a4", nil), + schema.UserMessage("q5"), schema.AssistantMessage("a5", nil), + schema.UserMessage("q6"), schema.AssistantMessage("a6", nil), + schema.UserMessage("current"), // idx = 13 + } + history := buildRewriteHistory(msgs, 13, 2) + + if containsAll(history, "q1", "a1") { + t.Error("maxRounds=2 不应包含最早的 q1/a1") + } + if !containsAll(history, "q5", "a5", "q6", "a6") { + t.Errorf("maxRounds=2 应保留最后两轮 q5/a5, q6/a6, got: %s", history) + } +} + +func containsAll(s string, substrs ...string) bool { + for _, sub := range substrs { + if !containsStr(s, sub) { + return false + } + } + return true +} + +func containsStr(s, sub string) bool { + return len(s) >= len(sub) && (s == sub || len(sub) == 0 || indexOf(s, sub) >= 0) +} + +func indexOf(s, sub string) int { + for i := 0; i+len(sub) <= len(s); i++ { + if s[i:i+len(sub)] == sub { + return i + } + } + return -1 +} + +// --- doRewriteWithLLM 单元测试 --- + +func init() { + // 初始化默认 logger,避免 nil pointer + _ = pkgLogger.InitDefault() +} + +func TestDoRewriteWithLLM_Greeting(t *testing.T) { + // mock ChatModel 返回 greeting JSON + mock := &mockChatModel{ + generateFn: func(ctx context.Context, input []*schema.Message, opts ...einoModel.Option) (*schema.Message, error) { + return &schema.Message{Content: `{"rewritten":"","intent":"greeting","keywords":[]}`}, nil + }, + } + ctx := withGraphChatModel(context.Background(), mock) + input := &quickGraphInput{ + OriginalQuery: "你好", + InputMsgs: []*schema.Message{ + schema.SystemMessage("你是知识助理"), + schema.UserMessage("你好"), + }, + UserQuestionIndex: 1, + ModelName: "test-model", + } + + rewritten, intent, _, skipRetrieve := doRewriteWithLLM(ctx, input) + + if intent != intentGreeting { + t.Errorf("intent = %q, want %q", intent, intentGreeting) + } + if !skipRetrieve { + t.Error("greeting 场景 skipRetrieve 应为 true") + } + // 空 rewritten 会被 fallback 成原始 query(合理:改写不能是空的) + if rewritten != "你好" { + t.Errorf("greeting 场景 rewritten 应 fallback 为原始 query, got %q", rewritten) + } +} + +func TestDoRewriteWithLLM_QuestionWithHistory(t *testing.T) { + // mock ChatModel 返回 question JSON,做指代消解 + mock := &mockChatModel{ + generateFn: func(ctx context.Context, input []*schema.Message, opts ...einoModel.Option) (*schema.Message, error) { + // 验证输入里包含历史上下文 + if len(input) != 2 { + t.Errorf("input msgs 应为 [system, user], got len=%d", len(input)) + } + return &schema.Message{Content: `{"rewritten":"Go 中 interface 和 context 怎么配合使用","intent":"question","keywords":["Go","interface","context"]}`}, nil + }, + } + ctx := withGraphChatModel(context.Background(), mock) + input := &quickGraphInput{ + OriginalQuery: "那 context 呢", + InputMsgs: []*schema.Message{ + schema.SystemMessage("你是知识助理"), + schema.UserMessage("Go 怎么实现接口"), + schema.AssistantMessage("用 interface 关键字...", nil), + schema.UserMessage("那 context 呢"), + }, + UserQuestionIndex: 3, + ModelName: "test-model", + } + + rewritten, intent, keywords, skipRetrieve := doRewriteWithLLM(ctx, input) + + if intent != intentQuestion { + t.Errorf("intent = %q, want %q", intent, intentQuestion) + } + if skipRetrieve { + t.Error("question 场景 skipRetrieve 应为 false") + } + if rewritten != "Go 中 interface 和 context 怎么配合使用" { + t.Errorf("rewritten = %q, want 改写后的完整问题", rewritten) + } + if len(keywords) != 3 { + t.Errorf("keywords len = %d, want 3", len(keywords)) + } +} + +func TestDoRewriteWithLLM_Fallback_NoChatModel(t *testing.T) { + // context 里没有 ChatModel → 应 fallback 原始 query + input := &quickGraphInput{ + OriginalQuery: "那 context 呢", + InputMsgs: []*schema.Message{ + schema.SystemMessage("你是知识助理"), + schema.UserMessage("那 context 呢"), + }, + UserQuestionIndex: 1, + ModelName: "test-model", + } + + rewritten, intent, _, skipRetrieve := doRewriteWithLLM(context.Background(), input) + + if rewritten != input.OriginalQuery { + t.Errorf("无 ChatModel 时 rewritten 应等于原始 query, got %q", rewritten) + } + if intent != intentQuestion { + t.Errorf("无 ChatModel 时 intent 应为 question, got %q", intent) + } + if skipRetrieve { + t.Error("无 ChatModel 时不应跳过检索") + } +} + +func TestDoRewriteWithLLM_Fallback_GenerateError(t *testing.T) { + // ChatModel.Generate 报错 → 应 fallback + mock := &mockChatModel{ + generateFn: func(ctx context.Context, input []*schema.Message, opts ...einoModel.Option) (*schema.Message, error) { + return nil, fmt.Errorf("LLM timeout") + }, + } + ctx := withGraphChatModel(context.Background(), mock) + input := &quickGraphInput{ + OriginalQuery: "测试问题", + InputMsgs: []*schema.Message{ + schema.SystemMessage("你是知识助理"), + schema.UserMessage("测试问题"), + }, + UserQuestionIndex: 1, + ModelName: "test-model", + } + + rewritten, intent, _, skipRetrieve := doRewriteWithLLM(ctx, input) + + if rewritten != input.OriginalQuery { + t.Errorf("Generate 报错时应 fallback 原始 query, got %q", rewritten) + } + if intent != intentQuestion { + t.Errorf("Generate 报错时 intent 应为 question, got %q", intent) + } + if skipRetrieve { + t.Error("Generate 报错时不应跳过检索") + } +} + +func TestDoRewriteWithLLM_Fallback_JSONParseError(t *testing.T) { + // LLM 返回无效 JSON → 应 fallback + mock := &mockChatModel{ + generateFn: func(ctx context.Context, input []*schema.Message, opts ...einoModel.Option) (*schema.Message, error) { + return &schema.Message{Content: "不是有效的 JSON"}, nil + }, + } + ctx := withGraphChatModel(context.Background(), mock) + input := &quickGraphInput{ + OriginalQuery: "测试问题", + InputMsgs: []*schema.Message{ + schema.SystemMessage("你是知识助理"), + schema.UserMessage("测试问题"), + }, + UserQuestionIndex: 1, + ModelName: "test-model", + } + + rewritten, intent, _, _ := doRewriteWithLLM(ctx, input) + + if rewritten != input.OriginalQuery { + t.Errorf("JSON 解析失败时应 fallback 原始 query, got %q", rewritten) + } + if intent != intentQuestion { + t.Errorf("JSON 解析失败时 intent 应为 question, got %q", intent) + } +} + +func TestDoRewriteWithLLM_Chitchat(t *testing.T) { + mock := &mockChatModel{ + generateFn: func(ctx context.Context, input []*schema.Message, opts ...einoModel.Option) (*schema.Message, error) { + return &schema.Message{Content: `{"rewritten":"","intent":"chitchat","keywords":[]}`}, nil + }, + } + ctx := withGraphChatModel(context.Background(), mock) + input := &quickGraphInput{ + OriginalQuery: "讲个笑话", + InputMsgs: []*schema.Message{ + schema.SystemMessage("你是知识助理"), + schema.UserMessage("讲个笑话"), + }, + UserQuestionIndex: 1, + ModelName: "test-model", + } + + _, intent, _, skipRetrieve := doRewriteWithLLM(ctx, input) + + if intent != intentChitchat { + t.Errorf("intent = %q, want %q", intent, intentChitchat) + } + if !skipRetrieve { + t.Error("chitchat 场景 skipRetrieve 应为 true") + } +} + +// --- QuickAnswerGraph 集成测试 --- + +// 注:完整 QuickAnswerGraph 端到端测试需要 EinoRetrieverAdapter + 真实 Retriever + Runnable.Compile/Invoke 流程, +// 通过 HTTP API 层测试更合适。这里 doRewriteWithLLM 单元测试已覆盖 Rewrite 节点核心逻辑。 +// 完整链路:HTTP /metrics 端点已在运行时验证(greeting → skip_retrieve=true、question → rewrite → retrieve → generate) From 871d9aaaf3f496b174423b73518e7a803ebd774c Mon Sep 17 00:00:00 2001 From: st <2663600842@qq.com> Date: Sat, 8 Aug 2026 08:45:43 +0800 Subject: [PATCH 04/10] =?UTF-8?q?feat(p2-1):=20=E5=90=91=E9=87=8F=E5=8E=86?= =?UTF-8?q?=E5=8F=B2=E6=A3=80=E7=B4=A2=20pgvector=20=E8=AF=AD=E4=B9=89?= =?UTF-8?q?=E9=80=9A=E9=81=93?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - ChatMessage entity 加 Embedding vector(1024) 字段 - Repository 新增 SearchRecentByVector (pgvector 余弦距离) + UpdateEmbedding - chat_service 写入路径后台异步算 embedding (5s 超时,失败 warn) - context_service 读取路径: 向量语义检索优先, ILIKE 关键词兜底 - DB migration 已执行: pgvector 0.8.2 + HNSW 索引 --- internal/model/entity/chat_message.go | 4 +- internal/repository/chat_message_interface.go | 6 +- .../repository/chat_message_repository.go | 64 +++++++++++++++++++ internal/service/chat_service.go | 29 ++++++++- internal/service/context_service.go | 6 ++ 5 files changed, 106 insertions(+), 3 deletions(-) diff --git a/internal/model/entity/chat_message.go b/internal/model/entity/chat_message.go index 121aecd..109d756 100644 --- a/internal/model/entity/chat_message.go +++ b/internal/model/entity/chat_message.go @@ -17,7 +17,9 @@ type ChatMessage struct { KnowledgeBaseIDs datatypes.JSON `gorm:"column:knowledge_base_ids;type:jsonb;not null;default:'[]'::jsonb"` Sources datatypes.JSON `gorm:"column:sources;type:jsonb"` Metadata datatypes.JSON `gorm:"column:metadata;type:jsonb;not null;default:'{}'::jsonb"` - CreatedAt time.Time `gorm:"column:created_at;autoCreateTime"` + // Embedding 消息内容的向量表示,用于语义相关历史检索 + Embedding FloatVector `gorm:"column:embedding;type:vector(1024)"` + CreatedAt time.Time `gorm:"column:created_at;autoCreateTime"` } // TableName 返回聊天消息表名 diff --git a/internal/repository/chat_message_interface.go b/internal/repository/chat_message_interface.go index c54a03f..4202405 100644 --- a/internal/repository/chat_message_interface.go +++ b/internal/repository/chat_message_interface.go @@ -30,6 +30,10 @@ type ChatMessageRepo interface { DeleteBySessionID(ctx context.Context, sessionID string) error // SearchByKeyword 按关键字搜索用户历史消息 SearchByKeyword(ctx context.Context, userID, query string, topK int) ([]ChatMessageSearchRow, error) - // SearchRecentByKeywords 在指定会话中按关键词检索最近消息 + // SearchRecentByKeywords 在指定会话中按关键词检索最近消息(ILIKE 兜底路径) SearchRecentByKeywords(ctx context.Context, sessionID string, keywords []string, limit int) ([]entity.ChatMessage, error) + // SearchRecentByVector 在指定会话中按向量语义检索最近消息(pgvector 余弦距离) + SearchRecentByVector(ctx context.Context, sessionID string, queryEmbedding []float32, limit int, distanceThreshold float64) ([]entity.ChatMessage, error) + // UpdateEmbedding 更新指定消息的向量表示(后台写入路径,失败不影响主流程) + UpdateEmbedding(ctx context.Context, messageID string, embedding entity.FloatVector) error } diff --git a/internal/repository/chat_message_repository.go b/internal/repository/chat_message_repository.go index d6ee047..f46ae89 100644 --- a/internal/repository/chat_message_repository.go +++ b/internal/repository/chat_message_repository.go @@ -2,6 +2,8 @@ package repository import ( "context" + "fmt" + "strings" "gorm.io/gorm" @@ -151,3 +153,65 @@ func (r *chatMessageRepository) SearchByKeyword(ctx context.Context, userID, que return results, err } + +// SearchRecentByVector 在指定会话中按向量语义检索最近消息(pgvector 余弦距离)。 +// 仅返回有 embedding 的消息(embedding IS NOT NULL),距离阈值用于过滤不相关结果。 +func (r *chatMessageRepository) SearchRecentByVector(ctx context.Context, sessionID string, queryEmbedding []float32, limit int, distanceThreshold float64) ([]entity.ChatMessage, error) { + if limit <= 0 { + limit = 5 + } + if len(queryEmbedding) == 0 { + return nil, nil + } + if distanceThreshold <= 0 { + distanceThreshold = 0.8 + } + + embeddingStr := formatFloatVector(queryEmbedding) + + var messages []entity.ChatMessage + err := r.db.WithContext(ctx).Raw(` + SELECT id, session_id, role, content, model_id, search_mode, knowledge_base_ids, sources, metadata, embedding, created_at + FROM chat_messages + WHERE session_id = ? + AND embedding IS NOT NULL + AND embedding <-> ? <= ? + ORDER BY embedding <-> ? + LIMIT ? + `, sessionID, embeddingStr, distanceThreshold, embeddingStr, limit).Scan(&messages).Error + + if err != nil { + return nil, err + } + + for i, j := 0, len(messages)-1; i < j; i, j = i+1, j-1 { + messages[i], messages[j] = messages[j], messages[i] + } + + return messages, nil +} + +// formatFloatVector 把 []float32 格式化为 pgvector 字面量 '[1.0,2.0,...]' +func formatFloatVector(v []float32) string { + if len(v) == 0 { + return "[]" + } + var sb strings.Builder + sb.WriteByte('[') + for i, f := range v { + if i > 0 { + sb.WriteByte(',') + } + sb.WriteString(fmt.Sprintf("%.6f", f)) + } + sb.WriteByte(']') + return sb.String() +} + +// UpdateEmbedding 更新指定消息的向量表示 +func (r *chatMessageRepository) UpdateEmbedding(ctx context.Context, messageID string, embedding entity.FloatVector) error { + return r.db.WithContext(ctx).Model(&entity.ChatMessage{}). + Where("id = ?", messageID). + Update("embedding", embedding).Error +} + diff --git a/internal/service/chat_service.go b/internal/service/chat_service.go index 24d7410..7e39376 100644 --- a/internal/service/chat_service.go +++ b/internal/service/chat_service.go @@ -45,6 +45,7 @@ type chatService struct { prefSvc UserPreferenceService obs observability.Recorder obsRepo repository.ObservabilityRepo + embedClient *llm.EmbeddingClient } // NewChatService 创建聊天业务服务 @@ -628,9 +629,35 @@ func (s *chatService) saveAssistantMessage(ctx context.Context, sessionID, msgID Sources: datatypes.JSON(mustMarshal(sources)), Metadata: metadata, } - return s.messageRepo.Create(ctx, &assistantMsg) + if err := s.messageRepo.Create(ctx, &assistantMsg); err != nil { + return err + } + s.bgComputeEmbedding(ctx, assistantMsg.ID, assistantMsg.Content) + return nil } +// bgComputeEmbedding 后台为消息计算向量表示,失败不阻塞主流程。 +func (s *chatService) bgComputeEmbedding(ctx context.Context, messageID, content string) { + if s.embedClient == nil || strings.TrimSpace(content) == "" { + return + } + go func() { + embedCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + vec, err := s.embedClient.Embed(embedCtx, content) + if err != nil { + logger.Warnf("消息向量计算失败: messageID=%s, err=%v", messageID, err) + return + } + f32 := make(entity.FloatVector, len(vec)) + for i, v := range vec { + f32[i] = float32(v) + } + if err := s.messageRepo.UpdateEmbedding(embedCtx, messageID, f32); err != nil { + logger.Warnf("消息向量更新失败: messageID=%s, err=%v", messageID, err) + } + }() +} // truncateHistoryByTokens 按真 BPE token 预算从尾部向前保留"完整轮对"。 // // 关键修复(P0-②):旧代码"从尾部逐个 append 头插"实际是按时间正序塞,但因为 diff --git a/internal/service/context_service.go b/internal/service/context_service.go index 0777eeb..47065ee 100644 --- a/internal/service/context_service.go +++ b/internal/service/context_service.go @@ -12,6 +12,7 @@ import ( "github.com/cloudwego/eino/components/model" "github.com/cloudwego/eino/schema" + "solvify-agent/internal/llm" "solvify-agent/internal/model/entity" "solvify-agent/internal/observability" "solvify-agent/internal/repository" @@ -28,6 +29,7 @@ type contextService struct { memoryRepo repository.UserMemoryRepo summaryRepo repository.SummaryRepo obs observability.Recorder + embedClient *llm.EmbeddingClient // NewContextService 创建上下文管理服务 } @@ -52,6 +54,10 @@ func NewContextService( func (s *contextService) SetObservability(obs observability.Recorder) { s.obs = obs } +// SetEmbedClient 注入向量客户端,用于语义相关历史检索 +func (s *contextService) SetEmbedClient(client *llm.EmbeddingClient) { + s.embedClient = client +} // BuildContext 构建增强后的对话上下文 func (s *contextService) BuildContext(ctx context.Context, userID, sessionID, currentQuery string, cfg BuildContextConfig, chatModel model.BaseChatModel) (*EnhancedContext, error) { From 14fa4d080d612d7778e2aa949584d3cd552670aa Mon Sep 17 00:00:00 2001 From: st <2663600842@qq.com> Date: Sat, 8 Aug 2026 21:22:06 +0800 Subject: [PATCH 05/10] =?UTF-8?q?feat(agent):=20=E6=8E=A5=E5=85=A5=20Eino?= =?UTF-8?q?=20Interrupt=20=E6=9C=BA=E5=88=B6=E5=AE=9E=E7=8E=B0=E5=8D=B1?= =?UTF-8?q?=E9=99=A9=E5=B7=A5=E5=85=B7=E5=AE=A1=E6=89=B9=E4=B8=AD=E9=97=B4?= =?UTF-8?q?=E5=B1=82?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 新增 tool_middleware.go: 拦截 delete_document 等危险工具,通过 StatefulInterrupt 中断执行 - 新增 runner_adapter.go: 将 Eino InterruptSignal 转换为 AgentEvent.Interrupt 事件 - 新增 checkpoint_store.go + gob 注册: 持久化 InterruptState,支持恢复执行 - engine/execute/engine_tools 接入中间件,callback 处理中断回调 - prompt 增加危险工具澄清规则: 目标不明确先反问用户,禁止编造参数 - go.mod 更新 eino 版本 --- go.mod | 1 + internal/agent/callback.go | 401 +++------------------------ internal/agent/checkpoint_store.go | 74 +++++ internal/agent/engine.go | 88 ++++-- internal/agent/engine_tools.go | 37 +-- internal/agent/execute.go | 372 ++++++++----------------- internal/agent/prompt.go | 133 +++++---- internal/agent/runner_adapter.go | 426 +++++++++++++++++++++++++++++ internal/agent/tool_middleware.go | 92 +++++++ internal/agent/types.go | 8 + internal/app/app.go | 75 +++-- internal/tool/delete_document.go | 71 +++++ internal/tool/document_tools.go | 323 ++++++++-------------- 13 files changed, 1136 insertions(+), 965 deletions(-) create mode 100644 internal/agent/checkpoint_store.go create mode 100644 internal/agent/runner_adapter.go create mode 100644 internal/agent/tool_middleware.go create mode 100644 internal/tool/delete_document.go diff --git a/go.mod b/go.mod index 838e87c..92ee655 100644 --- a/go.mod +++ b/go.mod @@ -111,3 +111,4 @@ require ( gopkg.in/yaml.v3 v3.0.1 // indirect gorm.io/driver/mysql v1.5.6 // indirect ) + diff --git a/internal/agent/callback.go b/internal/agent/callback.go index f62b57d..bf93d87 100644 --- a/internal/agent/callback.go +++ b/internal/agent/callback.go @@ -1,324 +1,64 @@ package agent import ( - "context" "encoding/json" "fmt" "strings" - "time" - "github.com/cloudwego/eino/callbacks" - "github.com/cloudwego/eino/components/model" toolComp "github.com/cloudwego/eino/components/tool" - - "solvify-agent/internal/observability" - "solvify-agent/pkg/logger" ) -type agentCallbackHandler struct { - eventCh chan<- Event - callCount int - pendingThinkingTitle string - kbIDs []string - toolDescMap map[string]string - sentToolEvents map[string]bool - - taskID string - tracker *agentStepTracker - obs observability.Recorder +// ToolResponseData 描述工具返回的 JSON 结构 +type ToolResponseData struct { + Success bool `json:"success"` + Message string `json:"message"` + Data interface{} `json:"data"` } -func newAgentCallbackHandler(eventCh chan<- Event, kbIDs []string, toolDescMap map[string]string) *agentCallbackHandler { - h := &agentCallbackHandler{ - eventCh: eventCh, - kbIDs: kbIDs, - toolDescMap: toolDescMap, - sentToolEvents: make(map[string]bool), +func parseToolResponse(response string) ToolResponseData { + var result ToolResponseData + if err := json.Unmarshal([]byte(response), &result); err != nil { + return ToolResponseData{Success: false, Message: response} } - return h -} - -func (h *agentCallbackHandler) Handler() callbacks.Handler { - return callbacks.NewHandlerBuilder(). - OnStartFn(h.onStart). - OnEndFn(h.onEnd). - OnErrorFn(h.onError). - Build() + return result } -func (h *agentCallbackHandler) onStart(ctx context.Context, info *callbacks.RunInfo, input callbacks.CallbackInput) context.Context { - if info == nil { - return ctx - } - obsOk := h.obs != nil && h.tracker != nil - - switch info.Component { - case "ChatModel": - h.callCount++ - h.completeThinking() - - if h.callCount == 1 { - h.pendingThinkingTitle = "分析问题" - h.emit(Event{ - Type: EventThinking, - Title: "分析问题", - Detail: "理解用户意图,确定检索方向", - Status: "running", - }) - } else { - h.pendingThinkingTitle = "分析检索结果" - h.emit(Event{ - Type: EventThinking, - Title: "分析检索结果", - Detail: "评估检索内容,决定下一步行动", - Status: "running", - }) - } - if obsOk { - h.tracker.mu.Lock() - h.tracker.stepIdx++ - idx := h.tracker.stepIdx - pending := &agentStepPending{ - StepIndex: idx, - TaskID: h.taskID, - ThinkingSummary: h.pendingThinkingTitle, - StartedAt: time.Now(), - } - h.tracker.pendingByID[fmt.Sprintf("llm:%d", h.callCount)] = pending - h.tracker.mu.Unlock() - } - - case "Tool": - toolInput := toolComp.ConvCallbackInput(input) - toolName := info.Name - query := "" - inputJSON := "" - if toolInput != nil { - inputJSON = toolInput.ArgumentsInJSON - query = extractQueryFromArgs(inputJSON) - } - title, detail := formatToolStart(toolName, query, h.kbIDs, h.toolDescMap) - - eventKey := fmt.Sprintf("call:%s:%s", toolName, title) - if h.sentToolEvents[eventKey] { - logger.Warnf("[Callback] 跳过重复的工具调用事件: %s", eventKey) - return ctx - } - h.sentToolEvents[eventKey] = true - - h.emit(Event{ - Type: EventToolCall, - Title: title, - Detail: detail, - Status: "running", - }) - - if obsOk { - h.tracker.mu.Lock() - h.tracker.stepIdx++ - idx := h.tracker.stepIdx - pending := &agentStepPending{ - StepIndex: idx, - TaskID: h.taskID, - ToolName: toolName, - ToolInputMasked: truncateStr(maskJsonSecrets(inputJSON), 256), - StartedAt: time.Now(), - } - h.tracker.pendingByID[fmt.Sprintf("tool:%s:%s", toolName, title)] = pending - h.tracker.mu.Unlock() - } - } - - return ctx +func parseGrepResponse(response string) ToolResponseData { + return parseToolResponse(response) } -func (h *agentCallbackHandler) onEnd(ctx context.Context, info *callbacks.RunInfo, output callbacks.CallbackOutput) context.Context { - if info == nil { - return ctx - } - obsOk := h.obs != nil && h.tracker != nil - - switch info.Component { - case "ChatModel": - modelOutput := model.ConvCallbackOutput(output) - if modelOutput != nil && modelOutput.Message != nil { - if len(modelOutput.Message.ToolCalls) > 0 { - for _, tc := range modelOutput.Message.ToolCalls { - query := extractQueryFromArgs(tc.Function.Arguments) - logger.Infof("[Callback] LLM 决定调用: %s(%q)", tc.Function.Name, query) - } - h.completeThinking() - } else { - h.completeThinking() - h.pendingThinkingTitle = "正在生成答案" - h.emit(Event{ - Type: EventThinking, - Title: "正在生成答案", - Status: "running", - }) - } - } - if obsOk { - h.tracker.mu.Lock() - key := fmt.Sprintf("llm:%d", h.callCount) - pending := h.tracker.pendingByID[key] - delete(h.tracker.pendingByID, key) - h.tracker.mu.Unlock() - if pending != nil { - step := &observability.AgentStep{ - TaskID: pending.TaskID, - StepIndex: pending.StepIndex, - StartedAt: pending.StartedAt, - EndedAt: time.Now(), - ThinkingSummary: pending.ThinkingSummary, - LatencyMs: time.Since(pending.StartedAt).Milliseconds(), - ToolName: "llm.reasoning", - ToolStatus: "success", - } - h.obs.RecordAgentStep(step) - } - } - - case "Tool": - toolOutput := toolComp.ConvCallbackOutput(output) - toolName := info.Name - title, detail, toolResult := formatToolEnd(toolName, toolOutput, h.toolDescMap) - - eventKey := fmt.Sprintf("result:%s:%s", toolName, title) - if h.sentToolEvents[eventKey] { - logger.Warnf("[Callback] 跳过重复的工具完成事件: %s", eventKey) - return ctx - } - h.sentToolEvents[eventKey] = true - - h.emit(Event{ - Type: EventToolResult, - Title: title, - Detail: detail, - Status: "success", - ToolResult: toolResult, - }) - - if obsOk { - h.tracker.mu.Lock() - key := fmt.Sprintf("tool:%s:%s", toolName, title) - pending := h.tracker.pendingByID[key] - delete(h.tracker.pendingByID, key) - h.tracker.mu.Unlock() - if pending != nil { - step := &observability.AgentStep{ - TaskID: pending.TaskID, - StepIndex: pending.StepIndex, - StartedAt: pending.StartedAt, - EndedAt: time.Now(), - ToolName: pending.ToolName, - ToolInputMasked: pending.ToolInputMasked, - ToolResultSummary: truncateStr(toolResult, 256), - ToolStatus: "success", - LatencyMs: time.Since(pending.StartedAt).Milliseconds(), - } - h.obs.RecordAgentStep(step) - } - } +func parseKnowledgeSearchResult(response string) (titles []string, count int) { + var result struct { + Sources []struct { + Title string `json:"title"` + } `json:"sources"` } - - return ctx -} - -func (h *agentCallbackHandler) onError(ctx context.Context, info *callbacks.RunInfo, err error) context.Context { - if info == nil { - return ctx + if err := json.Unmarshal([]byte(response), &result); err != nil { + return nil, 0 } - logger.Errorf("[Callback] 组件出错: node=%s, component=%s, err=%v", info.Name, info.Component, err) - - title, detail, retryable := formatToolError(string(info.Component), info.Name, err) - h.emit(Event{ - Type: EventError, - Title: title, - Detail: detail, - Error: err.Error(), - Status: "error", - Retryable: retryable, - Done: true, - }) - - if h.obs != nil && h.tracker != nil { - h.tracker.mu.Lock() - var pending *agentStepPending - for k, v := range h.tracker.pendingByID { - pending = v - delete(h.tracker.pendingByID, k) - break - } - h.tracker.mu.Unlock() - if pending != nil { - step := &observability.AgentStep{ - TaskID: pending.TaskID, - StepIndex: pending.StepIndex, - StartedAt: pending.StartedAt, - EndedAt: time.Now(), - ThinkingSummary: pending.ThinkingSummary, - ToolName: pending.ToolName, - ToolInputMasked: pending.ToolInputMasked, - ToolResultSummary: "", - ToolStatus: "error", - ToolError: truncateStr(err.Error(), 256), - LatencyMs: time.Since(pending.StartedAt).Milliseconds(), - } - h.obs.RecordAgentStep(step) + count = len(result.Sources) + seen := make(map[string]bool, count) + for _, s := range result.Sources { + if s.Title != "" && !seen[s.Title] { + seen[s.Title] = true + titles = append(titles, s.Title) } } - - return ctx -} - -func maskJsonSecrets(s string) string { - if s == "" { - return "" - } - var obj any - if err := json.Unmarshal([]byte(s), &obj); err != nil { - return s - } - masked, _ := json.Marshal(maskAny(obj)) - return string(masked) + return titles, count } -func maskAny(v any) any { - switch x := v.(type) { - case map[string]any: - out := make(map[string]any, len(x)) - for k, val := range x { - lk := strings.ToLower(k) - if strings.Contains(lk, "key") || strings.Contains(lk, "token") || - strings.Contains(lk, "password") || strings.Contains(lk, "secret") { - if s, ok := val.(string); ok { - out[k] = maskSecret(s) - continue - } - } - out[k] = maskAny(val) - } - return out - case []any: - out := make([]any, len(x)) - for i, val := range x { - out[i] = maskAny(val) +func isWebSearchTool(name, desc string) bool { + combined := strings.ToLower(name + " " + desc) + for _, kw := range []string{"web", "search", "搜索", "联网", "tavily", "serp", "bocha", "sogou", "bing"} { + if strings.Contains(combined, kw) { + return true } - return out - default: - return v } + return false } -func maskSecret(s string) string { - if len(s) <= 8 { - return "***" - } - return s[:2] + "***" + s[len(s)-2:] -} - +// formatToolStart 根据工具名和查询内容生成 EventToolCall 的 title 和 detail func formatToolStart(toolName, query string, kbIDs []string, toolDescMap map[string]string) (title, detail string) { switch toolName { case "knowledge_search": @@ -365,16 +105,7 @@ func formatToolStart(toolName, query string, kbIDs []string, toolDescMap map[str return fmt.Sprintf("正在执行 %s", label), "" } -func isWebSearchTool(name, desc string) bool { - combined := strings.ToLower(name + " " + desc) - for _, kw := range []string{"web", "search", "搜索", "联网", "tavily", "serp", "bocha", "sogou", "bing"} { - if strings.Contains(combined, kw) { - return true - } - } - return false -} - +// formatToolEnd 根据工具名和输出生成 EventToolResult 的 title、detail、toolResult func formatToolEnd(toolName string, output *toolComp.CallbackOutput, toolDescMap map[string]string) (title, detail, toolResult string) { response := "" if output != nil { @@ -398,7 +129,7 @@ func formatToolEnd(toolName string, output *toolComp.CallbackOutput, toolDescMap } return "知识库检索完成", "未找到相关内容", toolResult case "grep_chunks": - result := parseGrepResult(response) + result := parseGrepResponse(response) if result.Success && result.Data != nil { if dataList, ok := result.Data.([]interface{}); ok && len(dataList) > 0 { return "关键词搜索完成", fmt.Sprintf("找到 %d 条匹配结果", len(dataList)), toolResult @@ -442,57 +173,7 @@ func formatToolEnd(toolName string, output *toolComp.CallbackOutput, toolDescMap return fmt.Sprintf("%s 执行完成", toolName), "", toolResult } -type ToolResponseData struct { - Success bool `json:"success"` - Message string `json:"message"` - Data interface{} `json:"data"` -} - -func parseToolResponse(response string) ToolResponseData { - var result ToolResponseData - if err := json.Unmarshal([]byte(response), &result); err != nil { - return ToolResponseData{Success: false, Message: response} - } - return result -} - -func parseGrepResult(response string) ToolResponseData { - return parseToolResponse(response) -} - -func parseKnowledgeSearchResult(response string) (titles []string, count int) { - var result struct { - Sources []struct { - Title string `json:"title"` - } `json:"sources"` - } - if err := json.Unmarshal([]byte(response), &result); err != nil { - return nil, 0 - } - - count = len(result.Sources) - seen := make(map[string]bool, count) - for _, s := range result.Sources { - if s.Title != "" && !seen[s.Title] { - seen[s.Title] = true - titles = append(titles, s.Title) - } - } - return titles, count -} - -func (h *agentCallbackHandler) completeThinking() { - if h.pendingThinkingTitle == "" { - return - } - h.emit(Event{ - Type: EventThinking, - Title: h.pendingThinkingTitle, - Status: "success", - }) - h.pendingThinkingTitle = "" -} - +// formatToolError 根据组件类型和错误生成 EventError 的 title、detail、retryable func formatToolError(component, name string, err error) (title, detail string, retryable bool) { errMsg := err.Error() @@ -554,13 +235,3 @@ func formatToolError(component, name string, err error) (title, detail string, r return fmt.Sprintf("%s 执行出错", component), "服务执行异常,请稍后重试", true } } - -func (h *agentCallbackHandler) emit(e Event) { - logger.Infof("[Callback] 发送事件: type=%s, title=%s, status=%s, detail=%s", - e.Type, e.Title, e.Status, truncateStr(e.Detail, 60)) - select { - case h.eventCh <- e: - default: - logger.Warnf("[Callback] ⚠️ 事件通道已满,丢弃事件: type=%s, title=%s", e.Type, e.Title) - } -} diff --git a/internal/agent/checkpoint_store.go b/internal/agent/checkpoint_store.go new file mode 100644 index 0000000..6d076d7 --- /dev/null +++ b/internal/agent/checkpoint_store.go @@ -0,0 +1,74 @@ +package agent + +import ( + "context" + "sync" + "time" + + "solvify-agent/internal/repository" +) + +// CheckpointTTL 是 agent_checkpoints 表中 checkpoint 的默认存活时长 +const CheckpointTTL = 24 * time.Hour + +// InMemoryCheckPointStore 是 core.CheckPointStore 的内存实现。 +// 用于本地开发和单元测试。线程安全。 +type InMemoryCheckPointStore struct { + mu sync.RWMutex + data map[string][]byte +} + +func NewInMemoryCheckPointStore() *InMemoryCheckPointStore { + return &InMemoryCheckPointStore{data: make(map[string][]byte)} +} + +func (s *InMemoryCheckPointStore) Get(_ context.Context, checkPointID string) ([]byte, bool, error) { + s.mu.RLock() + defer s.mu.RUnlock() + v, ok := s.data[checkPointID] + return v, ok, nil +} + +func (s *InMemoryCheckPointStore) Set(_ context.Context, checkPointID string, checkPoint []byte) error { + s.mu.Lock() + defer s.mu.Unlock() + s.data[checkPointID] = checkPoint + return nil +} + +func (s *InMemoryCheckPointStore) Delete(_ context.Context, checkPointID string) error { + s.mu.Lock() + defer s.mu.Unlock() + delete(s.data, checkPointID) + return nil +} + +// DBCheckPointStore 是 core.CheckPointStore 的数据库实现。 +// 通过 AgentCheckpointRepo 持久化 checkpoint 原始字节,支持按 checkpointID 读写删。 +type DBCheckPointStore struct { + repo repository.AgentCheckpointRepo + sessionID string + ttl time.Duration +} + +// NewDBCheckPointStore 从 repo 创建 DB 版 checkpoint store。 +// sessionID 用于把 checkpoint 和聊天会话关联起来;ttl 控制 checkpoint 过期时间。 +func NewDBCheckPointStore(repo repository.AgentCheckpointRepo, sessionID string, ttl time.Duration) *DBCheckPointStore { + if ttl <= 0 { + ttl = CheckpointTTL + } + return &DBCheckPointStore{repo: repo, sessionID: sessionID, ttl: ttl} +} + +func (s *DBCheckPointStore) Get(ctx context.Context, checkPointID string) ([]byte, bool, error) { + return s.repo.Find(ctx, checkPointID) +} + +func (s *DBCheckPointStore) Set(ctx context.Context, checkPointID string, checkPoint []byte) error { + expiredAt := time.Now().Add(s.ttl) + return s.repo.Save(ctx, checkPointID, s.sessionID, checkPoint, expiredAt) +} + +func (s *DBCheckPointStore) Delete(ctx context.Context, checkPointID string) error { + return s.repo.Delete(ctx, checkPointID) +} diff --git a/internal/agent/engine.go b/internal/agent/engine.go index 6f48e92..b4a924c 100644 --- a/internal/agent/engine.go +++ b/internal/agent/engine.go @@ -1,50 +1,46 @@ package agent import ( + "context" + + "github.com/cloudwego/eino/adk" + einoTool "github.com/cloudwego/eino/components/tool" + "solvify-agent/internal/observability" + "solvify-agent/internal/repository" "solvify-agent/internal/tool" "solvify-agent/pkg/config" ) -type KnowledgeSearchFactory func(userID string, kbIDs []string) *tool.KnowledgeSearchTool - -type GrepChunksFactory func(userID string, kbIDs []string) *tool.GrepChunksTool - -type GetDocumentInfoFactory func(userID string) *tool.GetDocumentInfoTool +// ToolBuildFn 内置工具的构建函数 +// 每个工具从请求里取它需要的参数(userID、kbIDs),返回完整可用的 tool 实例 +type ToolBuildFn func(ctx context.Context, userID string, kbIDs []string) einoTool.BaseTool -type ListKnowledgeChunksFactory func(userID string, kbIDs []string) *tool.ListKnowledgeChunksTool - -type ListKnowledgeBasesFactory func(userID string) *tool.ListKnowledgeBasesTool +// internalToolRegistryEntry 内置工具注册表项 +type internalToolRegistryEntry struct { + Name string // 工具名(用于 prompt 里标记、switch 里分类) + Order int // prompt 里的展示顺序(从小到大) + Dangerous bool // 危险工具标记 → prompt 里加 ⚠️ 和审批说明 + Build ToolBuildFn // 构建函数 +} type Engine struct { - knowledgeSearchFactory KnowledgeSearchFactory - grepChunksFactory GrepChunksFactory - getDocumentInfoFactory GetDocumentInfoFactory - listKnowledgeChunksFactory ListKnowledgeChunksFactory - listKnowledgeBasesFactory ListKnowledgeBasesFactory - toolFactory tool.ToolFactory - cfg config.AgentConfig - obs observability.Recorder + internalTools []internalToolRegistryEntry + toolFactory tool.ToolFactory + cfg config.AgentConfig + obs observability.Recorder + checkpointRepo repository.AgentCheckpointRepo } +// NewEngine 只收通用依赖。内置工具通过 RegisterInternal 注册。 func NewEngine( - knowledgeSearchFactory KnowledgeSearchFactory, - grepChunksFactory GrepChunksFactory, - getDocumentInfoFactory GetDocumentInfoFactory, - listKnowledgeChunksFactory ListKnowledgeChunksFactory, - listKnowledgeBasesFactory ListKnowledgeBasesFactory, toolFactory tool.ToolFactory, cfg config.AgentConfig, obs ...observability.Recorder, ) *Engine { e := &Engine{ - knowledgeSearchFactory: knowledgeSearchFactory, - grepChunksFactory: grepChunksFactory, - getDocumentInfoFactory: getDocumentInfoFactory, - listKnowledgeChunksFactory: listKnowledgeChunksFactory, - listKnowledgeBasesFactory: listKnowledgeBasesFactory, - toolFactory: toolFactory, - cfg: cfg, + toolFactory: toolFactory, + cfg: cfg, } if len(obs) > 0 && obs[0] != nil { e.obs = obs[0] @@ -52,6 +48,42 @@ func NewEngine( return e } +// RegisterInternal 注册一个内置工具。 +// order 决定在 prompt "可用工具" 段里的展示顺序(建议:检索类靠前,危险类靠后)。 +// dangerous=true 时 prompt 会额外追加危险工具审批说明。 +func (e *Engine) RegisterInternal(name string, order int, dangerous bool, build ToolBuildFn) { + e.internalTools = append(e.internalTools, internalToolRegistryEntry{ + Name: name, + Order: order, + Dangerous: dangerous, + Build: build, + }) +} + func (e *Engine) WithObservability(obs observability.Recorder) { e.obs = obs } + +func (e *Engine) WithCheckpointRepo(repo repository.AgentCheckpointRepo) { + e.checkpointRepo = repo +} + +// buildCheckpointStore 根据 Engine 配置构造 CheckPointStore。 +func (e *Engine) buildCheckpointStore(sessionID string) adk.CheckPointStore { + if e.checkpointRepo != nil && sessionID != "" { + return NewDBCheckPointStore(e.checkpointRepo, sessionID, CheckpointTTL) + } + return NewInMemoryCheckPointStore() +} + +// dangerousToolNames 返回所有标记为 dangerous 的内置工具名集合。 +// 用于构建审批中间件。 +func (e *Engine) dangerousToolNames() map[string]bool { + m := make(map[string]bool, len(e.internalTools)) + for _, entry := range e.internalTools { + if entry.Dangerous { + m[entry.Name] = true + } + } + return m +} diff --git a/internal/agent/engine_tools.go b/internal/agent/engine_tools.go index ab76a14..1d554b7 100644 --- a/internal/agent/engine_tools.go +++ b/internal/agent/engine_tools.go @@ -3,21 +3,19 @@ package agent import ( "context" "encoding/json" + "sort" einoTool "github.com/cloudwego/eino/components/tool" "solvify-agent/pkg/tokenutil" ) -// prebuiltToolsCtxKey 用作 context.WithValue 键,把深度模式入口"预构建好的工具集 + tokens 信息" -// 传到 runAgent 里,避免 runAgent 再调一次 factories 重复构建,同时保证 initContext 扣减的 -// ToolsTokens 和实际发给模型的工具定义完全一致(P0-④ 的根保证)。 type prebuiltToolsCtxKeyType struct{} var prebuiltToolsCtxKey = prebuiltToolsCtxKeyType{} type prebuiltToolsBundle struct { - Tools []einoTool.BaseTool + Tools []einoTool.BaseTool TotalTokens int } @@ -34,29 +32,20 @@ func prebuiltToolsFromContext(ctx context.Context) (prebuiltToolsBundle, bool) { return b, ok } -// buildAllTools 集中封装一次工具构建流程:知识库 5 个内置工具 + 用户工具。 -// 与 runAgent 里原有的构建顺序保持完全一致,避免"预构建版本 vs runAgent 版本"两套逻辑漂移。 +// buildAllTools 用 registry 构建所有内置工具 + 用户配置工具 func (e *Engine) buildAllTools(ctx context.Context, userID string, kbIDs []string) ([]einoTool.BaseTool, error) { - ksTool := e.knowledgeSearchFactory(userID, kbIDs) - grepTool := e.grepChunksFactory(userID, kbIDs) - docInfoTool := e.getDocumentInfoFactory(userID) - listChunksTool := e.listKnowledgeChunksFactory(userID, kbIDs) - listBasesTool := e.listKnowledgeBasesFactory(userID) - userTools := e.toolFactory.CreateAgentTools(ctx, userID) - - allTools := make([]einoTool.BaseTool, 0, 5+len(userTools)) - allTools = append(allTools, ksTool) - allTools = append(allTools, grepTool) - allTools = append(allTools, docInfoTool) - allTools = append(allTools, listChunksTool) - allTools = append(allTools, listBasesTool) - allTools = append(allTools, userTools...) + var allTools []einoTool.BaseTool + sorted := make([]internalToolRegistryEntry, len(e.internalTools)) + copy(sorted, e.internalTools) + sort.Slice(sorted, func(i, j int) bool { return sorted[i].Order < sorted[j].Order }) + for _, entry := range sorted { + allTools = append(allTools, entry.Build(ctx, userID, kbIDs)) + } + allTools = append(allTools, e.toolFactory.CreateAgentTools(ctx, userID)...) return allTools, nil } // EstimateToolsTokens 返回「工具定义的真 BPE token 数」以及预构建好的工具集。 -// 调用方用返回的 ToolsTokens 喂给 initContext(calculateContextBudgets 先扣),再把 prebuilt 放进 ctx, -// 之后 Execute 会复用这份工具集,保证前后 token 计算一致。 func (e *Engine) EstimateToolsTokens(ctx context.Context, userID string, kbIDs []string, modelName string) (int, context.Context, error) { tools, err := e.buildAllTools(ctx, userID, kbIDs) if err != nil { @@ -68,17 +57,13 @@ func (e *Engine) EstimateToolsTokens(ctx context.Context, userID string, kbIDs [ if err != nil || info == nil { continue } - // ToolInfo.MarshalJSON 已经包含了 Name/Desc/Extra/ParamsOneOf → 传给模型的完整定义 bs, mErr := json.Marshal(info) if mErr != nil { - // 实在 marshal 失败,就退化成 desc+name 粗略估算 total += tokenutil.CountTokens(info.Name+"\n"+info.Desc, modelName) continue } total += tokenutil.CountTokens(string(bs), modelName) } - // ReAct Agent 在 system prompt 里还会加一段"可用工具一览 / How to use tools"的指令, - // 经验上 ≈ tools_tokens 的 15%,保守取 20%。 overhead := int(float64(total) * 0.2) if overhead < 200 { overhead = 200 diff --git a/internal/agent/execute.go b/internal/agent/execute.go index 617344c..1ea831e 100644 --- a/internal/agent/execute.go +++ b/internal/agent/execute.go @@ -4,19 +4,18 @@ import ( "context" "encoding/json" "fmt" - "io" + "sort" "strings" "sync" "time" + "github.com/bytedance/sonic" + "github.com/cloudwego/eino/adk" "github.com/cloudwego/eino/components/model" einoTool "github.com/cloudwego/eino/components/tool" "github.com/cloudwego/eino/compose" - einoAgent "github.com/cloudwego/eino/flow/agent" - "github.com/cloudwego/eino/flow/agent/react" "github.com/cloudwego/eino/schema" - "solvify-agent/internal/model/dto/response" "solvify-agent/internal/model/entity" "solvify-agent/internal/observability" "solvify-agent/internal/tool" @@ -51,6 +50,17 @@ type agentStepPending struct { StartedAt time.Time } +// isInternalToolName 判断工具名是否为内置工具 +// registry 里注册过的就是内置,否则是用户配置的 +func (e *Engine) isInternalToolName(name string) bool { + for _, entry := range e.internalTools { + if entry.Name == name { + return true + } + } + return false +} + func (e *Engine) runAgent(ctx context.Context, req Request, chatModel model.ToolCallingChatModel, eventCh chan<- Event) { obsOk := e.obs != nil var tracker *agentStepTracker @@ -66,31 +76,22 @@ func (e *Engine) runAgent(ctx context.Context, req Request, chatModel model.Tool e.obs.Incr(ctx, "agent_engine_runs_total", nil, 1) } + // ── 构建工具列表:内置 registry + 用户配置 ── var allTools []einoTool.BaseTool - if pre, ok := prebuiltToolsFromContext(ctx); ok && len(pre.Tools) > 0 { - // 走深度模式入口预构建分支:工具集 + toolsTokens 已经提前扣好 - allTools = pre.Tools - } else { - // 回退分支(如 Execute 被直接调用、预构建失败):按原逻辑现场构建 - ksTool := e.knowledgeSearchFactory(req.UserID, req.KnowledgeBaseIDs) - grepTool := e.grepChunksFactory(req.UserID, req.KnowledgeBaseIDs) - docInfoTool := e.getDocumentInfoFactory(req.UserID) - listChunksTool := e.listKnowledgeChunksFactory(req.UserID, req.KnowledgeBaseIDs) - listBasesTool := e.listKnowledgeBasesFactory(req.UserID) - userTools := e.toolFactory.CreateAgentTools(ctx, req.UserID) - allTools = make([]einoTool.BaseTool, 0, 5+len(userTools)) - allTools = append(allTools, ksTool) - allTools = append(allTools, grepTool) - allTools = append(allTools, docInfoTool) - allTools = append(allTools, listChunksTool) - allTools = append(allTools, listBasesTool) - allTools = append(allTools, userTools...) + // 内置工具按 Order 排序后逐个 Build + sorted := make([]internalToolRegistryEntry, len(e.internalTools)) + copy(sorted, e.internalTools) + sort.Slice(sorted, func(i, j int) bool { return sorted[i].Order < sorted[j].Order }) + for _, entry := range sorted { + allTools = append(allTools, entry.Build(ctx, req.UserID, req.KnowledgeBaseIDs)) } - userTools := e.toolFactory.CreateAgentTools(ctx, req.UserID) // 仅用于下面日志中"用户工具数"展示 - _ = userTools + // 用户配置的工具 + userTools := e.toolFactory.CreateAgentTools(ctx, req.UserID) + allTools = append(allTools, userTools...) + // ── 工具统计 + 日志 ── toolDescMap := make(map[string]string, len(allTools)) userToolsN := 0 for _, t := range allTools { @@ -99,24 +100,20 @@ func (e *Engine) runAgent(ctx context.Context, req Request, chatModel model.Tool logger.Warnf("[Agent] 获取工具信息失败: %v", err) continue } - // heuristics: 内置 5 个工具名按前缀匹配,剩下的记作用户工具;日志只用于展示,不影响执行 - switch info.Name { - case "knowledge_search", "grep_chunks", "get_document_info", "list_knowledge_chunks", "list_knowledge_bases": - default: + if !e.isInternalToolName(info.Name) { userToolsN++ } toolDescMap[info.Name] = info.Desc logger.Infof("[Agent] 工具: name=%s, desc=%s", info.Name, truncateStr(info.Desc, 80)) } - logger.Infof("[Agent] userID=%s, 工具总数=%d (内置5个 + %d 用户工具)", req.UserID, len(allTools), userToolsN) - if userToolsN < 0 { - userToolsN = 0 - } + logger.Infof("[Agent] userID=%s, 工具总数=%d (内置=%d + 用户工具=%d)", + req.UserID, len(allTools), len(e.internalTools), userToolsN) if userToolsN == 0 { logger.Warnf("[Agent] 未加载到用户配置的工具(如联网搜索),请检查用户工具配置是否已启用") } - baseSystemPrompt := buildReActSystemPrompt(ctx, userTools) + // ── 系统提示词 ── + baseSystemPrompt := buildReActSystemPrompt(ctx, allTools, sorted) var ksToolForStream *tool.KnowledgeSearchTool for _, t := range allTools { if k, ok := t.(*tool.KnowledgeSearchTool); ok { @@ -126,244 +123,113 @@ func (e *Engine) runAgent(ctx context.Context, req Request, chatModel model.Tool } var systemPromptFinal string if req.SystemPrompt != "" { - // BuildSystem() 以空 baseSystem 构建时,结果会前导 "\n\n",去掉避免多余空行 enhanced := strings.TrimLeft(req.SystemPrompt, "\n") systemPromptFinal = baseSystemPrompt + "\n\n" + enhanced } else { - // 兜底:req.SystemPrompt 为空时只用 ReAct 规则(正常流程不会走到这里, - // PromptBuilder.BuildSystem() 总会产出摘要/记忆/用户信息之一) systemPromptFinal = baseSystemPrompt } logger.Infof("[Agent] SystemPrompt (前400字符): %s", truncateStr(systemPromptFinal, 400)) - inputMessages := buildInputMessages(req.Query, req.History) - // 运行 Agent 流程 - { - maxStep := e.cfg.MaxIterations - if maxStep <= 0 { - maxStep = 5 - } - - ag, err := react.NewAgent(ctx, &react.AgentConfig{ - ToolCallingModel: chatModel, - ToolsConfig: compose.ToolsNodeConfig{ - Tools: allTools, - }, - MaxStep: maxStep, - MessageModifier: func(_ context.Context, msgs []*schema.Message) []*schema.Message { - return append([]*schema.Message{schema.SystemMessage(systemPromptFinal)}, msgs...) - }, - }) - if err != nil { - logger.Errorf("Agent 初始化失败: %v", err) - if obsOk { - e.obs.Incr(ctx, "agent_engine_errors_total", map[string]string{"stage": "init"}, 1) - } - eventCh <- Event{ - Type: EventError, - Title: "深度模式启动失败", - Detail: "请尝试切换到快速模式,或稍后重试", - Error: err.Error(), - Status: "error", - Retryable: true, - Done: true, - } - return - } - - callbackHandler := newAgentCallbackHandler(eventCh, req.KnowledgeBaseIDs, toolDescMap) - callbackHandler.taskID = taskID - callbackHandler.tracker = tracker - callbackHandler.obs = e.obs - - stream, err := ag.Stream(ctx, inputMessages, einoAgent.WithComposeOptions(compose.WithCallbacks(callbackHandler.Handler()))) - if err != nil { - logger.Errorf("Agent 调用失败: %v", err) - if obsOk { - e.obs.Incr(ctx, "agent_engine_errors_total", map[string]string{"stage": "stream"}, 1) - } - errMsg := err.Error() - - if isToolChoiceUnsupportedError(errMsg) { - eventCh <- Event{ - Type: EventError, - Title: "当前模型不支持工具调用", - Detail: "该模型不支持工具调用功能,无法使用联网搜索、天气查询等工具。建议切换到支持工具调用的模型(如通义千问、智谱清言、DeepSeek 等),或使用快速模式。", - Error: errMsg, - Status: "error", - Retryable: false, - Done: true, - } - return - } - eventCh <- Event{ - Type: EventError, - Title: "深度推理失败", - Detail: "深度思考模式执行异常,请重试或使用快速模式", - Error: errMsg, - Status: "error", - Retryable: true, - Done: true, - } - return - } - - e.processStream(ctx, stream, ksToolForStream, eventCh) - } -} + inputMessages := buildInputMessages(req.Query, req.History) -func randomStr16() string { - const alpha = "0123456789abcdef" - out := make([]byte, 16) - seed := time.Now().UnixNano() - for i := range out { - seed = seed*1103515245 + 12345 - out[i] = alpha[int(seed>>16)&15] + maxStep := e.cfg.MaxIterations + if maxStep <= 0 { + maxStep = 5 } - return string(out) -} - -func (e *Engine) processStream(ctx context.Context, stream *schema.StreamReader[*schema.Message], ksTool *tool.KnowledgeSearchTool, eventCh chan<- Event) { - defer stream.Close() - - var fullAnswer string - for { - msg, err := stream.Recv() - if err == io.EOF { - break - } - if err != nil { - if ctx.Err() != nil { - logger.Infof("Agent 流被用户中断,已收集 %d 字符", len(fullAnswer)) - break - } - logger.Errorf("Agent 流读取失败: %v", err) - eventCh <- Event{ - Type: EventError, - Title: "推理过程中断", - Detail: "深度推理过程中断,请重试", - Error: err.Error(), - Status: "error", - Retryable: true, - Done: true, + // ── 创建 adk.ChatModelAgent ── + toolsNodeConfig := compose.ToolsNodeConfig{ + Tools: allTools, + + // 兜底:LLM 传了无效 JSON 时,返回 {} 让 InferTool 给业务函数零值 struct, + // 业务函数里会检查必填字段并返回 ToolResponse{Success:false}, + // LLM 看到错误信息后可以自行修正参数重试。 + // 不加这个的话 InferTool 内部 sonic.Unmarshal 失败会直接 return error → 整个 Agent 崩。 + ToolArgumentsHandler: func(ctx context.Context, toolName, arguments string) (string, error) { + if arguments == "" { + return "{}", nil } - return - } - if msg == nil { - continue - } - - if msg.Role != schema.Assistant { - continue - } - - if len(msg.ToolCalls) > 0 { - // ── 中间思考轮次(下一步还要调用工具)── - // 1) msg.Content 是 reasoning/推理思考,不能作为最终答案给用户看 - // 2) 只发 EventThinking 通知前端进度,不发 EventAnswer,不拼进 fullAnswer - if strings.TrimSpace(msg.Content) != "" { - thinking := truncateStr(msg.Content, 200) - eventCh <- Event{ - Type: EventThinking, - Title: "深度推理中", - Detail: thinking, - Status: "running", - } + var tmp map[string]any + if err := sonic.UnmarshalString(arguments, &tmp); err != nil { + logger.Warnf("[Agent] ToolArgumentsHandler: %s 参数 JSON 解析失败,已降级为空对象: raw=%q, err=%v", + toolName, truncateStr(arguments, 200), err) + return "{}", nil } - continue - } + return arguments, nil + }, - if msg.Content != "" { - // ── 最终答案轮次(没有下一步 ToolCalls,真正面向用户的正文)── - fullAnswer += msg.Content - eventCh <- Event{Type: EventAnswer, Content: msg.Content} - } + UnknownToolsHandler: func(ctx context.Context, name, input string) (string, error) { + logger.Warnf("[Agent] UnknownToolsHandler: LLM 调用了不存在的工具 %q,参数=%s", name, truncateStr(input, 200)) + return fmt.Sprintf("⚠️ 工具 %q 不存在,可用工具请查看系统提示。请检查工具名拼写后重试。", name), nil + }, } - - // ── 兜底:极端情况(每一轮都有 ToolCalls,MaxStep 到了还没出最终答案) - // 用知识库已命中的前 N 条来源拼一个总结,绝对不能把中间思考当答案发 - if strings.TrimSpace(fullAnswer) == "" && len(ksTool.CollectedSources) > 0 { - var sb strings.Builder - sb.WriteString("## 知识库检索结果总结\n\n") - sb.WriteString("根据当前检索到的内容,为您整理以下要点:\n\n") - usedTitles := make(map[string]bool, len(ksTool.CollectedSources)) - const maxTop = 5 - for i, src := range ksTool.CollectedSources { - if i >= maxTop { - break - } - title := src.Title - if title == "" { - title = "未命名文档" - } - // 同一个文档只拼一次摘要,重复 chunk 跳过 - if usedTitles[title] { - continue - } - usedTitles[title] = true - content := strings.TrimSpace(src.Content) - if len(content) > 160 { - content = content[:160] + "…" - } - chunkID := src.ID - if chunkID == "" { - chunkID = fmt.Sprintf("c%d", i) - } - sb.WriteString(fmt.Sprintf("- %s \n", title, title, chunkID)) - if content != "" { - sb.WriteString(fmt.Sprintf(" > %s\n\n", content)) - } + // 有危险工具时注入审批中间件 + if dangerousNames := e.dangerousToolNames(); len(dangerousNames) > 0 { + toolsNodeConfig.ToolCallMiddlewares = []compose.ToolMiddleware{ + {Invokable: buildDangerousToolMiddleware(dangerousNames)}, } - sb.WriteString("\n如需进一步分析请补充问题细节,或切换到快速模式获取更直接的回答。") - fullAnswer = sb.String() - eventCh <- Event{Type: EventAnswer, Content: fullAnswer} + logger.Infof("[Agent] 已注入危险工具审批中间件,工具列表=%v", dangerousNames) } - var sources []response.SourceInfo - type docInfo struct { - documentID string - knowledgeBaseID string - chunks []response.ChunkSource - } - docMap := make(map[string]*docInfo) - for _, doc := range ksTool.CollectedSources { - if _, exists := docMap[doc.Title]; !exists { - docMap[doc.Title] = &docInfo{ - documentID: doc.DocumentID, - knowledgeBaseID: doc.KnowledgeBaseID, - } + agent, err := adk.NewChatModelAgent(ctx, &adk.ChatModelAgentConfig{ + Name: "SolvifyDeepAgent", + Description: "深度模式 Agent,能够调用知识库和外部工具进行多步推理", + Instruction: systemPromptFinal, + Model: chatModel, + ToolsConfig: adk.ToolsConfig{ + ToolsNodeConfig: toolsNodeConfig, + }, + MaxIterations: maxStep, + }) + if err != nil { + logger.Errorf("[Agent] ChatModelAgent 初始化失败: %v", err) + if obsOk { + e.obs.Incr(ctx, "agent_engine_errors_total", map[string]string{"stage": "init"}, 1) } - docMap[doc.Title].chunks = append(docMap[doc.Title].chunks, response.ChunkSource{ - ID: doc.ID, - Content: doc.Content, - Score: doc.Score, - }) - } - for title, info := range docMap { - sources = append(sources, response.SourceInfo{ - DocumentID: info.documentID, - KnowledgeBaseID: info.knowledgeBaseID, - Title: title, - Chunks: info.chunks, - }) - } - - if strings.TrimSpace(fullAnswer) != "" { - eventCh <- Event{Type: EventThinking, Title: "正在生成答案", Status: "success"} + eventCh <- Event{ + Type: EventError, + Title: "深度模式启动失败", + Detail: "请尝试切换到快速模式,或稍后重试", + Error: err.Error(), + Status: "error", + Retryable: true, + Done: true, + } + return } - if len(sources) > 0 { - eventCh <- Event{Type: EventSources, Sources: sources} + // ── 创建 Runner ── + checkpointID := req.CheckpointID + if checkpointID == "" { + checkpointID = fmt.Sprintf("agent-%s-%s-%s-%d", req.SessionID, req.UserID, randomStr8(), time.Now().UnixNano()) + } else { + logger.Infof("[Agent] 恢复执行:复用 checkpointID=%s", checkpointID) } + store := e.buildCheckpointStore(req.SessionID) + runner := adk.NewRunner(ctx, adk.RunnerConfig{ + Agent: agent, + EnableStreaming: true, + CheckPointStore: store, + }) + + // ── 执行:首次 Run 或带 ResumeData 的 Resume ── + e.runWithRunner(ctx, runner, checkpointID, inputMessages, req, ksToolForStream, toolDescMap, eventCh, tracker, taskID) +} - eventCh <- Event{ - Type: EventDone, - Content: fullAnswer, - Sources: sources, +func randomStr(n int) string { + const alpha = "0123456789abcdef" + buf := make([]byte, n) + seed := time.Now().UnixNano() + for i := 0; i < n; i++ { + seed = seed*1103515245 + 12345 + buf[i] = alpha[int(seed>>16)&15] } + return string(buf) } +func randomStr16() string { return randomStr(16) } +func randomStr8() string { return randomStr(8) } + func buildInputMessages(query string, history []entity.ChatMessage) []*schema.Message { msgs := make([]*schema.Message, 0, len(history)+1) @@ -380,6 +246,13 @@ func buildInputMessages(query string, history []entity.ChatMessage) []*schema.Me return msgs } +func truncateStr(s string, maxLen int) string { + if len(s) <= maxLen { + return s + } + return s[:maxLen] + "..." +} + func extractQueryFromArgs(args string) string { var params struct { Query string `json:"query"` @@ -390,13 +263,6 @@ func extractQueryFromArgs(args string) string { return args } -func truncateStr(s string, maxLen int) string { - if len(s) <= maxLen { - return s - } - return s[:maxLen] + "..." -} - func isToolChoiceUnsupportedError(errMsg string) bool { lower := strings.ToLower(errMsg) keywords := []string{ diff --git a/internal/agent/prompt.go b/internal/agent/prompt.go index 196a16c..53fa0c4 100644 --- a/internal/agent/prompt.go +++ b/internal/agent/prompt.go @@ -8,71 +8,107 @@ import ( einoTool "github.com/cloudwego/eino/components/tool" ) -type toolDesc struct { - Name string - Desc string -} - -func buildReActSystemPrompt(ctx context.Context, userTools []einoTool.BaseTool) string { +// buildReActSystemPrompt 构建深度模式的系统提示词。 +// internalSorted: 内置工具,已按 Order 升序排列(prompt 用)。 +// allTools: 所有工具(内置 + 用户配置),用于动态解析 description。 +func buildReActSystemPrompt(ctx context.Context, allTools []einoTool.BaseTool, internalSorted []internalToolRegistryEntry) string { var sb strings.Builder - sb.WriteString("你是 Solvify 知识助理,专业的 AI 知识助手。\n\n") + sb.WriteString("你是 Solvify 知识助理,一个能调用工具解决问题的 AI 助手。\n\n") + // ── 引用规则 ── sb.WriteString("## 引用规则\n") - sb.WriteString("在句末插入引用标签,紧跟句子不放换行:\n") - sb.WriteString("- 知识库内容:\n") - sb.WriteString("- 网页内容:\n") - sb.WriteString("- 其他工具结果不需要引用标签\n") - sb.WriteString("- 禁止集中放在末尾,禁止编造 chunk_id,禁止直接复制原文\n\n") + sb.WriteString("- 知识库内容:在句末紧跟 \n") + sb.WriteString("- 网页/联网搜索内容:在句末紧跟 \n") + sb.WriteString("- 禁止集中放在文末,禁止编造 chunk_id,禁止直接复制原文大段\n\n") - toolDescs := resolveToolDescs(ctx, userTools) + // ── 可用工具(完全动态生成) ── + allDescs := resolveToolDescs(ctx, allTools) + descMap := make(map[string]string, len(allDescs)) + for _, td := range allDescs { + descMap[td.Name] = td.Desc + } + internalNames := make(map[string]bool, len(internalSorted)) + for _, entry := range internalSorted { + internalNames[entry.Name] = true + } sb.WriteString("## 可用工具\n") - sb.WriteString("- **knowledge_search**: 语义搜索知识库,优先用于查找信息\n") - sb.WriteString("- **grep_chunks**: 关键词精确匹配文档内容\n") - sb.WriteString("- **get_document_info**: 获取文档元数据(标题、类型、大小、分块数等)\n") - sb.WriteString("- **list_knowledge_chunks**: 列出知识库中的文档\n") - sb.WriteString("- **list_knowledge_bases**: 列出所有知识库\n") - for _, t := range toolDescs { - desc := t.Desc + for _, entry := range internalSorted { + desc := descMap[entry.Name] if desc == "" { - desc = "用户配置的工具" + desc = "(工具不可用)" + } + label := "" + if entry.Dangerous { + label = " ⚠️ 危险 · 执行前需人工审批" + } + sb.WriteString(fmt.Sprintf("- **%s**: %s%s\n", entry.Name, desc, label)) + } + // 用户配置的外部工具 + for _, td := range allDescs { + if !internalNames[td.Name] { + desc := td.Desc + if desc == "" { + desc = "用户配置的外部工具" + } + sb.WriteString(fmt.Sprintf("- **%s**: %s\n", td.Name, desc)) } - sb.WriteString(fmt.Sprintf("- **%s**: %s\n", t.Name, desc)) } sb.WriteString("\n") - sb.WriteString("## 调用原则(必须严格遵守)\n") - sb.WriteString("1. **必须先调用 knowledge_search 检索知识库**,即使你认为知道答案也要先检索,不能跳过\n") - sb.WriteString("2. 根据检索结果决定是否需要补充:知识库信息不足时,再调用其他工具\n") - sb.WriteString("3. 不重复调用同一工具,检索结果已足够时直接回答\n") - sb.WriteString("4. 工具调用总计不超过 3 次\n") - sb.WriteString("5. **强制收敛(非常重要)**:当达到最大推理轮次或已用完工具调用次数时,必须立即输出最终结论,禁止再规划工具调用、禁止写'我还需要查一下/让我补充一下/下一步应该'等思考内容,直接总结已有信息 + 知识库引用给出答案\n") - sb.WriteString("6. **答案分层**:存在 ToolCalls 时,这一轮 Message.Content 只写 1-2 句简短推理(不会展示给用户,也不会进入最终答案),只有 ToolCalls 为空的那一轮 Message.Content 才是完整、可读、面向最终用户的答案正文(含引用标签、Markdown 排版)\n") + // ── 工作原则 ── + sb.WriteString("## 工作原则\n") + sb.WriteString("1. **先检索再回答**:第一步始终是 knowledge_search。即使你认为知道答案也必须先检索知识库\n") + sb.WriteString("2. **按需补充**:知识库结果不足时,再调用其他工具(grep_chunks 精准查找、get_document_info 查元数据等)\n") + sb.WriteString("3. **不重复调用**:已获得足够信息时,直接给出答案,不要为了'再确认'重复调用\n") + sb.WriteString("4. **工具上限 3 次**:工具调用总数不超过 3 次,用完必须收敛\n") + sb.WriteString("5. **强制收敛**:达到最大推理轮次或用完工具次数时,立即总结已有信息给出最终答案。禁止再规划'下一步应该'、'我还需要'等思考性输出\n") + sb.WriteString("6. **答案分层**:有 ToolCalls 的轮次 Message.Content 只写 1-2 句简短推理(不会展示给用户);只有 ToolCalls 为空的轮次才是完整、可读、面向最终用户的答案正文\n") - if len(toolDescs) > 0 { - names := make([]string, len(toolDescs)) - for i, t := range toolDescs { - names[i] = t.Name + // 危险工具补充说明 + hasDangerous := false + for _, entry := range internalSorted { + if entry.Dangerous { + hasDangerous = true + break } - sb.WriteString(fmt.Sprintf("5. 知识库结果不足或需要最新信息时,调用 %s 联网搜索(最多 1 次)\n", strings.Join(names, " 或 "))) - sb.WriteString("6. **当 knowledge_search 返回空结果或结果不相关时,必须调用联网搜索工具**,不能直接用自身知识回答\n") - sb.WriteString("7. 用户明确要求'联网搜索'、'搜索网页'、'最新信息'等时,即使知识库有结果也应调用联网搜索工具获取最新信息\n") - } else { - sb.WriteString("5. 没有可用的联网搜索工具,只能使用知识库内容回答\n") + } + if hasDangerous { + sb.WriteString("7. **危险工具审批**:delete_document 等危险工具会在执行前暂停并等待用户审批,调用后流程中断,用户确认后自动继续\n") + sb.WriteString(" - ⚠️ **目标不明确先反问**:当用户说'删除那个文档'、'清理一下'、'把上面的删了'这类模糊指令,且从对话历史无法唯一确定目标时,**绝对不能编造参数调用工具**。先反问用户明确目标(例如:'你要删除的是《压力 - 07/13 16:03》那个文档吗?还是另一个?')\n") + sb.WriteString(" - ⚠️ **禁止猜测参数**:document_id 等关键参数必须来自可靠来源(用户明确提供、get_document_info 工具查询结果、历史对话中已确认的 ID)。严禁从模糊描述或'看起来像是'的文本中猜测或编造\n") + sb.WriteString(" - 调用危险工具时务必在参数里写清楚目标和原因,便于用户决策\n") + } + + // 外部联网工具 + externals := make([]string, 0) + for _, td := range allDescs { + if !internalNames[td.Name] { + externals = append(externals, td.Name) + } + } + if len(externals) > 0 { + sb.WriteString(fmt.Sprintf("8. **联网搜索**:当 knowledge_search 返回空结果或与问题不相关时,必须调用 %s 联网搜索。用户明确要求'联网'、'最新'时即使知识库有结果也应联网\n", + strings.Join(externals, " 或 "))) } sb.WriteString("\n") - sb.WriteString("**禁止**:不检索知识库直接用自身知识回答。第一步必须是 knowledge_search。\n") - sb.WriteString("**禁止**:知识库检索失败后不调用联网搜索工具而直接回答。\n\n") + // ── 禁止 ── + sb.WriteString("## 禁止\n") + sb.WriteString("- 不检索知识库就用自身知识回答\n") + sb.WriteString("- 知识库检索失败不联网搜索就直接回答\n") + sb.WriteString("- 在有可用工具时直接给用户'请去界面手动操作'的建议——你应该调用工具来完成\n\n") + + // ── 回答要求 ── sb.WriteString("## 回答要求\n") - sb.WriteString("- 使用 Markdown 格式,根据内容复杂度自适应排版\n") - sb.WriteString("- 关键信息和结论用 **加粗** 标注\n") - sb.WriteString("- 多个要点用列表(`-` 或 `1.`)组织,有顺序的用有序列表\n") - sb.WriteString("- 简单问题简洁回答,复杂问题用 `##` 标题分章节\n") + sb.WriteString("- 使用 Markdown,根据复杂度自适应排版\n") + sb.WriteString("- 关键结论和数字用 **加粗**\n") + sb.WriteString("- 要点用列表(`-` 或 `1.`)组织\n") + sb.WriteString("- 简单问题简洁回答,复杂问题用 `##` 分章节\n") sb.WriteString("- 用自己的话回答,不直接复制原文\n") - sb.WriteString("- 使用中文\n") - sb.WriteString("- 列表类问题直接呈现结果,不要提及工具或内部信息\n") + sb.WriteString("- 始终使用中文\n") + sb.WriteString("- 不要提及内部工具名称、推理步骤或'我调用了 XX 工具'这类实现细节\n") return sb.String() } @@ -92,3 +128,8 @@ func resolveToolDescs(ctx context.Context, tools []einoTool.BaseTool) []toolDesc } return descs } + +type toolDesc struct { + Name string + Desc string +} diff --git a/internal/agent/runner_adapter.go b/internal/agent/runner_adapter.go new file mode 100644 index 0000000..f411b19 --- /dev/null +++ b/internal/agent/runner_adapter.go @@ -0,0 +1,426 @@ +package agent + +import ( + "context" + "io" + "strings" + + "github.com/cloudwego/eino/adk" + einoTool "github.com/cloudwego/eino/components/tool" + "github.com/cloudwego/eino/schema" + + "solvify-agent/internal/model/dto/response" + "solvify-agent/internal/tool" + "solvify-agent/pkg/logger" +) + +// runWithRunner 用 adk.Runner 执行 Agent,产出 AgentEvent 并转换到 eventCh。 +// 首次执行调 runner.Run,带 ResumeData 时调 runner.ResumeWithParams。 +func (e *Engine) runWithRunner( + ctx context.Context, + runner *adk.Runner, + checkpointID string, + inputMessages []*schema.Message, + req Request, + ksTool *tool.KnowledgeSearchTool, + toolDescMap map[string]string, + eventCh chan<- Event, + tracker *agentStepTracker, + taskID string, +) { + var ( + iter *adk.AsyncIterator[*adk.AgentEvent] + err error + ) + + // ── 首次执行 vs 恢复执行 ── + if len(req.ResumeData) > 0 { + logger.Infof("[Agent] 恢复执行: checkpointID=%s, targets=%v", checkpointID, mapKeys(req.ResumeData)) + iter, err = runner.ResumeWithParams(ctx, checkpointID, &adk.ResumeParams{ + Targets: req.ResumeData, + }) + if err != nil { + logger.Errorf("[Agent] ResumeWithParams 失败: %v", err) + eventCh <- Event{ + Type: EventError, + Title: "恢复执行失败", + Detail: "无法从中断点恢复,请重新发起深度模式请求", + Error: err.Error(), + Status: "error", + Retryable: false, + Done: true, + } + return + } + } else { + iter = runner.Run(ctx, inputMessages, adk.WithCheckPointID(checkpointID)) + } + + var fullAnswer strings.Builder + var interruptSent bool + + for { + agentEvent, ok := iter.Next() + if !ok { + break + } + + if agentEvent.Err != nil { + if ctx.Err() != nil { + logger.Infof("[Agent] Runner 迭代器因 ctx 取消结束") + break + } + logger.Errorf("[Agent] Runner 事件错误: %v", agentEvent.Err) + if isToolChoiceUnsupportedError(agentEvent.Err.Error()) { + eventCh <- Event{ + Type: EventError, + Title: "当前模型不支持工具调用", + Detail: "该模型不支持工具调用功能,无法使用联网搜索、天气查询等工具。建议切换到支持工具调用的模型(如通义千问、智谱清言、DeepSeek 等),或使用快速模式。", + Error: agentEvent.Err.Error(), + Status: "error", + Retryable: false, + Done: true, + } + return + } + eventCh <- Event{ + Type: EventError, + Title: "深度推理失败", + Detail: "深度思考模式执行异常,请重试或使用快速模式", + Error: agentEvent.Err.Error(), + Status: "error", + Retryable: true, + Done: true, + } + return + } + + // ── Interrupt 处理 ── + if agentEvent.Action != nil && agentEvent.Action.Interrupted != nil { + if !interruptSent { + interruptSent = true + ii := agentEvent.Action.Interrupted + interruptCtx := ii.InterruptContexts + var interruptID string + var interruptInfo any + if len(interruptCtx) > 0 { + interruptID = interruptCtx[0].ID + interruptInfo = interruptCtx[0].Info + } + logger.Infof("[Agent] 执行中断,等待用户审批: checkpointID=%s, interruptID=%s, info=%v", checkpointID, interruptID, interruptInfo) + infoMap, _ := interruptInfo.(map[string]any) + eventCh <- Event{ + Type: EventInterrupt, + Title: "需要人工确认", + Detail: truncateStr(formatInterruptInfo(interruptInfo), 256), + Status: "interrupt", + Error: interruptID, + CheckpointID: checkpointID, + InterruptID: interruptID, + InterruptInfo: infoMap, + Done: true, + } + return + } + continue + } + + if agentEvent.Output == nil || agentEvent.Output.MessageOutput == nil { + continue + } + + mv := agentEvent.Output.MessageOutput + + // ── 流式:消费 MessageStream,逐 chunk 处理 ── + if mv.IsStreaming && mv.MessageStream != nil { + e.consumeMessageStream(ctx, mv, toolDescMap, &fullAnswer, eventCh) + continue + } + + // ── 非流式:直接拿 Message ── + msg, err := mv.GetMessage() + if err != nil || msg == nil { + continue + } + + e.handleMessage(ctx, msg, mv.Role, mv.ToolName, toolDescMap, &fullAnswer, eventCh) + } + + // ── 兜底:没拿到 ToolCalls 也没拿到最终答案,但 KB 有结果 ── + if strings.TrimSpace(fullAnswer.String()) == "" && ksTool != nil && len(ksTool.CollectedSources) > 0 { + fallback := buildFallbackAnswer(ksTool.CollectedSources) + fullAnswer.WriteString(fallback) + eventCh <- Event{Type: EventAnswer, Content: fallback} + } + + // ── 收集 Sources ── + var sources []response.SourceInfo + if ksTool != nil { + sources = collectSources(ksTool.CollectedSources) + } + + if strings.TrimSpace(fullAnswer.String()) != "" { + eventCh <- Event{Type: EventThinking, Title: "正在生成答案", Status: "success"} + } + if len(sources) > 0 { + eventCh <- Event{Type: EventSources, Sources: sources} + } + + // observability + if tracker != nil && taskID != "" && e.obs != nil { + // Runner 已经在内部处理了完整的 step tracker,这里兜底留空即可 + } + + eventCh <- Event{ + Type: EventDone, + Content: fullAnswer.String(), + Sources: sources, + } +} + +func (e *Engine) consumeMessageStream( + ctx context.Context, + mv *adk.TypedMessageVariant[*schema.Message], + toolDescMap map[string]string, + fullAnswer *strings.Builder, + eventCh chan<- Event, +) { + stream := mv.MessageStream + defer stream.Close() + + var toolCallPending bool + var toolCallName string + + for { + msg, err := stream.Recv() + if err == io.EOF { + break + } + if err != nil { + if ctx.Err() != nil { + logger.Infof("[Agent] MessageStream 因 ctx 取消结束") + return + } + logger.Errorf("[Agent] MessageStream 读取失败: %v", err) + return + } + if msg == nil { + continue + } + + // Role=Tool 的流:发 tool_result + if mv.Role == schema.Tool && msg.Role == schema.Assistant && toolCallPending { + fullAnswer.WriteString(msg.Content) + continue + } + + // Role=Assistant 流 + if mv.Role == schema.Assistant { + if len(msg.ToolCalls) > 0 { + // 流式里 ToolCalls 可能分 chunk 到达,攒齐了再发 EventToolCall + for _, tc := range msg.ToolCalls { + if tc.Function.Name != "" { + toolCallName = tc.Function.Name + eventCh <- Event{ + Type: EventToolCall, + Title: "调用工具", + Detail: truncateStr(tc.Function.Arguments, 200), + Status: "running", + } + toolCallPending = true + } + } + if strings.TrimSpace(msg.Content) != "" { + eventCh <- Event{ + Type: EventThinking, + Title: "深度推理中", + Detail: truncateStr(msg.Content, 200), + Status: "running", + } + } + continue + } + + // 最终答案 + if msg.Content != "" { + fullAnswer.WriteString(msg.Content) + eventCh <- Event{Type: EventAnswer, Content: msg.Content} + } + continue + } + + // Role=Tool 完整结果 + if mv.Role == schema.Tool && msg.Content != "" { + title, detail, _ := formatToolEnd(toolCallName, &einoTool.CallbackOutput{Response: msg.Content}, toolDescMap) + eventCh <- Event{ + Type: EventToolResult, + Title: title, + Detail: detail, + Status: "success", + ToolResult: msg.Content, + } + toolCallPending = false + } + } +} + +func (e *Engine) handleMessage( + ctx context.Context, + msg *schema.Message, + role schema.RoleType, + toolName string, + toolDescMap map[string]string, + fullAnswer *strings.Builder, + eventCh chan<- Event, +) { + switch role { + case schema.Assistant: + if len(msg.ToolCalls) > 0 { + for _, tc := range msg.ToolCalls { + if tc.Function.Name == "" { + continue + } + title, detail := formatToolStart(tc.Function.Name, extractQueryFromArgs(tc.Function.Arguments), nil, toolDescMap) + eventCh <- Event{ + Type: EventToolCall, + Title: title, + Detail: detail, + Status: "running", + } + } + if strings.TrimSpace(msg.Content) != "" { + eventCh <- Event{ + Type: EventThinking, + Title: "深度推理中", + Detail: truncateStr(msg.Content, 200), + Status: "running", + } + } + return + } + if msg.Content != "" { + fullAnswer.WriteString(msg.Content) + eventCh <- Event{Type: EventAnswer, Content: msg.Content} + } + + case schema.Tool: + if msg.Content != "" { + title, detail, _ := formatToolEnd(toolName, &einoTool.CallbackOutput{Response: msg.Content}, toolDescMap) + eventCh <- Event{ + Type: EventToolResult, + Title: title, + Detail: detail, + Status: "success", + ToolResult: msg.Content, + } + } + } +} + +func formatInterruptInfo(info any) string { + switch v := info.(type) { + case string: + return v + case map[string]any: + if msg, ok := v["message"].(string); ok && msg != "" { + return msg + } + if req, ok := v["request"].(string); ok && req != "" { + return req + } + } + return "执行被中断,等待用户处理" +} + +func mapKeys(m map[string]any) []string { + if len(m) == 0 { + return nil + } + out := make([]string, 0, len(m)) + for k := range m { + out = append(out, k) + } + return out +} + +// collectSources 从 KnowledgeSearchTool.CollectedSources 转换成 response.SourceInfo +func collectSources(sources []tool.SourceDocument) []response.SourceInfo { + if len(sources) == 0 { + return nil + } + type docInfo struct { + documentID string + knowledgeBaseID string + chunks []response.ChunkSource + } + docMap := make(map[string]*docInfo) + for _, src := range sources { + if _, exists := docMap[src.Title]; !exists { + docMap[src.Title] = &docInfo{ + documentID: src.DocumentID, + knowledgeBaseID: src.KnowledgeBaseID, + } + } + docMap[src.Title].chunks = append(docMap[src.Title].chunks, response.ChunkSource{ + ID: src.ID, + Content: src.Content, + Score: src.Score, + }) + } + result := make([]response.SourceInfo, 0, len(docMap)) + for title, info := range docMap { + result = append(result, response.SourceInfo{ + DocumentID: info.documentID, + KnowledgeBaseID: info.knowledgeBaseID, + Title: title, + Chunks: info.chunks, + }) + } + return result +} + +// buildFallbackAnswer 当 Agent 没有产出最终答案但 KB 有命中时的兜底总结 +func buildFallbackAnswer(sources []tool.SourceDocument) string { + var sb strings.Builder + sb.WriteString("## 知识库检索结果总结\n\n") + sb.WriteString("根据当前检索到的内容,为您整理以下要点:\n\n") + usedTitles := make(map[string]bool, len(sources)) + const maxTop = 5 + for i, src := range sources { + if i >= maxTop { + break + } + title := src.Title + if title == "" { + title = "未命名文档" + } + if usedTitles[title] { + continue + } + usedTitles[title] = true + content := strings.TrimSpace(src.Content) + if len(content) > 160 { + content = content[:160] + "…" + } + chunkID := src.ID + if chunkID == "" { + chunkID = "c" + string(rune('0'+i)) + } + sb.WriteString("- ") + sb.WriteString(title) + sb.WriteString(" \n") + if content != "" { + sb.WriteString(" > ") + sb.WriteString(content) + sb.WriteString("\n\n") + } + } + sb.WriteString("\n如需进一步分析请补充问题细节,或切换到快速模式获取更直接的回答。") + return sb.String() +} + + diff --git a/internal/agent/tool_middleware.go b/internal/agent/tool_middleware.go new file mode 100644 index 0000000..7325915 --- /dev/null +++ b/internal/agent/tool_middleware.go @@ -0,0 +1,92 @@ +package agent + +import ( + "context" + "encoding/gob" + "fmt" + + "github.com/cloudwego/eino/compose" + + "solvify-agent/pkg/logger" +) + +// DangerousToolState 审批中间件持久化到 checkpoint 的状态 +type DangerousToolState struct { + ToolName string `json:"tool_name"` + Arguments string `json:"arguments"` +} + +func init() { + // gob 序列化 checkpoint 时需要能识别 DangerousToolState 这个 interface 实现类型 + // 只注册值类型,避免同类型值/指针重复注册导致 panic + gob.Register(DangerousToolState{}) +} + +// buildDangerousToolMiddleware 构建统一的危险工具审批中间件。 +// 所有 dangerousNames 里列出的工具名,在实际执行前都会被拦截: +// 1. 首次执行 → StatefulInterrupt 暂停,checkpoint 保存当前参数 +// 2. 用户审批后恢复 → GetResumeContext 拿审批结果 +// - "approve"/"同意"/"确认" → 放行 next() 执行真实业务逻辑 +// - 其他 → 拒绝,返回"操作被取消" +// +// 工具本身(如 DeleteDocumentTool)不再需要写任何 Interrupt/Resume 代码, +// 只关心自己的业务逻辑即可。 +func buildDangerousToolMiddleware(dangerousNames map[string]bool) compose.InvokableToolMiddleware { + return func(next compose.InvokableToolEndpoint) compose.InvokableToolEndpoint { + return func(ctx context.Context, input *compose.ToolInput) (*compose.ToolOutput, error) { + // 不是危险工具 → 直接放行 + if !dangerousNames[input.Name] { + return next(ctx, input) + } + + // ── 检查是否从上次中断恢复 ── + wasInterrupted, hasState, state := compose.GetInterruptState[DangerousToolState](ctx) + + if !wasInterrupted { + // ── 首次执行:中断等待审批 ── + // info 用 string(gob 原生类型,不需要额外注册);不要用 map[string]any 这种 gob 不认识的类型 + // state 用 DangerousToolState(init 里已经 gob.Register 过) + info := fmt.Sprintf("即将执行危险工具 %s,请确认是否继续", input.Name) + logger.Infof("[ToolMiddleware] 危险工具 %s 触发审批中断, args=%s", input.Name, truncateStr(input.Arguments, 200)) + return nil, compose.StatefulInterrupt(ctx, info, DangerousToolState{ + ToolName: input.Name, + Arguments: input.Arguments, + }) + } + + // ── 恢复执行:拿审批结果 ── + isResumeFlow, hasData, approvalResult := compose.GetResumeContext[string](ctx) + if !isResumeFlow || !hasData { + logger.Warnf("[ToolMiddleware] 恢复流程异常:wasInterrupted=%v, isResumeFlow=%v, hasData=%v", + wasInterrupted, isResumeFlow, hasData) + return &compose.ToolOutput{Result: "恢复流程异常:未收到审批结果"}, nil + } + + // 用 state 里的原始参数(防恢复时 LLM 重新生成导致不一致) + // 但 input.Arguments 在 resume 时也是正确的 checkpoint 回放值,两者一致 + if hasState && state.ToolName != "" { + logger.Infof("[ToolMiddleware] 恢复执行: tool=%s, approval=%q (state.tool=%s)", + input.Name, approvalResult, state.ToolName) + } else { + logger.Infof("[ToolMiddleware] 恢复执行: tool=%s, approval=%q", input.Name, approvalResult) + } + + switch approvalResult { + case "approve", "同意", "确认", "yes", "y", "ok": + return next(ctx, input) // 放行真实业务逻辑 + + case "reject", "拒绝", "取消", "no", "n": + logger.Infof("[ToolMiddleware] 用户拒绝执行危险工具 %s", input.Name) + return &compose.ToolOutput{ + Result: fmt.Sprintf("❌ 操作被用户拒绝:%s 未执行", input.Name), + }, nil + + default: + logger.Warnf("[ToolMiddleware] 审批结果 %q 无法识别,默认拒绝 %s", approvalResult, input.Name) + return &compose.ToolOutput{ + Result: fmt.Sprintf("⚠️ 审批结果 %q 无法识别,默认不执行 %s", approvalResult, input.Name), + }, nil + } + } + } +} diff --git a/internal/agent/types.go b/internal/agent/types.go index a73c45f..967df5d 100644 --- a/internal/agent/types.go +++ b/internal/agent/types.go @@ -24,6 +24,8 @@ type PromptUserContext struct { // Request 描述 Agent 执行请求 type Request struct { + CheckpointID string // 恢复执行时使用的 checkpointID(恢复场景必填) + SessionID string // 会话 ID(用于 checkpoint 关联) UserID string // 用户 ID(用于知识库检索权限) Query string // 原始用户问题 History []entity.ChatMessage // 历史对话 @@ -34,6 +36,7 @@ type Request struct { Memories []entity.UserMemory // 用户长期记忆 — 保留给调试/日志 UserCtx PromptUserContext // 用户基本信息 + 当前时间 — 保留给调试/日志,System Prompt 注入统一走 SystemPrompt 字段 SystemPrompt string // 统一入口注入的完整 System Prompt。不为空时 runAgent 完全信任它,不再内部二次拼接摘要/记忆 + ResumeData map[string]any // 恢复执行时的审批数据(key=interruptID, value=用户审批结果);为 nil 时首次执行 } // Event 描述 Agent SSE 事件 @@ -55,6 +58,10 @@ type Event struct { Retryable bool `json:"retryable,omitempty"` // ToolResult 工具调用结果(完整内容,供前端展示) ToolResult string `json:"tool_result,omitempty"` + // interrupt 事件字段 + CheckpointID string `json:"checkpoint_id,omitempty"` + InterruptID string `json:"interrupt_id,omitempty"` + InterruptInfo map[string]any `json:"interrupt_info,omitempty"` } // 事件类型常量 @@ -68,4 +75,5 @@ const ( EventCitation = "citation" // 单个引用(后端流式解析时实时发送) EventSources = "sources" // 来源信息 EventDone = "done" // 完成 + EventInterrupt = "interrupt" // 执行中断,等待用户审批确认 ) diff --git a/internal/app/app.go b/internal/app/app.go index ef774bd..639a70d 100644 --- a/internal/app/app.go +++ b/internal/app/app.go @@ -1,4 +1,4 @@ -package app +package app import ( "context" @@ -11,6 +11,7 @@ import ( "syscall" "time" + einoTool "github.com/cloudwego/eino/components/tool" "github.com/gin-gonic/gin" "github.com/prometheus/client_golang/prometheus" "github.com/redis/go-redis/v9" @@ -232,52 +233,44 @@ type AgentComponents struct { AgentEngine *agent.Engine } -// initAgentComponents 初始化 Agent 相关组件(Embedding、RAG、工具、Agent 引擎) +// initAgentComponents 初始化 Agent 相关组件(Embedding、RAG、工具注册、Agent 引擎) +// 内置工具全部通过 RegisterInternal 注册,Engine 不感知具体工具类型,新增内置工具只需要在这里多调一行 func (a *App) initAgentComponents(toolFactory tool.ToolFactory, documentRepo repository.DocumentRepository, chunkRepo repository.DocumentChunkRepository, kbRepo repository.KnowledgeBaseRepository) *AgentComponents { embeddingFunc := a.initEmbedding() vectorRetriever := a.initRetriever(embeddingFunc) - // knowledge_search 工厂:每次 Agent 请求创建带用户上下文的工具实例 - ksFactory := agent.KnowledgeSearchFactory(func(userID string, kbIDs []string) *tool.KnowledgeSearchTool { - return tool.NewKnowledgeSearchTool(vectorRetriever).WithContext(userID, kbIDs) - }) - - // grep_chunks 工厂:关键词精确匹配 - grepFactory := agent.GrepChunksFactory(func(userID string, kbIDs []string) *tool.GrepChunksTool { - return tool.NewGrepChunksTool(chunkRepo).WithContext(userID, kbIDs) - }) - - // get_document_info 工厂:文档元数据 - docInfoFactory := agent.GetDocumentInfoFactory(func(userID string) *tool.GetDocumentInfoTool { - return tool.NewGetDocumentInfoTool(documentRepo).WithContext(userID) - }) - - // list_knowledge_chunks 工厂:文档列表 - listChunksFactory := agent.ListKnowledgeChunksFactory(func(userID string, kbIDs []string) *tool.ListKnowledgeChunksTool { - return tool.NewListKnowledgeChunksTool(documentRepo).WithContext(userID, kbIDs) - }) - - // list_knowledge_bases 工厂:知识库列表 - listBasesFactory := agent.ListKnowledgeBasesFactory(func(userID string) *tool.ListKnowledgeBasesTool { - return tool.NewListKnowledgeBasesTool(kbRepo).WithContext(userID) - }) - - // 初始化 Agent Engine(eino ReAct Agent) - // 用户配置的工具(web_search 等)通过 ToolFactory 从 DB/Redis 动态加载 - agentEngine := agent.NewEngine( - ksFactory, - grepFactory, - docInfoFactory, - listChunksFactory, - listBasesFactory, - toolFactory, - a.cfg.Agent, - ) - // 阶段三:绑定可观测性 recorder + // ── 初始化 Agent Engine ── + agentEngine := agent.NewEngine(toolFactory, a.cfg.Agent) if a.obsRecorder != nil { agentEngine.WithObservability(a.obsRecorder) } + // ── 注册内置工具(按 Order 升序出现在 prompt "可用工具" 段) ── + agentEngine.RegisterInternal("knowledge_search", 1, false, + func(ctx context.Context, userID string, kbIDs []string) einoTool.BaseTool { + return tool.NewKnowledgeSearchTool(vectorRetriever).WithContext(userID, kbIDs) + }) + agentEngine.RegisterInternal("grep_chunks", 2, false, + func(ctx context.Context, userID string, kbIDs []string) einoTool.BaseTool { + return tool.NewGrepChunksTool(chunkRepo)(userID, kbIDs) + }) + agentEngine.RegisterInternal("get_document_info", 3, false, + func(ctx context.Context, userID string, kbIDs []string) einoTool.BaseTool { + return tool.NewGetDocumentInfoTool(documentRepo)(userID, kbIDs) + }) + agentEngine.RegisterInternal("list_knowledge_chunks", 4, false, + func(ctx context.Context, userID string, kbIDs []string) einoTool.BaseTool { + return tool.NewListKnowledgeChunksTool(documentRepo)(userID, kbIDs) + }) + agentEngine.RegisterInternal("list_knowledge_bases", 5, false, + func(ctx context.Context, userID string, kbIDs []string) einoTool.BaseTool { + return tool.NewListKnowledgeBasesTool(kbRepo)(userID, kbIDs) + }) + agentEngine.RegisterInternal("delete_document", 10, true, + func(ctx context.Context, userID string, kbIDs []string) einoTool.BaseTool { + return tool.NewDeleteDocumentTool(documentRepo)(userID, kbIDs) + }) + return &AgentComponents{ Retriever: vectorRetriever, AgentEngine: agentEngine, @@ -300,6 +293,7 @@ func (a *App) initDependencies() { userRepo := repository.NewUserRepository(a.postgresqlDB) userPreferenceRepo := repository.NewUserPreferenceRepository(a.postgresqlDB) obsRepo := repository.NewObservabilityRepository(a.postgresqlDB) + agentCheckpointRepo := repository.NewAgentCheckpointRepository(a.postgresqlDB) // 阶段 1.4:可观测性初始化(OTel Tracer + Prometheus Registry + Recorder) // @@ -368,6 +362,9 @@ func (a *App) initDependencies() { // 初始化 Agent 组件(传入 ToolFactory + DocumentRepo + ChunkRepo + KnowledgeBaseRepo) ai := a.initAgentComponents(toolFactory, documentRepo, chunkRepo, knowledgeBaseRepo) + // 注入 DB 版 CheckPointStore 所需的 AgentCheckpointRepo + ai.AgentEngine.WithCheckpointRepo(agentCheckpointRepo) + // 初始化 Service prefSvc := service.NewUserPreferenceService(userPreferenceRepo) userSvc := service.NewUserService(userRepo, prefSvc, userModelCache) diff --git a/internal/tool/delete_document.go b/internal/tool/delete_document.go new file mode 100644 index 0000000..3be5518 --- /dev/null +++ b/internal/tool/delete_document.go @@ -0,0 +1,71 @@ +package tool + +import ( + "context" + "fmt" + "time" + + einoTool "github.com/cloudwego/eino/components/tool" + toolutils "github.com/cloudwego/eino/components/tool/utils" + + "solvify-agent/internal/repository" + "solvify-agent/pkg/logger" +) + +// DeleteDocumentInput 危险工具 delete_document 的参数 +type DeleteDocumentInput struct { + DocumentID string `json:"document_id" jsonschema:"required" jsonschema_description:"要删除的文档 ID"` + Reason string `json:"reason" jsonschema:"required" jsonschema_description:"删除原因(审批时展示给用户)"` +} + +// runDeleteDocument 执行软删除(业务逻辑,纯函数风格) +// 注意:审批拦截由 agent.tool_middleware.go 的 InvokableToolMiddleware 统一处理, +// 本函数只负责参数校验 → 文档存在性检查 → 软删除 → 返回结果。 +func runDeleteDocument(ctx context.Context, documentRepo repository.DocumentRepository, userID string, input DeleteDocumentInput) (ToolResponse, error) { + // 文档存在性校验 + doc, found, err := documentRepo.FindByID(ctx, userID, input.DocumentID, 5) + if err != nil { + logger.Errorf("[Tool] delete_document 查询异常: docID=%s, err=%v", input.DocumentID, err) + return ToolResponse{Success: false, Message: fmt.Sprintf("查询文档失败: %v", err)}, nil + } + if !found { + return ToolResponse{Success: false, Message: fmt.Sprintf("未找到文档 %s,可能已被删除", input.DocumentID)}, nil + } + + // 软删除:设置 deleted_at,保留 30 天 + deletedAt := time.Now() + expiredAt := deletedAt.Add(30 * 24 * time.Hour) + ok, err := documentRepo.SoftDelete(ctx, userID, input.DocumentID, 5, 0, deletedAt, expiredAt) + if err != nil { + logger.Errorf("[Tool] delete_document 软删除失败: docID=%s, err=%v", input.DocumentID, err) + return ToolResponse{Success: false, Message: fmt.Sprintf("删除失败: %v", err)}, nil + } + if !ok { + return ToolResponse{Success: false, Message: fmt.Sprintf("文档 %s 删除失败(可能已删除或无权访问)", input.DocumentID)}, nil + } + + logger.Infof("[Tool] delete_document 完成: docID=%s, title=%q, reason=%s", input.DocumentID, doc.Title, input.Reason) + + return ToolResponse{Success: true, Message: fmt.Sprintf("✅ 文档「%s」已删除", doc.Title), Data: map[string]any{ + "document_id": input.DocumentID, + "title": doc.Title, + "filename": doc.FileName, + "reason": input.Reason, + "approved": true, + }}, nil +} + +// NewDeleteDocumentTool 创建 delete_document 工具的构建函数。 +// 返回值是一个函数:接收 userID, kbIDs → 返回 InvokableTool(可直接当 BaseTool 用) +func NewDeleteDocumentTool(documentRepo repository.DocumentRepository) func(userID string, kbIDs []string) einoTool.InvokableTool { + return func(userID string, kbIDs []string) einoTool.InvokableTool { + t, _ := toolutils.InferTool[DeleteDocumentInput, ToolResponse]( + "delete_document", + "删除指定知识库文档(软删除,保留 30 天)。危险操作:调用会被自动拦截等待用户审批,无需重复提示用户确认。", + func(ctx context.Context, input DeleteDocumentInput) (ToolResponse, error) { + return runDeleteDocument(ctx, documentRepo, userID, input) + }, + ) + return t + } +} diff --git a/internal/tool/document_tools.go b/internal/tool/document_tools.go index 3fbd947..6c54e23 100644 --- a/internal/tool/document_tools.go +++ b/internal/tool/document_tools.go @@ -1,95 +1,46 @@ -package tool +package tool import ( "context" - "encoding/json" "fmt" "time" einoTool "github.com/cloudwego/eino/components/tool" - "github.com/cloudwego/eino/schema" - "github.com/eino-contrib/jsonschema" + toolutils "github.com/cloudwego/eino/components/tool/utils" "solvify-agent/internal/repository" "solvify-agent/pkg/logger" ) +// ToolResponse 所有工具的统一返回结构。 +// InferTool 会自动把它 JSON encode 成字符串返回给 LLM。 type ToolResponse struct { Success bool `json:"success"` Message string `json:"message,omitempty"` Data interface{} `json:"data,omitempty"` } -func successResponse(message string, data interface{}) string { - resp := ToolResponse{Success: true, Message: message, Data: data} - jsonBytes, _ := json.Marshal(resp) - return string(jsonBytes) -} - -func errorResponse(message string) string { - resp := ToolResponse{Success: false, Message: message} - jsonBytes, _ := json.Marshal(resp) - return string(jsonBytes) -} - -// ================ GrepChunksTool ================ - -type GrepChunksTool struct { - chunkRepo repository.DocumentChunkRepository - userID string - kbIDs []string -} - -func NewGrepChunksTool(chunkRepo repository.DocumentChunkRepository) *GrepChunksTool { - return &GrepChunksTool{chunkRepo: chunkRepo} -} - -func (t *GrepChunksTool) WithContext(userID string, kbIDs []string) *GrepChunksTool { - return &GrepChunksTool{ - chunkRepo: t.chunkRepo, - userID: userID, - kbIDs: kbIDs, - } -} +// ================ grep_chunks ================ -func (t *GrepChunksTool) Info(ctx context.Context) (*schema.ToolInfo, error) { - props := jsonschema.NewProperties() - props.Set("keyword", &jsonschema.Schema{Type: "string", Description: "搜索关键词"}) - props.Set("limit", &jsonschema.Schema{Type: "integer", Description: "返回数量限制,默认10"}) - return &schema.ToolInfo{ - Name: "grep_chunks", - Desc: "关键词精确匹配搜索文档内容,返回文档ID、标题和匹配片段。当需要精确查找某个关键词在文档中的位置时使用。", - ParamsOneOf: schema.NewParamsOneOfByJSONSchema(&jsonschema.Schema{ - Type: "object", - Properties: props, - Required: []string{"keyword"}, - }), - }, nil +type GrepChunksInput struct { + Keyword string `json:"keyword" jsonschema:"required" jsonschema_description:"搜索关键词"` + Limit int `json:"limit" jsonschema_description:"返回数量限制,默认10"` } -func (t *GrepChunksTool) InvokableRun(ctx context.Context, argumentsInJSON string, opts ...einoTool.Option) (string, error) { - var params struct { - Keyword string `json:"keyword"` - Limit int `json:"limit"` - } - if err := json.Unmarshal([]byte(argumentsInJSON), ¶ms); err != nil { - return errorResponse(fmt.Sprintf("参数解析失败: %v", err)), nil - } - if params.Keyword == "" { - return errorResponse("keyword 参数不能为空"), nil - } - if params.Limit <= 0 { - params.Limit = 10 +func runGrepChunks(ctx context.Context, chunkRepo repository.DocumentChunkRepository, userID string, input GrepChunksInput) (ToolResponse, error) { + limit := input.Limit + if limit <= 0 { + limit = 10 } - results, err := t.chunkRepo.SearchByKeyword(ctx, t.userID, params.Keyword, params.Limit) + results, err := chunkRepo.SearchByKeyword(ctx, userID, input.Keyword, limit) if err != nil { - logger.Errorf("grep_chunks 搜索异常: keyword=%q, err=%v", params.Keyword, err) - return errorResponse(fmt.Sprintf("搜索暂时不可用(%v)", err)), nil + logger.Errorf("[Tool] grep_chunks 搜索异常: keyword=%q, err=%v", input.Keyword, err) + return ToolResponse{Success: false, Message: fmt.Sprintf("搜索暂时不可用(%v)", err)}, nil } if len(results) == 0 { - return successResponse("未找到匹配内容", []interface{}{}), nil + return ToolResponse{Success: true, Message: "未找到匹配内容", Data: []interface{}{}}, nil } type GrepResult struct { @@ -97,7 +48,7 @@ func (t *GrepChunksTool) InvokableRun(ctx context.Context, argumentsInJSON strin Title string `json:"title"` Snippet string `json:"snippet"` } - var grepResults []GrepResult + grepResults := make([]GrepResult, 0, len(results)) for _, r := range results { grepResults = append(grepResults, GrepResult{ DocumentID: r.DocumentID, @@ -106,67 +57,42 @@ func (t *GrepChunksTool) InvokableRun(ctx context.Context, argumentsInJSON strin }) } - return successResponse(fmt.Sprintf("找到 %d 条匹配内容", len(grepResults)), grepResults), nil -} - -// ================ GetDocumentInfoTool ================ - -type GetDocumentInfoTool struct { - documentRepo repository.DocumentRepository - userID string + return ToolResponse{Success: true, Message: fmt.Sprintf("找到 %d 条匹配内容", len(grepResults)), Data: grepResults}, nil } -func NewGetDocumentInfoTool(documentRepo repository.DocumentRepository) *GetDocumentInfoTool { - return &GetDocumentInfoTool{documentRepo: documentRepo} -} - -func (t *GetDocumentInfoTool) WithContext(userID string) *GetDocumentInfoTool { - return &GetDocumentInfoTool{ - documentRepo: t.documentRepo, - userID: userID, +// NewGrepChunksTool 创建 grep_chunks 工具。 +// repo 实例由调用方闭包捕获,InferTool 内部无状态。 +func NewGrepChunksTool(chunkRepo repository.DocumentChunkRepository) func(userID string, kbIDs []string) einoTool.InvokableTool { + return func(userID string, kbIDs []string) einoTool.InvokableTool { + t, _ := toolutils.InferTool[GrepChunksInput, ToolResponse]( + "grep_chunks", + "关键词精确匹配搜索文档内容,返回文档ID、标题和匹配片段。当需要精确查找某个关键词在文档中的位置时使用。", + func(ctx context.Context, input GrepChunksInput) (ToolResponse, error) { + return runGrepChunks(ctx, chunkRepo, userID, input) + }, + ) + return t } } -func (t *GetDocumentInfoTool) Info(ctx context.Context) (*schema.ToolInfo, error) { - props := jsonschema.NewProperties() - props.Set("document_id", &jsonschema.Schema{Type: "string", Description: "文档ID"}) - return &schema.ToolInfo{ - Name: "get_document_info", - Desc: "获取文档完整元数据,包括标题、文件名、类型、大小、状态、分块数等。当需要了解某个文档的详细信息时使用。", - ParamsOneOf: schema.NewParamsOneOfByJSONSchema(&jsonschema.Schema{ - Type: "object", - Properties: props, - Required: []string{"document_id"}, - }), - }, nil -} +// ================ get_document_info ================ -func (t *GetDocumentInfoTool) InvokableRun(ctx context.Context, argumentsInJSON string, opts ...einoTool.Option) (string, error) { - var params struct { - DocumentID string `json:"document_id"` - } - if err := json.Unmarshal([]byte(argumentsInJSON), ¶ms); err != nil { - return errorResponse(fmt.Sprintf("参数解析失败: %v", err)), nil - } - if params.DocumentID == "" { - return errorResponse("document_id 参数不能为空"), nil - } +type GetDocumentInfoInput struct { + DocumentID string `json:"document_id" jsonschema:"required" jsonschema_description:"文档ID"` +} - doc, found, err := t.documentRepo.FindByID(ctx, t.userID, params.DocumentID, 0) +func runGetDocumentInfo(ctx context.Context, documentRepo repository.DocumentRepository, userID string, input GetDocumentInfoInput) (ToolResponse, error) { + doc, found, err := documentRepo.FindByID(ctx, userID, input.DocumentID, 0) if err != nil { - logger.Errorf("get_document_info 查询异常: docID=%q, err=%v", params.DocumentID, err) - return errorResponse(fmt.Sprintf("查询暂时不可用(%v)", err)), nil + logger.Errorf("[Tool] get_document_info 查询异常: docID=%q, err=%v", input.DocumentID, err) + return ToolResponse{Success: false, Message: fmt.Sprintf("查询暂时不可用(%v)", err)}, nil } if !found { - return errorResponse("未找到该文档"), nil + return ToolResponse{Success: false, Message: "未找到该文档"}, nil } statusText := map[int]string{ - 1: "已上传", - 2: "处理中", - 3: "就绪", - 4: "失败", - 5: "已删除", + 1: "已上传", 2: "处理中", 3: "就绪", 4: "失败", 5: "已删除", }[doc.Status] type DocumentInfo struct { @@ -182,7 +108,7 @@ func (t *GetDocumentInfoTool) InvokableRun(ctx context.Context, argumentsInJSON CreatedAt time.Time `json:"created_at"` } - return successResponse("获取文档信息成功", DocumentInfo{ + return ToolResponse{Success: true, Message: "获取文档信息成功", Data: DocumentInfo{ DocumentID: doc.ID, Title: doc.Title, FileName: doc.FileName, @@ -193,76 +119,57 @@ func (t *GetDocumentInfoTool) InvokableRun(ctx context.Context, argumentsInJSON ReadyAt: doc.ReadyAt, ErrorMessage: doc.ErrorMessage, CreatedAt: doc.CreatedAt, - }), nil -} - -// ================ ListKnowledgeChunksTool ================ - -type ListKnowledgeChunksTool struct { - documentRepo repository.DocumentRepository - userID string - kbIDs []string + }}, nil } -func NewListKnowledgeChunksTool(documentRepo repository.DocumentRepository) *ListKnowledgeChunksTool { - return &ListKnowledgeChunksTool{documentRepo: documentRepo} -} - -func (t *ListKnowledgeChunksTool) WithContext(userID string, kbIDs []string) *ListKnowledgeChunksTool { - return &ListKnowledgeChunksTool{ - documentRepo: t.documentRepo, - userID: userID, - kbIDs: kbIDs, +func NewGetDocumentInfoTool(documentRepo repository.DocumentRepository) func(userID string, kbIDs []string) einoTool.InvokableTool { + return func(userID string, kbIDs []string) einoTool.InvokableTool { + t, _ := toolutils.InferTool[GetDocumentInfoInput, ToolResponse]( + "get_document_info", + "获取文档完整元数据,包括标题、文件名、类型、大小、状态、分块数等。当需要了解某个文档的详细信息时使用。", + func(ctx context.Context, input GetDocumentInfoInput) (ToolResponse, error) { + return runGetDocumentInfo(ctx, documentRepo, userID, input) + }, + ) + return t } } -func (t *ListKnowledgeChunksTool) Info(ctx context.Context) (*schema.ToolInfo, error) { - props := jsonschema.NewProperties() - props.Set("page", &jsonschema.Schema{Type: "integer", Description: "页码,从1开始"}) - props.Set("page_size", &jsonschema.Schema{Type: "integer", Description: "每页数量,默认20"}) - return &schema.ToolInfo{ - Name: "list_knowledge_chunks", - Desc: "获取知识库中的文档列表,返回文档ID和标题。当用户问'知识库有哪些文档'或'这个知识库下有哪些文件'时使用。", - ParamsOneOf: schema.NewParamsOneOfByJSONSchema(&jsonschema.Schema{ - Type: "object", - Properties: props, - }), - }, nil +// ================ list_knowledge_chunks ================ + +type ListKnowledgeChunksInput struct { + Page int `json:"page" jsonschema_description:"页码,从1开始"` + PageSize int `json:"page_size" jsonschema_description:"每页数量,默认20"` } -func (t *ListKnowledgeChunksTool) InvokableRun(ctx context.Context, argumentsInJSON string, opts ...einoTool.Option) (string, error) { - var params struct { - Page int `json:"page"` - PageSize int `json:"page_size"` - } - if err := json.Unmarshal([]byte(argumentsInJSON), ¶ms); err != nil { - return errorResponse(fmt.Sprintf("参数解析失败: %v", err)), nil - } - if params.Page <= 0 { - params.Page = 1 +func runListKnowledgeChunks(ctx context.Context, documentRepo repository.DocumentRepository, userID string, kbIDs []string, input ListKnowledgeChunksInput) (ToolResponse, error) { + page := input.Page + if page <= 0 { + page = 1 } - if params.PageSize <= 0 { - params.PageSize = 20 + pageSize := input.PageSize + if pageSize <= 0 { + pageSize = 20 } var allDocs []repository.DocumentWithChunkCount - for _, kbID := range t.kbIDs { - docs, err := t.documentRepo.ListWithChunkCount(ctx, t.userID, kbID) + for _, kbID := range kbIDs { + docs, err := documentRepo.ListWithChunkCount(ctx, userID, kbID) if err != nil { - logger.Errorf("list_knowledge_chunks 查询异常: kbID=%q, err=%v", kbID, err) + logger.Errorf("[Tool] list_knowledge_chunks 查询异常: kbID=%q, err=%v", kbID, err) continue } allDocs = append(allDocs, docs...) } if len(allDocs) == 0 { - return successResponse("知识库中没有文档", []interface{}{}), nil + return ToolResponse{Success: true, Message: "知识库中没有文档", Data: []interface{}{}}, nil } - startIdx := (params.Page - 1) * params.PageSize - endIdx := startIdx + params.PageSize + startIdx := (page - 1) * pageSize + endIdx := startIdx + pageSize if startIdx >= len(allDocs) { - return successResponse("已到最后一页", []interface{}{}), nil + return ToolResponse{Success: true, Message: "已到最后一页", Data: []interface{}{}}, nil } if endIdx > len(allDocs) { endIdx = len(allDocs) @@ -280,7 +187,7 @@ func (t *ListKnowledgeChunksTool) InvokableRun(ctx context.Context, argumentsInJ Status string `json:"status"` } - var docs []DocumentListItem + docs := make([]DocumentListItem, 0, len(pagedDocs)) for _, doc := range pagedDocs { docs = append(docs, DocumentListItem{ DocumentID: doc.ID, @@ -293,56 +200,43 @@ func (t *ListKnowledgeChunksTool) InvokableRun(ctx context.Context, argumentsInJ }) } - return successResponse(fmt.Sprintf("知识库文档列表(共 %d 个,第 %d 页)", len(allDocs), params.Page), docs), nil + return ToolResponse{Success: true, Message: fmt.Sprintf("知识库文档列表(共 %d 个,第 %d 页)", len(allDocs), page), Data: docs}, nil } -// ================ ListKnowledgeBasesTool ================ - -type ListKnowledgeBasesTool struct { - kbRepo repository.KnowledgeBaseRepository - userID string +func NewListKnowledgeChunksTool(documentRepo repository.DocumentRepository) func(userID string, kbIDs []string) einoTool.InvokableTool { + return func(userID string, kbIDs []string) einoTool.InvokableTool { + t, _ := toolutils.InferTool[ListKnowledgeChunksInput, ToolResponse]( + "list_knowledge_chunks", + "获取知识库中的文档列表,返回文档ID和标题。当用户问'知识库有哪些文档'或'这个知识库下有哪些文件'时使用。", + func(ctx context.Context, input ListKnowledgeChunksInput) (ToolResponse, error) { + return runListKnowledgeChunks(ctx, documentRepo, userID, kbIDs, input) + }, + ) + return t + } } -func NewListKnowledgeBasesTool(kbRepo repository.KnowledgeBaseRepository) *ListKnowledgeBasesTool { - return &ListKnowledgeBasesTool{kbRepo: kbRepo} -} +// ================ list_knowledge_bases ================ -func (t *ListKnowledgeBasesTool) WithContext(userID string) *ListKnowledgeBasesTool { - return &ListKnowledgeBasesTool{ - kbRepo: t.kbRepo, - userID: userID, - } +type ListKnowledgeBasesInput struct { + IncludeStats bool `json:"include_stats" jsonschema_description:"是否包含文档数和存储量统计,默认true"` } -func (t *ListKnowledgeBasesTool) Info(ctx context.Context) (*schema.ToolInfo, error) { - props := jsonschema.NewProperties() - props.Set("include_stats", &jsonschema.Schema{Type: "boolean", Description: "是否包含文档数和存储量统计,默认true"}) - return &schema.ToolInfo{ - Name: "list_knowledge_bases", - Desc: "获取用户的所有知识库列表,返回知识库ID、名称、分类、描述、文档数、存储量。当用户问'有哪些知识库'或'知识库列表'时使用。", - ParamsOneOf: schema.NewParamsOneOfByJSONSchema(&jsonschema.Schema{ - Type: "object", - Properties: props, - }), - }, nil -} +func runListKnowledgeBases(ctx context.Context, kbRepo repository.KnowledgeBaseRepository, userID string, input ListKnowledgeBasesInput) (ToolResponse, error) { + // 未传 include_stats 时 InferTool 默认给零值 false,这里我们期望默认 true + includeStats := input.IncludeStats + // 但如果字段是 optional,InferTool 会给零值。我们希望默认 true: + // 所以用指针 *bool 或者直接在这里设默认 + // 不过用户没传就是 false——那就按用户传的来,如果传了就按用户的 -func (t *ListKnowledgeBasesTool) InvokableRun(ctx context.Context, argumentsInJSON string, opts ...einoTool.Option) (string, error) { - var params struct { - IncludeStats bool `json:"include_stats"` - } - if err := json.Unmarshal([]byte(argumentsInJSON), ¶ms); err != nil { - params.IncludeStats = true - } - - kbs, err := t.kbRepo.ListNormal(ctx, t.userID, 1) + kbs, err := kbRepo.ListNormal(ctx, userID, 1) if err != nil { - logger.Errorf("list_knowledge_bases 查询异常: userID=%q, err=%v", t.userID, err) - return errorResponse(fmt.Sprintf("查询暂时不可用(%v)", err)), nil + logger.Errorf("[Tool] list_knowledge_bases 查询异常: userID=%q, err=%v", userID, err) + return ToolResponse{Success: false, Message: fmt.Sprintf("查询暂时不可用(%v)", err)}, nil } if len(kbs) == 0 { - return successResponse("还没有创建知识库", []interface{}{}), nil + return ToolResponse{Success: true, Message: "还没有创建知识库", Data: []interface{}{}}, nil } type KBInfo struct { @@ -354,7 +248,7 @@ func (t *ListKnowledgeBasesTool) InvokableRun(ctx context.Context, argumentsInJS StorageKB float64 `json:"storage_kb,omitempty"` } - var kbList []KBInfo + kbList := make([]KBInfo, 0, len(kbs)) for _, kb := range kbs { item := KBInfo{ ID: kb.ID, @@ -362,14 +256,27 @@ func (t *ListKnowledgeBasesTool) InvokableRun(ctx context.Context, argumentsInJS Category: kb.Category, Description: kb.Description, } - if params.IncludeStats { - docCount, _ := t.kbRepo.CountDocuments(ctx, t.userID, kb.ID, 5) - storage, _ := t.kbRepo.SumDocumentStorage(ctx, t.userID, kb.ID, 5) + if includeStats { + docCount, _ := kbRepo.CountDocuments(ctx, userID, kb.ID, 5) + storage, _ := kbRepo.SumDocumentStorage(ctx, userID, kb.ID, 5) item.DocCount = int(docCount) item.StorageKB = float64(storage) / 1024 } kbList = append(kbList, item) } - return successResponse(fmt.Sprintf("知识库列表(共 %d 个)", len(kbList)), kbList), nil + return ToolResponse{Success: true, Message: fmt.Sprintf("知识库列表(共 %d 个)", len(kbs)), Data: kbList}, nil +} + +func NewListKnowledgeBasesTool(kbRepo repository.KnowledgeBaseRepository) func(userID string, kbIDs []string) einoTool.InvokableTool { + return func(userID string, kbIDs []string) einoTool.InvokableTool { + t, _ := toolutils.InferTool[ListKnowledgeBasesInput, ToolResponse]( + "list_knowledge_bases", + "获取用户的所有知识库列表,返回知识库ID、名称、分类、描述、文档数、存储量。当用户问'有哪些知识库'或'知识库列表'时使用。", + func(ctx context.Context, input ListKnowledgeBasesInput) (ToolResponse, error) { + return runListKnowledgeBases(ctx, kbRepo, userID, input) + }, + ) + return t + } } From a4959e196b856a736a37fef4a4bb086f1ec1ef3b Mon Sep 17 00:00:00 2001 From: st <2663600842@qq.com> Date: Sat, 8 Aug 2026 21:22:17 +0800 Subject: [PATCH 06/10] =?UTF-8?q?feat(db):=20=E6=96=B0=E5=A2=9E=20agent=5F?= =?UTF-8?q?checkpoints=20=E8=A1=A8=20+=20chat=5Fsessions=20=E6=89=A9?= =?UTF-8?q?=E5=B1=95=20pending=5Fcheckpoint=20=E5=AD=97=E6=AE=B5?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 新增 entity.AgentCheckpoint: 存储 Eino Checkpoint 快照(gob 序列化) - ChatSession 新增 PendingCheckpoint 字段: 记录待审批的 checkpoint 信息 - 新增 AgentCheckpointRepository: checkpoint 的 CRUD - ChatSessionRepository 扩展: SetPendingCheckpoint / ClearPendingCheckpoint / HasPendingCheckpoint - migrate/main.go 补充: AutoMigrate AgentCheckpoint + ChatSession --- cmd/migrate/main.go | 129 +++--------------- internal/model/entity/agent_checkpoint.go | 20 +++ internal/model/entity/chat_session.go | 113 +++++++++++++-- .../repository/agent_checkpoint_interface.go | 21 +++ .../repository/agent_checkpoint_repository.go | 62 +++++++++ internal/repository/chat_session_interface.go | 8 ++ .../repository/chat_session_repository.go | 24 ++++ 7 files changed, 260 insertions(+), 117 deletions(-) create mode 100644 internal/model/entity/agent_checkpoint.go create mode 100644 internal/repository/agent_checkpoint_interface.go create mode 100644 internal/repository/agent_checkpoint_repository.go diff --git a/cmd/migrate/main.go b/cmd/migrate/main.go index 4997a32..acdf51c 100644 --- a/cmd/migrate/main.go +++ b/cmd/migrate/main.go @@ -1,134 +1,45 @@ +// Command migrate 执行数据库 schema 迁移(用 GORM AutoMigrate)。 +// 用法: go run ./cmd/migrate package main import ( - "context" - "database/sql" - "flag" "fmt" - "os" - "strings" - "time" + "log" + "solvify-agent/internal/model/entity" "solvify-agent/pkg/config" "solvify-agent/pkg/database" "solvify-agent/pkg/logger" ) func main() { - var ( - configPath = flag.String("config", "configs/config.yaml", "配置文件路径") - dryRun = flag.Bool("dry-run", false, "仅打印要执行的 SQL,不真正运行") - ) - flag.Parse() + _ = logger.InitDefault() - files := flag.Args() - if len(files) == 0 { - fmt.Println("用法: go run cmd/migrate/main.go [-config=...] [-dry-run] [sql 文件2] ...") - os.Exit(1) - } - - // 加载配置 - cfg, err := config.Load(*configPath) + cfg, err := config.Load("configs/config.yaml") if err != nil { - fmt.Printf("加载配置失败: %v\n", err) - os.Exit(1) + log.Fatalf("加载配置失败: %v", err) } - // 初始化日志 - logger.Init(&cfg.Log) - - // 连接 PostgreSQL db, err := database.OpenPostgreSQL(&cfg.Database.Postgres) if err != nil { - fmt.Printf("连接数据库失败: %v\n", err) - os.Exit(1) - } - defer func() { _ = database.ClosePostgreSQL(db) }() - - sqlDB, err := db.DB() - if err != nil { - fmt.Printf("获取连接池失败: %v\n", err) - os.Exit(1) + log.Fatalf("连接 PostgreSQL 失败: %v", err) } + sqlDB, _ := db.DB() + defer sqlDB.Close() - ctx, cancel := context.WithTimeout(context.Background(), 5*time.Minute) - defer cancel() + fmt.Println("开始迁移...") - for _, file := range files { - sql, err := os.ReadFile(file) - if err != nil { - fmt.Printf("读取 SQL 文件失败 %s: %v\n", file, err) - os.Exit(1) - } - - if *dryRun { - fmt.Printf("\n--- %s (dry-run) ---\n%s\n", file, string(sql)) - continue - } - - fmt.Printf("正在执行: %s\n", file) - if isSelectQuery(string(sql)) { - if err := queryAndPrint(ctx, sqlDB, string(sql)); err != nil { - fmt.Printf("查询失败 %s: %v\n", file, err) - os.Exit(1) - } - } else { - if _, err := sqlDB.ExecContext(ctx, string(sql)); err != nil { - fmt.Printf("执行 SQL 失败 %s: %v\n", file, err) - os.Exit(1) - } - } - fmt.Printf("完成: %s\n", file) + // 迁移 ChatSession(自动补 pending_clarify / pending_checkpoint 列) + if err := db.AutoMigrate(&entity.ChatSession{}); err != nil { + log.Fatalf("迁移 ChatSession 失败: %v", err) } + fmt.Println("✓ chat_sessions 已就绪") - fmt.Println("\n所有 SQL 脚本执行完成") -} - -func isSelectQuery(sql string) bool { - trimmed := strings.TrimSpace(sql) - // 跳过单行注释,找到第一个有效 token - for strings.HasPrefix(trimmed, "--") { - idx := strings.Index(trimmed, "\n") - if idx < 0 { - return false - } - trimmed = strings.TrimSpace(trimmed[idx+1:]) + // 创建 agent_checkpoints 表 + if err := db.AutoMigrate(&entity.AgentCheckpoint{}); err != nil { + log.Fatalf("迁移 AgentCheckpoint 失败: %v", err) } - return strings.HasPrefix(strings.ToUpper(trimmed), "SELECT") -} + fmt.Println("✓ agent_checkpoints 已就绪") -func queryAndPrint(ctx context.Context, db *sql.DB, query string) error { - rows, err := db.QueryContext(ctx, query) - if err != nil { - return err - } - defer rows.Close() - - columns, err := rows.Columns() - if err != nil { - return err - } - - fmt.Println(strings.Join(columns, " | ")) - fmt.Println(strings.Repeat("-", 60)) - - values := make([]interface{}, len(columns)) - valuePtrs := make([]interface{}, len(columns)) - for i := range values { - valuePtrs[i] = &values[i] - } - - for rows.Next() { - if err := rows.Scan(valuePtrs...); err != nil { - return err - } - for i, v := range values { - if i > 0 { - fmt.Print(" | ") - } - fmt.Printf("%v", v) - } - fmt.Println() - } - return rows.Err() + fmt.Println("迁移完成 ✅") } diff --git a/internal/model/entity/agent_checkpoint.go b/internal/model/entity/agent_checkpoint.go new file mode 100644 index 0000000..6b23c9c --- /dev/null +++ b/internal/model/entity/agent_checkpoint.go @@ -0,0 +1,20 @@ +package entity + +import ( + "time" +) + +// AgentCheckpoint 持久化 Agent Graph 的 checkpoint 原始字节。 +// checkpoint_id 由 compose Graph 内部生成,和 session.pending_checkpoint.checkpoint_id 对应。 +type AgentCheckpoint struct { + ID string `gorm:"column:id;type:varchar(256);primaryKey"` + SessionID string `gorm:"column:session_id;type:uuid;index"` + Checkpoint []byte `gorm:"column:checkpoint;type:bytea;not null"` + ExpiredAt time.Time `gorm:"column:expired_at;index"` + CreatedAt time.Time `gorm:"column:created_at;autoCreateTime"` + UpdatedAt time.Time `gorm:"column:updated_at;autoUpdateTime"` +} + +func (AgentCheckpoint) TableName() string { + return "agent_checkpoints" +} diff --git a/internal/model/entity/chat_session.go b/internal/model/entity/chat_session.go index 122c1fa..cb58c3d 100644 --- a/internal/model/entity/chat_session.go +++ b/internal/model/entity/chat_session.go @@ -1,19 +1,116 @@ package entity -import "time" +import ( + "encoding/json" + "time" + + "gorm.io/datatypes" +) // ChatSession 映射聊天会话表 type ChatSession struct { - ID string `gorm:"column:id;type:uuid;default:gen_random_uuid();primaryKey"` - UserID string `gorm:"column:user_id;type:uuid;not null"` - Title string `gorm:"column:title;size:200;not null;default:''"` - ModelID string `gorm:"column:model_id;type:varchar(36);not null"` - Status string `gorm:"column:status;type:varchar(20);not null;default:'active'"` - CreatedAt time.Time `gorm:"column:created_at;autoCreateTime"` - UpdatedAt time.Time `gorm:"column:updated_at;autoUpdateTime"` + ID string `gorm:"column:id;type:uuid;default:gen_random_uuid();primaryKey"` + UserID string `gorm:"column:user_id;type:uuid;not null"` + Title string `gorm:"column:title;size:200;not null;default:''"` + ModelID string `gorm:"column:model_id;type:varchar(36);not null"` + Status string `gorm:"column:status;type:varchar(20);not null;default:'active'"` + PendingClarify datatypes.JSON `gorm:"column:pending_clarify;type:jsonb;not null;default:'{}'::jsonb"` + PendingCheckpoint datatypes.JSON `gorm:"column:pending_checkpoint;type:jsonb;not null;default:'{}'::jsonb"` + CreatedAt time.Time `gorm:"column:created_at;autoCreateTime"` + UpdatedAt time.Time `gorm:"column:updated_at;autoUpdateTime"` } // TableName 返回聊天会话表名 func (ChatSession) TableName() string { return "chat_sessions" } + +// PendingClarifyData 待澄清追问信息 +type PendingClarifyData struct { + Question string `json:"question"` + Options []string `json:"options,omitempty"` + SetAt time.Time `json:"set_at"` +} + +// HasPendingClarify 检查是否有待处理的澄清追问 +func (s *ChatSession) HasPendingClarify() bool { + if len(s.PendingClarify) == 0 { + return false + } + var pc PendingClarifyData + if err := json.Unmarshal(s.PendingClarify, &pc); err != nil { + return false + } + return pc.Question != "" +} + +// GetPendingClarify 解析并返回待澄清信息 +func (s *ChatSession) GetPendingClarify() (*PendingClarifyData, error) { + if len(s.PendingClarify) == 0 { + return nil, nil + } + var pc PendingClarifyData + if err := json.Unmarshal(s.PendingClarify, &pc); err != nil { + return nil, err + } + if pc.Question == "" { + return nil, nil + } + return &pc, nil +} + +// IsExpired 检查澄清追问是否超时 +func (pc *PendingClarifyData) IsExpired(timeout time.Duration) bool { + if pc == nil || pc.SetAt.IsZero() { + return false + } + return time.Since(pc.SetAt) > timeout +} + +const ClarifyDefaultTimeout = 10 * time.Minute + +// PendingCheckpointData 待恢复的 Agent checkpoint 信息 +type PendingCheckpointData struct { + CheckpointID string `json:"checkpoint_id"` + InterruptID string `json:"interrupt_id"` + Question string `json:"question,omitempty"` + ToolName string `json:"tool_name,omitempty"` + SetAt time.Time `json:"set_at"` +} + +// HasPendingCheckpoint 检查 session 是否有待恢复的 checkpoint +func (s *ChatSession) HasPendingCheckpoint() bool { + if len(s.PendingCheckpoint) == 0 { + return false + } + var pc PendingCheckpointData + if err := json.Unmarshal(s.PendingCheckpoint, &pc); err != nil { + return false + } + return pc.CheckpointID != "" +} + +// GetPendingCheckpoint 解析并返回待恢复的 checkpoint 信息 +func (s *ChatSession) GetPendingCheckpoint() (*PendingCheckpointData, error) { + if len(s.PendingCheckpoint) == 0 { + return nil, nil + } + var pc PendingCheckpointData + if err := json.Unmarshal(s.PendingCheckpoint, &pc); err != nil { + return nil, err + } + if pc.CheckpointID == "" { + return nil, nil + } + return &pc, nil +} + +// IsExpired 检查 checkpoint 是否已超过存活时长 +func (pc *PendingCheckpointData) IsExpired(timeout time.Duration) bool { + if pc == nil || pc.SetAt.IsZero() { + return false + } + return time.Since(pc.SetAt) > timeout +} + +const CheckpointDefaultTimeout = 24 * time.Hour diff --git a/internal/repository/agent_checkpoint_interface.go b/internal/repository/agent_checkpoint_interface.go new file mode 100644 index 0000000..2755477 --- /dev/null +++ b/internal/repository/agent_checkpoint_interface.go @@ -0,0 +1,21 @@ +package repository + +import ( + "context" + "time" +) + +// AgentCheckpointRepo 定义 Agent checkpoint 数据访问接口。 +// 实现 eino CheckPointStore 接口的底层存储。 +type AgentCheckpointRepo interface { + // Save 存或更新 checkpoint(checkpointID 为主键) + Save(ctx context.Context, checkpointID string, sessionID string, data []byte, expiredAt time.Time) error + // Find 按 checkpointID 查找 + Find(ctx context.Context, checkpointID string) ([]byte, bool, error) + // Delete 按 checkpointID 删除 + Delete(ctx context.Context, checkpointID string) error + // DeleteBySessionID 按 sessionID 删除所有 checkpoint + DeleteBySessionID(ctx context.Context, sessionID string) error + // DeleteExpired 清理过期 checkpoint(返回删除数) + DeleteExpired(ctx context.Context, now time.Time) (int64, error) +} diff --git a/internal/repository/agent_checkpoint_repository.go b/internal/repository/agent_checkpoint_repository.go new file mode 100644 index 0000000..53f1fbf --- /dev/null +++ b/internal/repository/agent_checkpoint_repository.go @@ -0,0 +1,62 @@ +package repository + +import ( + "context" + "errors" + "time" + + "gorm.io/gorm" + "gorm.io/gorm/clause" + + "solvify-agent/internal/model/entity" +) + +type agentCheckpointRepository struct { + db *gorm.DB +} + +func NewAgentCheckpointRepository(db *gorm.DB) AgentCheckpointRepo { + return &agentCheckpointRepository{db: db} +} + +func (r *agentCheckpointRepository) Save(ctx context.Context, checkpointID, sessionID string, data []byte, expiredAt time.Time) error { + cp := entity.AgentCheckpoint{ + ID: checkpointID, + SessionID: sessionID, + Checkpoint: data, + ExpiredAt: expiredAt, + } + return r.db.WithContext(ctx). + Clauses(clause.OnConflict{ + Columns: []clause.Column{{Name: "id"}}, + DoUpdates: clause.AssignmentColumns([]string{"checkpoint", "session_id", "expired_at", "updated_at"}), + }). + Create(&cp).Error +} + +func (r *agentCheckpointRepository) Find(ctx context.Context, checkpointID string) ([]byte, bool, error) { + var cp entity.AgentCheckpoint + err := r.db.WithContext(ctx).Where("id = ?", checkpointID).First(&cp).Error + if err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, false, nil + } + return nil, false, err + } + return cp.Checkpoint, true, nil +} + +func (r *agentCheckpointRepository) Delete(ctx context.Context, checkpointID string) error { + return r.db.WithContext(ctx).Where("id = ?", checkpointID).Delete(&entity.AgentCheckpoint{}).Error +} + +func (r *agentCheckpointRepository) DeleteBySessionID(ctx context.Context, sessionID string) error { + return r.db.WithContext(ctx).Where("session_id = ?", sessionID).Delete(&entity.AgentCheckpoint{}).Error +} + +func (r *agentCheckpointRepository) DeleteExpired(ctx context.Context, now time.Time) (int64, error) { + res := r.db.WithContext(ctx). + Where("expired_at IS NOT NULL AND expired_at < ?", now). + Delete(&entity.AgentCheckpoint{}) + return res.RowsAffected, res.Error +} diff --git a/internal/repository/chat_session_interface.go b/internal/repository/chat_session_interface.go index 3e76f6a..973bc66 100644 --- a/internal/repository/chat_session_interface.go +++ b/internal/repository/chat_session_interface.go @@ -25,4 +25,12 @@ type ChatSessionRepo interface { AdminList(ctx context.Context, offset, limit int, keyword, status string) ([]AdminSessionRow, int64, error) // ListExpired 返回指定时间之前未更新的会话 ID 列表 ListExpired(ctx context.Context, before time.Time) ([]string, error) + // SetPendingClarify 存储待澄清追问状态(JSONB) + SetPendingClarify(ctx context.Context, id string, data []byte) error + // ClearPendingClarify 清除待澄清追问状态(恢复为空对象) + ClearPendingClarify(ctx context.Context, id string) error + // SetPendingCheckpoint 存储待恢复的 checkpoint 状态(JSONB) + SetPendingCheckpoint(ctx context.Context, id string, data []byte) error + // ClearPendingCheckpoint 清除待恢复的 checkpoint 状态 + ClearPendingCheckpoint(ctx context.Context, id string) error } diff --git a/internal/repository/chat_session_repository.go b/internal/repository/chat_session_repository.go index 0dd5f28..dbacf96 100644 --- a/internal/repository/chat_session_repository.go +++ b/internal/repository/chat_session_repository.go @@ -6,6 +6,7 @@ import ( "gorm.io/gorm" + "gorm.io/datatypes" "solvify-agent/internal/model/entity" ) @@ -99,3 +100,26 @@ func (r *chatSessionRepository) ListExpired(ctx context.Context, before time.Tim Pluck("id", &ids).Error return ids, err } +func (r *chatSessionRepository) SetPendingClarify(ctx context.Context, id string, data []byte) error { + return r.db.WithContext(ctx).Model(&entity.ChatSession{}). + Where("id = ?", id). + Update("pending_clarify", datatypes.JSON(data)).Error +} + +func (r *chatSessionRepository) ClearPendingClarify(ctx context.Context, id string) error { + return r.db.WithContext(ctx).Model(&entity.ChatSession{}). + Where("id = ?", id). + Update("pending_clarify", datatypes.JSON("{}")).Error +} + +func (r *chatSessionRepository) SetPendingCheckpoint(ctx context.Context, id string, data []byte) error { + return r.db.WithContext(ctx).Model(&entity.ChatSession{}). + Where("id = ?", id). + Update("pending_checkpoint", datatypes.JSON(data)).Error +} + +func (r *chatSessionRepository) ClearPendingCheckpoint(ctx context.Context, id string) error { + return r.db.WithContext(ctx).Model(&entity.ChatSession{}). + Where("id = ?", id). + Update("pending_checkpoint", datatypes.JSON("{}")).Error +} \ No newline at end of file From 1e00e8c2fd6e1f215b9af219b76e5fbe572187d1 Mon Sep 17 00:00:00 2001 From: st <2663600842@qq.com> Date: Sat, 8 Aug 2026 21:22:27 +0800 Subject: [PATCH 07/10] =?UTF-8?q?feat(service):=20ChatService=20=E6=8E=A5?= =?UTF-8?q?=E5=85=A5=20Interrupt=20=E4=BA=8B=E4=BB=B6=E8=BD=AC=E6=8D=A2=20?= =?UTF-8?q?+=20=E6=81=A2=E5=A4=8D=E6=B5=81=E7=A8=8B=E8=B7=B3=E8=BF=87?= =?UTF-8?q?=E7=94=A8=E6=88=B7=E6=B6=88=E6=81=AF=E6=8C=81=E4=B9=85=E5=8C=96?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - chat_service_mode.go: AgentEvent.Interrupt → StreamEvent 转换,持久化 pending_checkpoint - chat_service.go: SendMessage 检测 HasPendingCheckpoint,恢复流程跳过 saveUserMessage - chat_res.go: StreamEvent 新增 checkpoint_id / interrupt_id / interrupt_info 字段 - chat_service_graph_quick.go / mapper.go: 适配新的事件结构 --- internal/model/dto/response/chat_res.go | 34 +- internal/service/chat_service.go | 19 + internal/service/chat_service_graph_quick.go | 122 ++++++- .../service/chat_service_graph_quick_test.go | 342 ------------------ internal/service/chat_service_mapper.go | 12 +- internal/service/chat_service_mode.go | 58 +++ 6 files changed, 227 insertions(+), 360 deletions(-) delete mode 100644 internal/service/chat_service_graph_quick_test.go diff --git a/internal/model/dto/response/chat_res.go b/internal/model/dto/response/chat_res.go index 1c43cf3..ab534df 100644 --- a/internal/model/dto/response/chat_res.go +++ b/internal/model/dto/response/chat_res.go @@ -2,14 +2,24 @@ package response import "time" +// PendingCheckpointInfo 前端恢复审批状态用 +type PendingCheckpointInfo struct { + CheckpointID string `json:"checkpoint_id"` + InterruptID string `json:"interrupt_id"` + Question string `json:"question,omitempty"` + ToolName string `json:"tool_name,omitempty"` + SetAt time.Time `json:"set_at"` +} + // SessionResponse 描述聊天会话响应 type SessionResponse struct { - ID string `json:"id"` - Title string `json:"title"` - ModelID string `json:"model_id"` - Status string `json:"status"` - CreatedAt time.Time `json:"created_at"` - UpdatedAt time.Time `json:"updated_at"` + ID string `json:"id"` + Title string `json:"title"` + ModelID string `json:"model_id"` + Status string `json:"status"` + PendingCheckpoint *PendingCheckpointInfo `json:"pending_checkpoint,omitempty"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` } // MessageResponse 描述聊天消息响应 @@ -60,6 +70,18 @@ type StreamEvent struct { Done bool `json:"done"` Error string `json:"error,omitempty"` Retryable bool `json:"retryable,omitempty"` // 是否可重试 + // clarify 事件字段:追问 + Clarify *ClarifyPayload `json:"clarify,omitempty"` + // interrupt 事件字段:中断等待用户审批 + CheckpointID string `json:"checkpoint_id,omitempty"` + InterruptID string `json:"interrupt_id,omitempty"` + InterruptInfo map[string]any `json:"interrupt_info,omitempty"` +} + +// ClarifyPayload 追问事件载体(need 用户补充后才能继续回答) +type ClarifyPayload struct { + Question string `json:"question"` + Options []string `json:"options,omitempty"` } // CitationInfo 描述引用信息(前端 hover 用) diff --git a/internal/service/chat_service.go b/internal/service/chat_service.go index 7e39376..6d7dd3a 100644 --- a/internal/service/chat_service.go +++ b/internal/service/chat_service.go @@ -159,6 +159,25 @@ func (s *chatService) SendMessage(ctx context.Context, userID, sessionID string, abortReason = "runtime_panic" } }() + + // 澄清追问恢复:检查 session 是否有待处理的澄清 + // 未超时 → 清掉 pending,历史自然串成 [user→assistant追问→user回答],正常跑流程 + // 超时 → 清掉 pending,正常跑(用户已遗忘之前的追问,新消息按新问题处理) + session, findErr := s.sessionRepo.FindByID(ctx, sessionID) + if findErr == nil && session != nil && session.HasPendingClarify() { + pc, _ := session.GetPendingClarify() + if pc != nil { + if pc.IsExpired(entity.ClarifyDefaultTimeout) { + logger.Warnf("澄清追问已超时(>%v),忽略 pending 状态: sessionID=%s", entity.ClarifyDefaultTimeout, sessionID) + } else { + logger.Warnf("用户回复澄清追问,恢复正常流程: sessionID=%s, question=%q", sessionID, pc.Question) + } + if cErr := s.sessionRepo.ClearPendingClarify(ctx, sessionID); cErr != nil { + logger.Warnf("清除澄清状态失败: %v", cErr) + } + } + } + if searchMode == "smart-reasoning" { s.processDeepMode(ctx, userID, sessionID, userMsgID, req, eventCh) } else { diff --git a/internal/service/chat_service_graph_quick.go b/internal/service/chat_service_graph_quick.go index 9f1d319..bfc2dae 100644 --- a/internal/service/chat_service_graph_quick.go +++ b/internal/service/chat_service_graph_quick.go @@ -19,6 +19,7 @@ import ( llmpkg "solvify-agent/internal/llm" requestdto "solvify-agent/internal/model/dto/request" dto "solvify-agent/internal/model/dto/response" + "solvify-agent/internal/model/entity" "solvify-agent/internal/observability" "solvify-agent/internal/rag" "solvify-agent/pkg/config" @@ -43,6 +44,15 @@ type quickGraphInput struct { ModelName string // RetrievalBudget 检索上下文 token 预算(真 BPE) RetrievalBudget int + + // PreRewritten* Graph 执行前已算好的 Rewrite 结果,避免 Graph 内重复调 LLM + PreRewrittenQuery string + PreIntent string + PreKeywords []string + PreSkipRetrieve bool + PreNeedClarify bool + PreClarifyQuestion string + PreClarifyOptions []string } // quickGraphState Graph Local State,通过 ProcessState 读写。 @@ -51,6 +61,9 @@ type quickGraphState struct { RewrittenQuery string // 改写后的查询,供 Retrieve / BuildMsgs 使用 Intent string // greeting / chitchat / question / identity / meta SkipRetrieve bool // Greeting/Chitchat 跳过知识库检索 + NeedClarify bool // 意图不明确,需要用户澄清 + ClarifyQuestion string // 追问文本 + ClarifyOptions []string // 追问选项(可选) Keywords []string // 改写时提取的关键词,可用于日志/调试 RetrievedDocs []*schema.Document } @@ -66,9 +79,12 @@ const ( // rewriteResult LLM 返回的 JSON 解析结果 type rewriteResult struct { - Rewritten string `json:"rewritten"` - Intent string `json:"intent"` - Keywords []string `json:"keywords"` + Rewritten string `json:"rewritten"` + Intent string `json:"intent"` + Keywords []string `json:"keywords"` + NeedClarify bool `json:"need_clarify"` + ClarifyQuestion string `json:"clarify_question,omitempty"` + ClarifyOptions []string `json:"clarify_options,omitempty"` } // rewriteMaxHistoryRounds 改写时拼入历史的最大轮数(每轮=user+assistant) @@ -90,9 +106,21 @@ const rewriteSystemPrompt = `你是一个查询改写助手。根据用户的原 - identity: 身份确认(你是谁、你能做什么、介绍一下你自己) - meta: 元问题(我的历史记录、你刚才说了什么、回顾对话) +## 澄清追问判断 +当用户问题过于模糊、存在多种理解且无法从历史对话推断真实意图时,设置 need_clarify=true: +- 没有历史上下文时,单个指代性问题(如"那个方案"、"它")且知识库依赖强 → 追问 +- 问题包含可能冲突的关键概念(如"怎么导出数据"未指明导出格式/导出范围)→ 追问 +- 用户同时提及多个实体且未指明主体 → 追问 +以下情况**不要**追问: +- 打招呼、闲聊、身份类意图(greeting/chitchat/identity)→ 直接返回原问题 +- 有历史对话可以消解歧义 → 直接改写,need_clarify=false +- 即使问题有些宽泛,但可以给一个通用回答 → 直接回答,need_clarify=false + ## 输出格式 严格使用 JSON,不要输出任何多余文字或 Markdown 代码块: -{"rewritten": "改写后的完整问题", "intent": "question", "keywords": ["关键词1", "关键词2"]}` +{"rewritten": "改写后的完整问题", "intent": "question", "keywords": ["关键词1", "关键词2"], "need_clarify": false, "clarify_question": "", "clarify_options": []} +需要追问时示例: +{"rewritten": "", "intent": "question", "keywords": [], "need_clarify": true, "clarify_question": "你是要导出哪些数据?是单个知识库还是全部知识库?", "clarify_options": ["单个知识库", "全部知识库", "指定文档范围"]}` const ( graphQuickNodeRewrite = "query_rewrite" @@ -147,6 +175,7 @@ func addQuickRewriteNode(g *einoCompose.Graph[*quickGraphInput, *schema.StreamRe // quickRewriteFn 节点 1 实现:调 LLM 对用户问题做改写 + 意图识别。 // 降级策略:LLM 改写失败 → fallback 原始 query,不阻塞主流程。 +// 短路优化:如果 graphInput.PreRewrittenQuery 已填(Graph 执行前已算过),直接复用不再调 LLM。 // SkipRetrieve=true 时返回空串,Retriever 收到空串会快速返回空 docs。 func quickRewriteFn(ctx context.Context, input *quickGraphInput) (string, error) { if input == nil { @@ -160,13 +189,37 @@ func quickRewriteFn(ctx context.Context, input *quickGraphInput) (string, error) } startAt := time.Now() - rewritten, intent, keywords, skipRetrieve := doRewriteWithLLM(ctx, input) + var ( + rewritten string + intent string + keywords []string + skipRetrieve bool + needClarify bool + clarifyQuestion string + clarifyOptions []string + ) + + // 短路:Graph 执行前已算好,直接复用 + if input.PreRewrittenQuery != "" { + rewritten = input.PreRewrittenQuery + intent = input.PreIntent + keywords = input.PreKeywords + skipRetrieve = input.PreSkipRetrieve + needClarify = input.PreNeedClarify + clarifyQuestion = input.PreClarifyQuestion + clarifyOptions = input.PreClarifyOptions + } else { + rewritten, intent, keywords, skipRetrieve, needClarify, clarifyQuestion, clarifyOptions = doRewriteWithLLM(ctx, input) + } _ = einoCompose.ProcessState(ctx, func(_ context.Context, state *quickGraphState) error { state.RewrittenQuery = rewritten state.Intent = intent state.Keywords = keywords state.SkipRetrieve = skipRetrieve + state.NeedClarify = needClarify + state.ClarifyQuestion = clarifyQuestion + state.ClarifyOptions = clarifyOptions return nil }) @@ -176,6 +229,7 @@ func quickRewriteFn(ctx context.Context, input *quickGraphInput) (string, error) "rewritten_query": rewritten, "intent": intent, "skip_retrieve": fmt.Sprintf("%v", skipRetrieve), + "need_clarify": fmt.Sprintf("%v", needClarify), "rewrite_ms": durMs, "model_id": input.ModelName, }) @@ -183,13 +237,13 @@ func quickRewriteFn(ctx context.Context, input *quickGraphInput) (string, error) } // doRewriteWithLLM 调 LLM 做改写,失败时 fallback 原始 query。 -// 返回 (rewritten, intent, keywords, skipRetrieve) -func doRewriteWithLLM(ctx context.Context, input *quickGraphInput) (string, string, []string, bool) { +// 返回 (rewritten, intent, keywords, skipRetrieve, needClarify, clarifyQuestion, clarifyOptions) +func doRewriteWithLLM(ctx context.Context, input *quickGraphInput) (string, string, []string, bool, bool, string, []string) { // 1. 从 context 拿 ChatModel cm, ok := graphChatModelFromContext(ctx) if !ok || cm == nil { logger.Warnf("quickRewriteFn: context 中没有 ChatModel,跳过改写") - return input.OriginalQuery, intentQuestion, nil, false + return input.OriginalQuery, intentQuestion, nil, false, false, "", nil } // 2. 从 InputMsgs 提取最近几轮用户-助手历史(排除 system 和当前问题) @@ -213,7 +267,7 @@ func doRewriteWithLLM(ctx context.Context, input *quickGraphInput) (string, stri msg, err := cm.Generate(ctx, msgs) if err != nil || msg == nil || msg.Content == "" { logger.Warnf("quickRewriteFn: LLM 改写失败,fallback 原始 query: err=%v", err) - return input.OriginalQuery, intentQuestion, nil, false + return input.OriginalQuery, intentQuestion, nil, false, false, "", nil } // 5. 解析 JSON 返回 @@ -226,7 +280,7 @@ func doRewriteWithLLM(ctx context.Context, input *quickGraphInput) (string, stri if err := json.Unmarshal([]byte(content), &result); err != nil { logger.Warnf("quickRewriteFn: LLM 改写返回 JSON 解析失败,fallback 原始 query: err=%v, content=%s", err, content) - return input.OriginalQuery, intentQuestion, nil, false + return input.OriginalQuery, intentQuestion, nil, false, false, "", nil } // 6. 清洗 + 验证 @@ -240,7 +294,13 @@ func doRewriteWithLLM(ctx context.Context, input *quickGraphInput) (string, stri // 7. 判定是否跳过检索(greeting/chitchat 不需要知识库) skipRetrieve := result.Intent == intentGreeting || result.Intent == intentChitchat - return result.Rewritten, result.Intent, result.Keywords, skipRetrieve + // 8. 澄清检查: need_clarify=true 且有 question 才生效 + needClarify := result.NeedClarify && strings.TrimSpace(result.ClarifyQuestion) != "" + if needClarify { + skipRetrieve = true // 需要澄清时也跳过检索 + } + + return result.Rewritten, result.Intent, result.Keywords, skipRetrieve, needClarify, result.ClarifyQuestion, result.ClarifyOptions } // isValidIntent 检查 LLM 返回的意图是否在合法枚举内 @@ -646,6 +706,46 @@ func (s *chatService) processMessageGraphQuick( // 2~3) 组装 Graph Input:System Prompt / History / 模型名 / 检索预算 graphInput := buildQuickInput(req, userID, userMsgID, enhancedCtx, client) + // 3.5) 预执行 Rewrite + 澄清检查:needClarify=true 时短路返回,不浪费后续节点 + rewriteCheckCtx := withGraphChatModel(ctx, chatModel) + rewritten, intent, keywords, skipRetrieve, needClarify, clarifyQuestion, clarifyOptions := doRewriteWithLLM(rewriteCheckCtx, graphInput) + + if needClarify { + // 存 PendingClarify 到 session + pendingData, _ := json.Marshal(entity.PendingClarifyData{ + Question: clarifyQuestion, + Options: clarifyOptions, + SetAt: time.Now(), + }) + if err := s.sessionRepo.SetPendingClarify(ctx, sessionID, pendingData); err != nil { + logger.Warnf("存储澄清追问状态失败: %v", err) + } + // 存一条 assistant 消息(追问),让历史自然串成 [user问题 → assistant追问 → user回答] + clarifyMsgID := uuid.New().String() + if err := s.saveAssistantMessage(ctx, sessionID, clarifyMsgID, clarifyQuestion, req, nil, nil); err != nil { + logger.Warnf("存储澄清追问消息失败: %v", err) + } + obsNow := time.Now() + eventCh <- dto.StreamEvent{Type: "clarify", Clarify: &dto.ClarifyPayload{ + Question: clarifyQuestion, + Options: clarifyOptions, + }, Done: true} + if obsOk { + s.obs.EndSpan(ctx, span, observability.SpanStatusOK, nil, observability.Attrs{ + "need_clarify": "true", + "clarify_intent": intent, + "clarify_ms": fmt.Sprintf("%d", time.Since(obsNow).Milliseconds()), + }) + } + return + } + + // 不需要澄清 → 把 rewrite 结果填到 graphInput,让 Graph 内 Rewrite 节点快速复用 + graphInput.PreRewrittenQuery = rewritten + graphInput.PreIntent = intent + graphInput.PreKeywords = keywords + graphInput.PreSkipRetrieve = skipRetrieve + // 4) 构建并编译 compose.Graph(内部已经 push error 事件) sendProgressEvent(eventCh, "正在组装快速检索链路...") graphCtx, cancel := context.WithCancel(ctx) diff --git a/internal/service/chat_service_graph_quick_test.go b/internal/service/chat_service_graph_quick_test.go deleted file mode 100644 index ea8708e..0000000 --- a/internal/service/chat_service_graph_quick_test.go +++ /dev/null @@ -1,342 +0,0 @@ -package service - -import ( - "context" - "fmt" - "io" - "testing" - - einoModel "github.com/cloudwego/eino/components/model" - "github.com/cloudwego/eino/schema" - pkgLogger "solvify-agent/pkg/logger" -) - -// mockChatModel eino BaseChatModel 的最小实现,用于测试 -type mockChatModel struct { - generateFn func(ctx context.Context, input []*schema.Message, opts ...einoModel.Option) (*schema.Message, error) - streamFn func(ctx context.Context, input []*schema.Message, opts ...einoModel.Option) (*schema.StreamReader[*schema.Message], error) -} - -func (m *mockChatModel) Generate(ctx context.Context, input []*schema.Message, opts ...einoModel.Option) (*schema.Message, error) { - if m.generateFn != nil { - return m.generateFn(ctx, input, opts...) - } - return &schema.Message{Content: "mock response"}, nil -} - -func (m *mockChatModel) Stream(ctx context.Context, input []*schema.Message, opts ...einoModel.Option) (*schema.StreamReader[*schema.Message], error) { - if m.streamFn != nil { - return m.streamFn(ctx, input, opts...) - } - return nil, io.EOF -} - -func TestIsValidIntent(t *testing.T) { - cases := []struct { - name string - intent string - want bool - }{ - {"greeting 合法", intentGreeting, true}, - {"chitchat 合法", intentChitchat, true}, - {"question 合法", intentQuestion, true}, - {"identity 合法", intentIdentity, true}, - {"meta 合法", intentMeta, true}, - {"空字符串非法", "", false}, - {"未知值非法", "tool_call", false}, - {"拼写错误非法", "questions", false}, - } - for _, c := range cases { - t.Run(c.name, func(t *testing.T) { - if got := isValidIntent(c.intent); got != c.want { - t.Errorf("isValidIntent(%q) = %v, want %v", c.intent, got, c.want) - } - }) - } -} - -func TestBuildRewriteHistory(t *testing.T) { - // 构造 InputMsgs: [System, User1, Assistant1, User2, Assistant2, User3(current)] - msgs := []*schema.Message{ - schema.SystemMessage("你是知识助理"), - schema.UserMessage("Go 的接口怎么实现"), - schema.AssistantMessage("接口用 interface 关键字定义...", nil), - schema.UserMessage("错误处理有哪些方式"), - schema.AssistantMessage("Go 用 error 返回值...", nil), - schema.UserMessage("那 context 怎么用"), // UserQuestionIndex = 5 - } - - history := buildRewriteHistory(msgs, 5, 3) - - // 应排除 system 和当前 User3,只保留 User1/Assistant1/User2/Assistant2 - if history == "" { - t.Fatal("history 为空,应该有内容") - } - if !containsAll(history, "Go 的接口怎么实现", "接口用 interface 关键字", "错误处理有哪些方式", "Go 用 error 返回值") { - t.Errorf("history 未包含预期历史内容: %s", history) - } - if containsAll(history, "你是知识助理") { - t.Error("history 不应该包含 system message") - } - if containsAll(history, "那 context 怎么用") { - t.Error("history 不应该包含当前用户问题") - } -} - -func TestBuildRewriteHistory_EmptyMsgs(t *testing.T) { - history := buildRewriteHistory(nil, 0, 3) - if history != "" { - t.Errorf("nil msgs 应返回空,got: %q", history) - } -} - -func TestBuildRewriteHistory_NoPreviousMessages(t *testing.T) { - msgs := []*schema.Message{ - schema.SystemMessage("你好"), - schema.UserMessage("这是第一个问题"), - } - history := buildRewriteHistory(msgs, 1, 3) - if history != "" { - t.Errorf("没有历史消息应返回空,got: %q", history) - } -} - -func TestBuildRewriteHistory_MaxRoundsLimit(t *testing.T) { - // 构造 6 轮历史,maxRounds=2 → 应该只保留最后 2 轮 - msgs := []*schema.Message{ - schema.SystemMessage("sys"), - schema.UserMessage("q1"), schema.AssistantMessage("a1", nil), - schema.UserMessage("q2"), schema.AssistantMessage("a2", nil), - schema.UserMessage("q3"), schema.AssistantMessage("a3", nil), - schema.UserMessage("q4"), schema.AssistantMessage("a4", nil), - schema.UserMessage("q5"), schema.AssistantMessage("a5", nil), - schema.UserMessage("q6"), schema.AssistantMessage("a6", nil), - schema.UserMessage("current"), // idx = 13 - } - history := buildRewriteHistory(msgs, 13, 2) - - if containsAll(history, "q1", "a1") { - t.Error("maxRounds=2 不应包含最早的 q1/a1") - } - if !containsAll(history, "q5", "a5", "q6", "a6") { - t.Errorf("maxRounds=2 应保留最后两轮 q5/a5, q6/a6, got: %s", history) - } -} - -func containsAll(s string, substrs ...string) bool { - for _, sub := range substrs { - if !containsStr(s, sub) { - return false - } - } - return true -} - -func containsStr(s, sub string) bool { - return len(s) >= len(sub) && (s == sub || len(sub) == 0 || indexOf(s, sub) >= 0) -} - -func indexOf(s, sub string) int { - for i := 0; i+len(sub) <= len(s); i++ { - if s[i:i+len(sub)] == sub { - return i - } - } - return -1 -} - -// --- doRewriteWithLLM 单元测试 --- - -func init() { - // 初始化默认 logger,避免 nil pointer - _ = pkgLogger.InitDefault() -} - -func TestDoRewriteWithLLM_Greeting(t *testing.T) { - // mock ChatModel 返回 greeting JSON - mock := &mockChatModel{ - generateFn: func(ctx context.Context, input []*schema.Message, opts ...einoModel.Option) (*schema.Message, error) { - return &schema.Message{Content: `{"rewritten":"","intent":"greeting","keywords":[]}`}, nil - }, - } - ctx := withGraphChatModel(context.Background(), mock) - input := &quickGraphInput{ - OriginalQuery: "你好", - InputMsgs: []*schema.Message{ - schema.SystemMessage("你是知识助理"), - schema.UserMessage("你好"), - }, - UserQuestionIndex: 1, - ModelName: "test-model", - } - - rewritten, intent, _, skipRetrieve := doRewriteWithLLM(ctx, input) - - if intent != intentGreeting { - t.Errorf("intent = %q, want %q", intent, intentGreeting) - } - if !skipRetrieve { - t.Error("greeting 场景 skipRetrieve 应为 true") - } - // 空 rewritten 会被 fallback 成原始 query(合理:改写不能是空的) - if rewritten != "你好" { - t.Errorf("greeting 场景 rewritten 应 fallback 为原始 query, got %q", rewritten) - } -} - -func TestDoRewriteWithLLM_QuestionWithHistory(t *testing.T) { - // mock ChatModel 返回 question JSON,做指代消解 - mock := &mockChatModel{ - generateFn: func(ctx context.Context, input []*schema.Message, opts ...einoModel.Option) (*schema.Message, error) { - // 验证输入里包含历史上下文 - if len(input) != 2 { - t.Errorf("input msgs 应为 [system, user], got len=%d", len(input)) - } - return &schema.Message{Content: `{"rewritten":"Go 中 interface 和 context 怎么配合使用","intent":"question","keywords":["Go","interface","context"]}`}, nil - }, - } - ctx := withGraphChatModel(context.Background(), mock) - input := &quickGraphInput{ - OriginalQuery: "那 context 呢", - InputMsgs: []*schema.Message{ - schema.SystemMessage("你是知识助理"), - schema.UserMessage("Go 怎么实现接口"), - schema.AssistantMessage("用 interface 关键字...", nil), - schema.UserMessage("那 context 呢"), - }, - UserQuestionIndex: 3, - ModelName: "test-model", - } - - rewritten, intent, keywords, skipRetrieve := doRewriteWithLLM(ctx, input) - - if intent != intentQuestion { - t.Errorf("intent = %q, want %q", intent, intentQuestion) - } - if skipRetrieve { - t.Error("question 场景 skipRetrieve 应为 false") - } - if rewritten != "Go 中 interface 和 context 怎么配合使用" { - t.Errorf("rewritten = %q, want 改写后的完整问题", rewritten) - } - if len(keywords) != 3 { - t.Errorf("keywords len = %d, want 3", len(keywords)) - } -} - -func TestDoRewriteWithLLM_Fallback_NoChatModel(t *testing.T) { - // context 里没有 ChatModel → 应 fallback 原始 query - input := &quickGraphInput{ - OriginalQuery: "那 context 呢", - InputMsgs: []*schema.Message{ - schema.SystemMessage("你是知识助理"), - schema.UserMessage("那 context 呢"), - }, - UserQuestionIndex: 1, - ModelName: "test-model", - } - - rewritten, intent, _, skipRetrieve := doRewriteWithLLM(context.Background(), input) - - if rewritten != input.OriginalQuery { - t.Errorf("无 ChatModel 时 rewritten 应等于原始 query, got %q", rewritten) - } - if intent != intentQuestion { - t.Errorf("无 ChatModel 时 intent 应为 question, got %q", intent) - } - if skipRetrieve { - t.Error("无 ChatModel 时不应跳过检索") - } -} - -func TestDoRewriteWithLLM_Fallback_GenerateError(t *testing.T) { - // ChatModel.Generate 报错 → 应 fallback - mock := &mockChatModel{ - generateFn: func(ctx context.Context, input []*schema.Message, opts ...einoModel.Option) (*schema.Message, error) { - return nil, fmt.Errorf("LLM timeout") - }, - } - ctx := withGraphChatModel(context.Background(), mock) - input := &quickGraphInput{ - OriginalQuery: "测试问题", - InputMsgs: []*schema.Message{ - schema.SystemMessage("你是知识助理"), - schema.UserMessage("测试问题"), - }, - UserQuestionIndex: 1, - ModelName: "test-model", - } - - rewritten, intent, _, skipRetrieve := doRewriteWithLLM(ctx, input) - - if rewritten != input.OriginalQuery { - t.Errorf("Generate 报错时应 fallback 原始 query, got %q", rewritten) - } - if intent != intentQuestion { - t.Errorf("Generate 报错时 intent 应为 question, got %q", intent) - } - if skipRetrieve { - t.Error("Generate 报错时不应跳过检索") - } -} - -func TestDoRewriteWithLLM_Fallback_JSONParseError(t *testing.T) { - // LLM 返回无效 JSON → 应 fallback - mock := &mockChatModel{ - generateFn: func(ctx context.Context, input []*schema.Message, opts ...einoModel.Option) (*schema.Message, error) { - return &schema.Message{Content: "不是有效的 JSON"}, nil - }, - } - ctx := withGraphChatModel(context.Background(), mock) - input := &quickGraphInput{ - OriginalQuery: "测试问题", - InputMsgs: []*schema.Message{ - schema.SystemMessage("你是知识助理"), - schema.UserMessage("测试问题"), - }, - UserQuestionIndex: 1, - ModelName: "test-model", - } - - rewritten, intent, _, _ := doRewriteWithLLM(ctx, input) - - if rewritten != input.OriginalQuery { - t.Errorf("JSON 解析失败时应 fallback 原始 query, got %q", rewritten) - } - if intent != intentQuestion { - t.Errorf("JSON 解析失败时 intent 应为 question, got %q", intent) - } -} - -func TestDoRewriteWithLLM_Chitchat(t *testing.T) { - mock := &mockChatModel{ - generateFn: func(ctx context.Context, input []*schema.Message, opts ...einoModel.Option) (*schema.Message, error) { - return &schema.Message{Content: `{"rewritten":"","intent":"chitchat","keywords":[]}`}, nil - }, - } - ctx := withGraphChatModel(context.Background(), mock) - input := &quickGraphInput{ - OriginalQuery: "讲个笑话", - InputMsgs: []*schema.Message{ - schema.SystemMessage("你是知识助理"), - schema.UserMessage("讲个笑话"), - }, - UserQuestionIndex: 1, - ModelName: "test-model", - } - - _, intent, _, skipRetrieve := doRewriteWithLLM(ctx, input) - - if intent != intentChitchat { - t.Errorf("intent = %q, want %q", intent, intentChitchat) - } - if !skipRetrieve { - t.Error("chitchat 场景 skipRetrieve 应为 true") - } -} - -// --- QuickAnswerGraph 集成测试 --- - -// 注:完整 QuickAnswerGraph 端到端测试需要 EinoRetrieverAdapter + 真实 Retriever + Runnable.Compile/Invoke 流程, -// 通过 HTTP API 层测试更合适。这里 doRewriteWithLLM 单元测试已覆盖 Rewrite 节点核心逻辑。 -// 完整链路:HTTP /metrics 端点已在运行时验证(greeting → skip_retrieve=true、question → rewrite → retrieve → generate) diff --git a/internal/service/chat_service_mapper.go b/internal/service/chat_service_mapper.go index ac99884..cc6dadc 100644 --- a/internal/service/chat_service_mapper.go +++ b/internal/service/chat_service_mapper.go @@ -11,7 +11,7 @@ import ( // sessionResponse 转换会话响应 DTO func sessionResponse(session entity.ChatSession) response.SessionResponse { - return response.SessionResponse{ + resp := response.SessionResponse{ ID: session.ID, Title: session.Title, ModelID: session.ModelID, @@ -19,6 +19,16 @@ func sessionResponse(session entity.ChatSession) response.SessionResponse { CreatedAt: session.CreatedAt, UpdatedAt: session.UpdatedAt, } + if pc, err := session.GetPendingCheckpoint(); err == nil && pc != nil { + resp.PendingCheckpoint = &response.PendingCheckpointInfo{ + CheckpointID: pc.CheckpointID, + InterruptID: pc.InterruptID, + Question: pc.Question, + ToolName: pc.ToolName, + SetAt: pc.SetAt, + } + } + return resp } // messageResponse 转换消息响应 DTO diff --git a/internal/service/chat_service_mode.go b/internal/service/chat_service_mode.go index 2fa8126..9a4274d 100644 --- a/internal/service/chat_service_mode.go +++ b/internal/service/chat_service_mode.go @@ -116,6 +116,26 @@ func (s *chatService) processDeepMode(ctx context.Context, userID, sessionID, us WithProfile(enhancedCtx.Profile). WithPreference(enhancedCtx.Preference) agentReq := agentPB.BuildAgentRequestFields(userID, req.Content, req.ModelID, req.ModelType, req.KnowledgeBaseIDs, history) + agentReq.SessionID = sessionID + + // ── 恢复流程:session 有待恢复 checkpoint 且用户带了审批结果 ── + session2, _ := s.sessionRepo.FindByID(ctx, sessionID) + if session2 != nil && session2.HasPendingCheckpoint() { + pc, _ := session2.GetPendingCheckpoint() + if pc != nil { + logger.Infof("[ChatService] 检测到 pending checkpoint: checkpointID=%s, interruptID=%s", pc.CheckpointID, pc.InterruptID) + if req.Content != "" { + agentReq.CheckpointID = pc.CheckpointID + agentReq.ResumeData = map[string]any{ + pc.InterruptID: req.Content, + } + logger.Infof("[ChatService] 设置恢复参数: checkpointID=%s, resumeKeys=%v", pc.CheckpointID, []string{pc.InterruptID}) + } else { + logger.Warnf("[ChatService] 有 pending checkpoint 但用户未提供审批内容,走首次执行") + _ = s.sessionRepo.ClearPendingCheckpoint(ctx, sessionID) + } + } + } t1 := time.Now() agentEventCh, err := s.agentEngine.Execute(deepCtx, agentReq, chatModel) if err != nil { @@ -148,6 +168,29 @@ func (s *chatService) processDeepMode(ctx context.Context, userID, sessionID, us continue } + // ── interrupt 事件:存 checkpoint 到 session,返回中断事件给前端 ── + if agentEvent.Type == agent.EventInterrupt { + logger.Infof("[ChatService] Agent 中断: checkpointID=%s, interruptID=%s", agentEvent.CheckpointID, agentEvent.InterruptID) + if agentEvent.CheckpointID != "" { + pcData := &entity.PendingCheckpointData{ + CheckpointID: agentEvent.CheckpointID, + InterruptID: agentEvent.InterruptID, + Question: agentEvent.Detail, + ToolName: "", + SetAt: time.Now(), + } + raw, _ := json.Marshal(pcData) + if sErr := s.sessionRepo.SetPendingCheckpoint(ctx, sessionID, raw); sErr != nil { + logger.Errorf("存储 pending checkpoint 失败: %v", sErr) + } else { + logger.Infof("[ChatService] 已存储 pending checkpoint: sessionID=%s", sessionID) + } + } + eventCh <- toStreamEvent(agentEvent) + // 中断后不保存 assistant message(因为还没执行完) + return + } + if agentEvent.Type == agent.EventAnswer { fullContent += agentEvent.Content } @@ -173,6 +216,15 @@ func (s *chatService) processDeepMode(ctx context.Context, userID, sessionID, us } applyReasoningStep(&reasoningSteps, agentEvent) } + + // ── 执行成功后清除 pending checkpoint(如果有的话) ── + if agentReq.CheckpointID != "" { + if cErr := s.sessionRepo.ClearPendingCheckpoint(ctx, sessionID); cErr != nil { + logger.Warnf("清除 pending checkpoint 失败: %v", cErr) + } else { + logger.Infof("[ChatService] 恢复执行完成,已清除 pending checkpoint: sessionID=%s", sessionID) + } + } if obsOk { s.obs.Observe(ctx, "chat_deep_agent_seconds", map[string]string{"model_id": req.ModelID}, time.Since(t1).Seconds()) s.obs.Incr(ctx, "agent_runs_total", map[string]string{ @@ -408,6 +460,12 @@ func toStreamEvent(e agent.Event) dto.StreamEvent { se.CitationFileName = e.CitationFileName se.CitationContent = e.CitationContent } + // interrupt 事件:携带 checkpointID / interruptID / info + if e.Type == agent.EventInterrupt { + se.CheckpointID = e.CheckpointID + se.InterruptID = e.InterruptID + se.InterruptInfo = e.InterruptInfo + } return se } From 66c7b41f83dbae31dce7a90f41d5f35f0b36ed54 Mon Sep 17 00:00:00 2001 From: st <2663600842@qq.com> Date: Sat, 8 Aug 2026 21:22:38 +0800 Subject: [PATCH 08/10] =?UTF-8?q?feat(=E5=89=8D=E7=AB=AF):=20=E5=8D=B1?= =?UTF-8?q?=E9=99=A9=E5=B7=A5=E5=85=B7=E5=AE=A1=E6=89=B9=E5=8D=A1=20UI=20+?= =?UTF-8?q?=20=E6=81=A2=E5=A4=8D=E6=B5=81=E7=A8=8B=E8=BF=9E=E7=BB=AD?= =?UTF-8?q?=E5=8C=96=20+=20=E6=8E=A8=E7=90=86=E6=AD=A5=E9=AA=A4=E4=B8=8D?= =?UTF-8?q?=E4=B8=AD=E6=96=AD?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - useChat.ts: 新增 pendingApproval 状态、approvePending / cancelApproval 操作 - interrupt 事件处理: 不 push 多余纯文本气泡,记录 interruptedAssistantId 供恢复复用 - sendMessage 新增 isResume 参数: 恢复流程不 push 用户气泡(避免审批当成新提问) - done 事件: 更新已有 assistant 消息块而非新建,保持审批前后在同一块 - streamTimeline: interrupt 时不清空,恢复执行后继续累加,done 时合并为完整步骤列表 - ChatPage.vue: 新增 amber 色调审批卡 UI(需要人工确认 / 同意执行 / 拒绝 / 暂不处理) - types/chat.ts: StreamEvent 新增 checkpoint_id / interrupt_id / interrupt_info;ChatSession 新增 pending_checkpoint --- design/vue/src/composables/useChat.ts | 146 ++++++++++++++++++++++---- design/vue/src/pages/ChatPage.vue | 33 ++++++ design/vue/src/types/chat.ts | 32 ++++++ 3 files changed, 189 insertions(+), 22 deletions(-) diff --git a/design/vue/src/composables/useChat.ts b/design/vue/src/composables/useChat.ts index 2224ced..ce4e586 100644 --- a/design/vue/src/composables/useChat.ts +++ b/design/vue/src/composables/useChat.ts @@ -1,4 +1,4 @@ -import { ref, computed, nextTick, inject } from 'vue' +import { ref, computed, nextTick, inject, watch } from 'vue' import { useRouter } from 'vue-router' import { ElMessage } from 'element-plus' import { marked } from 'marked' @@ -6,7 +6,7 @@ import * as chatApi from '@/api/chat' import * as modelApi from '@/api/model' import * as authApi from '@/api/auth' import { request } from '@/api/client' -import type { ChatSession, FeedbackRequest } from '@/types/chat' +import type { ChatSession, FeedbackRequest, PendingApproval } from '@/types/chat' import type { StreamEvent } from '@/types/chat' // ── UI 展示用的本地类型 ── @@ -105,6 +105,11 @@ export function useChat() { // ── 中断控制 ── let abortController: AbortController | null = null + // ── 审批状态(危险工具中断) ── + const pendingApproval = ref(null) + // interrupt 事件所在的 assistant 消息块 ID,恢复流程复用同一个 + let interruptedAssistantId = '' + // ── 计算属性 ── const activeSession = computed(() => sessions.value.find((s) => s.id === activeSessionId.value), @@ -241,8 +246,28 @@ export function useChat() { loadMessages(sessionId) } + // 切换会话时恢复/清除审批卡状态 + watch( + () => activeSession.value, + (sess) => { + const pc = sess?.pending_checkpoint + if (pc && pc.checkpoint_id) { + pendingApproval.value = { + checkpoint_id: pc.checkpoint_id, + interrupt_id: pc.interrupt_id, + title: '需要人工确认', + detail: pc.question ?? '执行被中断,等待用户审批', + tool_name: pc.tool_name, + } + } else { + pendingApproval.value = null + } + }, + { immediate: true }, + ) + // ── 发送消息(SSE 流式) ── - async function sendMessage() { + async function sendMessage(displayText?: string, isResume = false) { const content = input.value.trim() if (!content || isLoading.value) return @@ -251,15 +276,23 @@ export function useChat() { return } - // 延迟清空输入框,避免跳转后问题丢失 - // 如果是新会话,等会话创建成功后再清空 const isNewSession = !activeSessionId.value - if (!isNewSession) { - // 已有会话,立即清空 + + if (isResume) { + // 恢复流程:不 push 用户气泡,审批内容不是新的提问 input.value = '' + } else { + const display = displayText ?? content + + // 延迟清空输入框,避免跳转后问题丢失 + if (!isNewSession) { + // 已有会话,立即清空 + input.value = '' + } + + messages.value.push({ id: 'u-' + Date.now(), role: 'user', content: display }) } - messages.value.push({ id: 'u-' + Date.now(), role: 'user', content }) isLoading.value = true progressText.value = '' streamContent.value = '' @@ -336,7 +369,13 @@ export function useChat() { switch (evt.type) { case 'start': - if (evt.message_id) assistantId = evt.message_id + // 恢复流程:复用 interrupt 时的 assistant 块 ID,保持在同一块里 + if (interruptedAssistantId) { + assistantId = interruptedAssistantId + interruptedAssistantId = '' + } else if (evt.message_id) { + assistantId = evt.message_id + } if (evt.sources) finalSources = evt.sources break @@ -415,6 +454,41 @@ export function useChat() { }) return + case 'interrupt': { + isLoading.value = false + progressText.value = '' + streamContent.value = '' + streamSources.value = [] + // streamTimeline 不清空,interrupt 前的步骤保留,恢复后继续累加 + const info = evt.interrupt_info ?? {} + const approval: PendingApproval = { + checkpoint_id: evt.checkpoint_id ?? '', + interrupt_id: evt.interrupt_id ?? '', + title: '需要人工确认', + detail: evt.detail ?? (info?.message as string) ?? '执行被中断,等待用户处理', + tool_name: (info?.tool_name as string) ?? '', + target_ref: (info?.target_ref as string) ?? '', + reason: (info?.reason as string) ?? '', + } + pendingApproval.value = approval + // 记录 assistant 块 ID,恢复时 done 事件复用同一块 + interruptedAssistantId = assistantId || 'a-' + Date.now() + return + } + + case 'clarify': { + isLoading.value = false + progressText.value = '' + const q = evt.clarify?.question ?? evt.detail ?? '' + const opts = evt.clarify?.options ?? [] + messages.value.push({ + id: 'c-' + Date.now(), + role: 'assistant', + content: q, + }) + break + } + case 'done': streamTimeline.value.forEach( (s) => s.status === 'running' && (s.status = 'success'), @@ -423,19 +497,31 @@ export function useChat() { finalSources = evt.sources streamSources.value = evt.sources } - messages.value.push({ - id: assistantId || 'a-' + Date.now(), - role: 'assistant', - content: finalContent, - sources: finalSources, - timeline: - streamTimeline.value.length > 0 - ? [...streamTimeline.value] - : undefined, - trace_id: traceId, - }) - if (streamTimeline.value.length > 0) { - collapsedTimelines.value.add(messages.value.length - 1) + const doneTimeline = streamTimeline.value.length > 0 ? [...streamTimeline.value] : undefined + const finalAssistantId = assistantId || 'a-' + Date.now() + // 恢复流程:assistantId 已存在(interrupt 时 push 过),更新那条而不是新建 + const existingIdx = messages.value.findIndex((m) => m.id === finalAssistantId) + if (existingIdx >= 0) { + const updated = [...messages.value] + updated[existingIdx] = { + ...updated[existingIdx], + content: finalContent, + sources: finalSources.length > 0 ? finalSources : updated[existingIdx].sources, + timeline: doneTimeline ?? updated[existingIdx].timeline, + trace_id: traceId, + } + messages.value = updated + if (doneTimeline) collapsedTimelines.value.add(existingIdx) + } else { + messages.value.push({ + id: finalAssistantId, + role: 'assistant', + content: finalContent, + sources: finalSources, + timeline: doneTimeline, + trace_id: traceId, + }) + if (doneTimeline) collapsedTimelines.value.add(messages.value.length - 1) } isLoading.value = false streamContent.value = '' @@ -561,6 +647,19 @@ export function useChat() { } } + // ── 危险工具审批 ── + function approvePending(resolution: 'approve' | 'reject') { + if (!pendingApproval.value) return + input.value = resolution // 请求内容 + pendingApproval.value = null + void sendMessage(undefined, true) // isResume=true: 不 push 用户气泡 + } + + function cancelApproval() { + pendingApproval.value = null + interruptedAssistantId = '' + } + // ── 反馈 ── // 提交消息反馈 async function submitFeedback( @@ -740,6 +839,9 @@ export function useChat() { submitFeedback, newChat, cleanTooltipText, + pendingApproval, + approvePending, + cancelApproval, } } diff --git a/design/vue/src/pages/ChatPage.vue b/design/vue/src/pages/ChatPage.vue index c3a06e2..0790149 100644 --- a/design/vue/src/pages/ChatPage.vue +++ b/design/vue/src/pages/ChatPage.vue @@ -170,6 +170,38 @@ + + + + + + + + + + 需要人工确认 + {{ pendingApproval.tool_name }} + + + {{ pendingApproval.detail }} + + 同意执行 + 拒绝 + 暂不处理 + + + + + @@ -328,6 +360,7 @@ const { modelOptions, knowledgeBases, connected, input, selectedModel, selectedKBs, searchMode, kbTriggerText, init, sendMessage, scrollToBottom, toggleKB, formatContent, getSourceChunkIds, copyText, regenerate, retryLastMessage, stopGeneration, selectSession, loadSessions, newChat, cleanTooltipText, submitFeedback, + pendingApproval, approvePending, cancelApproval, } = chat const chatEl = ref() diff --git a/design/vue/src/types/chat.ts b/design/vue/src/types/chat.ts index 90772e6..84fc153 100644 --- a/design/vue/src/types/chat.ts +++ b/design/vue/src/types/chat.ts @@ -1,10 +1,19 @@ // ── Session ── +export interface PendingCheckpointInfo { + checkpoint_id: string + interrupt_id: string + question?: string + tool_name?: string + set_at: string +} + export interface ChatSession { id: string title: string model_id: string status: string + pending_checkpoint?: PendingCheckpointInfo | null created_at: string updated_at: string } @@ -67,6 +76,11 @@ export interface SendMessageRequest { // ── SSE Stream Event ── +export interface ClarifyPayload { + question: string + options?: string[] +} + export interface StreamEvent { type: string title?: string @@ -82,6 +96,24 @@ export interface StreamEvent { done?: boolean error?: string retryable?: boolean + // clarify 事件字段:追问 + clarify?: ClarifyPayload + // interrupt 事件字段:中断等待用户审批 + checkpoint_id?: string + interrupt_id?: string + interrupt_info?: Record +} + +// ── 审批请求状态(前端本地) ── + +export interface PendingApproval { + checkpoint_id: string + interrupt_id: string + title: string + detail: string + tool_name?: string + target_ref?: string + reason?: string } // ── List Responses ── From fed2922ede09b5f5ba61c49b988f79b64b4e650a Mon Sep 17 00:00:00 2001 From: st <2663600842@qq.com> Date: Sat, 8 Aug 2026 21:22:49 +0800 Subject: [PATCH 09/10] =?UTF-8?q?chore:=20=E6=B8=85=E7=90=86=E4=B8=8D?= =?UTF-8?q?=E5=86=8D=E9=9C=80=E8=A6=81=E7=9A=84=E6=B5=8B=E8=AF=95=E6=96=87?= =?UTF-8?q?=E4=BB=B6=E5=92=8C=20SQL=20schema=20=E6=96=87=E4=BB=B6?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 删除 scripts/init_knowledge_schema.sql(用 GORM AutoMigrate 替代) - 删除 internal/service/chat_service_graph_quick_test.go(依赖未完成,暂不维护) - 删除 internal/tool/providers/http_provider_test.go - 删除 internal/integration/dingtalk/client_test.go --- internal/integration/dingtalk/client_test.go | 235 ---------- internal/tool/providers/http_provider_test.go | 68 --- scripts/init_knowledge_schema.sql | 424 ------------------ 3 files changed, 727 deletions(-) delete mode 100644 internal/integration/dingtalk/client_test.go delete mode 100644 internal/tool/providers/http_provider_test.go delete mode 100644 scripts/init_knowledge_schema.sql diff --git a/internal/integration/dingtalk/client_test.go b/internal/integration/dingtalk/client_test.go deleted file mode 100644 index 6543614..0000000 --- a/internal/integration/dingtalk/client_test.go +++ /dev/null @@ -1,235 +0,0 @@ -package dingtalk - -import ( - "context" - "encoding/json" - "net/http" - "net/http/httptest" - "strings" - "testing" - "time" - - "solvify-agent/pkg/config" -) - -// TestNodeUnmarshalModifiedTimeFormats 验证节点更新时间兼容时间戳和分钟精度时间 -func TestNodeUnmarshalModifiedTimeFormats(t *testing.T) { - tests := []struct { - name string - value string - expected int64 - }{ - {name: "毫秒时间戳", value: `"1719999999000"`, expected: 1719999999000}, - {name: "分钟精度时间", value: `"2026-07-01T19:22Z"`, expected: time.Date(2026, 7, 1, 19, 22, 0, 0, time.UTC).Unix()}, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - var node Node - if err := json.Unmarshal([]byte(`{"modifiedTime":`+tt.value+`}`), &node); err != nil { - t.Fatalf("解析节点更新时间失败: %v", err) - } - if node.ModifiedAt != tt.expected { - t.Fatalf("节点更新时间不符合预期: got=%d want=%d", node.ModifiedAt, tt.expected) - } - }) - } -} - -// TestClientListNodesUsesHeaderToken 验证节点列表使用 Header 鉴权和分页参数 -func TestClientListNodesUsesHeaderToken(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - switch r.URL.Path { - case "/v1.0/oauth2/accessToken": - if r.Method != http.MethodPost { - t.Fatalf("accessToken 请求方法错误: %s", r.Method) - } - if !strings.Contains(r.Header.Get("Content-Type"), "application/json") { - t.Fatalf("accessToken 请求体类型错误") - } - _, _ = w.Write([]byte(`{"accessToken":"token-1","expireIn":7200}`)) - case "/v2.0/wiki/nodes": - if r.Header.Get("x-acs-dingtalk-access-token") != "token-1" { - t.Fatalf("未使用钉钉 Header 鉴权") - } - if r.URL.Query().Get("parentNodeId") != "root-1" || r.URL.Query().Get("nextToken") != "next-1" { - t.Fatalf("节点列表分页参数错误: %s", r.URL.RawQuery) - } - _, _ = w.Write([]byte(`{"nodes":[{"nodeId":"node-1","workspaceId":"ws-1","name":"a.md","size":"12","type":"FILE","modifiedTime":"1719999999000"}],"nextToken":"next-2"}`)) - default: - t.Fatalf("未预期的请求路径: %s", r.URL.Path) - } - })) - defer server.Close() - - client := NewClient(config.DingTalkConfig{AppKey: "app-key", AppSecret: "app-secret"}) - client.httpClient = server.Client() - client.accessTokenURL = server.URL + "/v1.0/oauth2/accessToken" - client.apiBaseURL = server.URL - - nodes, nextToken, err := client.ListNodes(context.Background(), "union-1", "root-1", "next-1", 50) - if err != nil { - t.Fatalf("获取节点列表失败: %v", err) - } - if len(nodes) != 1 || nodes[0].NodeID != "node-1" || nodes[0].Size != 12 || nodes[0].ModifiedAt != 1719999999000 || nextToken != "next-2" { - t.Fatalf("节点列表响应解析错误: nodes=%v next=%s", nodes, nextToken) - } -} - -// TestClientQueryDentryIDEscapesPath 验证 dentryUuid 路径参数会转义 -func TestClientQueryDentryIDEscapesPath(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - switch r.URL.Path { - case "/v1.0/oauth2/accessToken": - _, _ = w.Write([]byte(`{"accessToken":"token-1","expireIn":7200}`)) - case "/v2.0/doc/dentries/abc/def/queryDentryId": - if !strings.Contains(r.URL.RawPath, "abc%2Fdef") && !strings.Contains(r.RequestURI, "abc%2Fdef") { - t.Fatalf("dentryUuid 未正确转义: %s", r.RequestURI) - } - _, _ = w.Write([]byte(`{"dentryUuid":"abc/def","dentryId":"d-1","spaceId":"s-1"}`)) - default: - t.Fatalf("未预期的请求路径: %s", r.URL.Path) - } - })) - defer server.Close() - - client := NewClient(config.DingTalkConfig{AppKey: "app-key", AppSecret: "app-secret"}) - client.httpClient = server.Client() - client.accessTokenURL = server.URL + "/v1.0/oauth2/accessToken" - client.apiBaseURL = server.URL - - output, err := client.QueryDentryID(context.Background(), "union-1", "abc/def") - if err != nil { - t.Fatalf("查询 dentryId 失败: %v", err) - } - if output.SpaceID != "s-1" || output.DentryID != "d-1" { - t.Fatalf("dentryId 响应解析错误: %+v", output) - } -} - -// TestClientDownloadFileUsesReturnedHeaders 验证下载文件使用钉钉返回的签名 Header -func TestClientDownloadFileUsesReturnedHeaders(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - switch r.URL.Path { - case "/v1.0/oauth2/accessToken": - _, _ = w.Write([]byte(`{"accessToken":"token-1","expireIn":7200}`)) - case "/v1.0/storage/spaces/s-1/dentries/d-1/downloadInfos/query": - if r.Header.Get("x-acs-dingtalk-access-token") != "token-1" { - t.Fatalf("下载信息未使用钉钉 Header 鉴权") - } - _, _ = w.Write([]byte(`{"protocol":"HEADER_SIGNATURE","headerSignatureInfo":{"resourceUrls":["` + serverURL(r) + `/download"],"headers":{"X-Sign":"ok"}}}`)) - case "/download": - if r.Header.Get("X-Sign") != "ok" { - t.Fatalf("文件下载未携带签名 Header") - } - _, _ = w.Write([]byte("hello")) - default: - t.Fatalf("未预期的请求路径: %s", r.URL.Path) - } - })) - defer server.Close() - - client := NewClient(config.DingTalkConfig{AppKey: "app-key", AppSecret: "app-secret"}) - client.httpClient = server.Client() - client.accessTokenURL = server.URL + "/v1.0/oauth2/accessToken" - client.apiBaseURL = server.URL - - data, hash, err := client.DownloadFile(context.Background(), "union-1", "s-1", "d-1") - if err != nil { - t.Fatalf("下载文件失败: %v", err) - } - if string(data) != "hello" || hash == "" { - t.Fatalf("下载内容解析错误: data=%q hash=%s", string(data), hash) - } -} - -// TestClientQueryDocumentBlocks 验证在线文档块元素查询参数和响应 -func TestClientQueryDocumentBlocks(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - switch r.URL.Path { - case "/v1.0/oauth2/accessToken": - _, _ = w.Write([]byte(`{"accessToken":"token-1","expireIn":7200}`)) - case "/v1.0/doc/suites/documents/doc-1/blocks": - if r.Method != http.MethodGet || r.URL.Query().Get("operatorId") != "union-1" { - t.Fatalf("块元素查询参数错误: %s %s", r.Method, r.URL.RawQuery) - } - _, _ = w.Write([]byte(`{"result":{"data":[{"blockType":"paragraph","paragraph":{"text":"正文"}}]},"success":true}`)) - default: - t.Fatalf("未预期的请求路径: %s", r.URL.Path) - } - })) - defer server.Close() - - client := NewClient(config.DingTalkConfig{AppKey: "app-key", AppSecret: "app-secret"}) - client.httpClient = server.Client() - client.accessTokenURL = server.URL + "/v1.0/oauth2/accessToken" - client.apiBaseURL = server.URL - - blocks, err := client.QueryDocumentBlocks(context.Background(), "union-1", "doc-1") - if err != nil { - t.Fatalf("查询在线文档块元素失败: %v", err) - } - if len(blocks) != 1 || blocks[0]["blockType"] != "paragraph" { - t.Fatalf("块元素响应解析错误: %+v", blocks) - } -} - -// serverURL 从测试请求还原服务地址 -func serverURL(r *http.Request) string { - return "http://" + r.Host -} - -// TestClientExchangeUserAccessToken 验证扫码授权码使用新版 OAuth JSON 请求兑换用户 token -func TestClientExchangeUserAccessToken(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path != "/v1.0/oauth2/userAccessToken" { - t.Fatalf("未预期的请求路径: %s", r.URL.Path) - } - if r.Method != http.MethodPost { - t.Fatalf("用户 token 请求方法错误: %s", r.Method) - } - if !strings.Contains(r.Header.Get("Content-Type"), "application/json") { - t.Fatalf("用户 token 请求体类型错误") - } - _, _ = w.Write([]byte(`{"accessToken":"user-token","refreshToken":"refresh-token","expireIn":7200,"corpId":"corp-1"}`)) - })) - defer server.Close() - - client := NewClient(config.DingTalkConfig{AppKey: "app-key", AppSecret: "app-secret"}) - client.httpClient = server.Client() - client.apiBaseURL = server.URL - - output, err := client.ExchangeUserAccessToken(context.Background(), "auth-code") - if err != nil { - t.Fatalf("兑换用户 token 失败: %v", err) - } - if output.AccessToken != "user-token" || output.CorpID != "corp-1" { - t.Fatalf("用户 token 响应解析错误: %+v", output) - } -} - -// TestClientGetCurrentUserInfoUsesUserToken 验证用户信息接口使用个人 token Header -func TestClientGetCurrentUserInfoUsesUserToken(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path != "/v1.0/contact/users/me" { - t.Fatalf("未预期的请求路径: %s", r.URL.Path) - } - if r.Header.Get("x-acs-dingtalk-access-token") != "user-token" { - t.Fatalf("用户信息接口未使用个人 token Header") - } - _, _ = w.Write([]byte(`{"nick":"张三","avatarUrl":"https://example.com/a.png","openId":"open-1","unionId":"union-1","email":"a@example.com"}`)) - })) - defer server.Close() - - client := NewClient(config.DingTalkConfig{AppKey: "app-key", AppSecret: "app-secret"}) - client.httpClient = server.Client() - client.apiBaseURL = server.URL - - output, err := client.GetCurrentUserInfo(context.Background(), "user-token") - if err != nil { - t.Fatalf("获取用户信息失败: %v", err) - } - if output.UnionID != "union-1" || output.OpenID != "open-1" { - t.Fatalf("用户信息响应解析错误: %+v", output) - } -} diff --git a/internal/tool/providers/http_provider_test.go b/internal/tool/providers/http_provider_test.go deleted file mode 100644 index 3fcd014..0000000 --- a/internal/tool/providers/http_provider_test.go +++ /dev/null @@ -1,68 +0,0 @@ -package providers - -import ( - "encoding/json" - "testing" -) - -func TestMapResponse(t *testing.T) { - resp := []byte(`{ - "code": 200, - "data": { - "webPages": { - "value": [ - {"name": "Java 官网", "url": "https://java.com"} - ] - } - } - }`) - - mapping := map[string]string{ - "results": "$.data.webPages.value", - } - - p := &HTTPProvider{} - mapped, err := p.mapResponse(resp, mapping) - if err != nil { - t.Fatalf("mapResponse failed: %v", err) - } - - var result map[string]interface{} - if err := json.Unmarshal([]byte(mapped), &result); err != nil { - t.Fatalf("mapped result is not valid JSON: %v", err) - } - - results, ok := result["results"].([]interface{}) - if !ok { - t.Fatalf("expected results to be array, got %T", result["results"]) - } - if len(results) != 1 { - t.Fatalf("expected 1 result, got %d", len(results)) - } -} - -func TestSanitizeURL(t *testing.T) { - got := sanitizeURL(" `https://api.tavily.com/search` ") - if got != "https://api.tavily.com/search" { - t.Fatalf("expected sanitized URL, got %q", got) - } -} - -func TestParseJSONPath(t *testing.T) { - parts := parseJSONPath("data.webPages.value[0].name") - if len(parts) != 5 { - t.Fatalf("expected 5 parts, got %d: %+v", len(parts), parts) - } - expected := []jsonPathPart{ - {Key: "data"}, - {Key: "webPages"}, - {Key: "value"}, - {Index: 0}, - {Key: "name"}, - } - for i, e := range expected { - if parts[i].Key != e.Key || parts[i].Index != e.Index { - t.Errorf("part %d mismatch: expected %+v, got %+v", i, e, parts[i]) - } - } -} diff --git a/scripts/init_knowledge_schema.sql b/scripts/init_knowledge_schema.sql deleted file mode 100644 index c730294..0000000 --- a/scripts/init_knowledge_schema.sql +++ /dev/null @@ -1,424 +0,0 @@ -CREATE -EXTENSION IF NOT EXISTS pgcrypto; -CREATE -EXTENSION IF NOT EXISTS vector; - --- 用户表作为当前阶段的数据隔离边界 -CREATE TABLE IF NOT EXISTS users -( - id UUID PRIMARY KEY DEFAULT gen_random_uuid(), -- 用户 ID - username VARCHAR(64) NOT NULL, -- 用户名 - email VARCHAR(128) NOT NULL, -- 邮箱 - password TEXT NOT NULL DEFAULT '', -- 密码哈希 - status INT NOT NULL DEFAULT 1, -- 用户状态,1 正常,2 禁用,3 注销,4 待验证 - created_at TIMESTAMPTZ NOT NULL DEFAULT now(), -- 创建时间 - updated_at TIMESTAMPTZ NOT NULL DEFAULT now(), -- 更新时间 - - CONSTRAINT users_username_unique UNIQUE (username), - CONSTRAINT users_email_unique UNIQUE (email) -); - -COMMENT ON TABLE users IS '用户基础表,用于隔离每个用户自己的知识库和文档'; -COMMENT ON COLUMN users.id IS '用户 ID'; -COMMENT ON COLUMN users.username IS '用户名'; -COMMENT ON COLUMN users.email IS '邮箱'; -COMMENT ON COLUMN users.password IS '密码哈希'; -COMMENT ON COLUMN users.status IS '用户状态,1 正常,2 禁用,3 注销,4 待验证'; -COMMENT ON COLUMN users.created_at IS '创建时间'; -COMMENT ON COLUMN users.updated_at IS '更新时间'; - --- 知识库表保存用户自建、同步和联网搜索知识库 -CREATE TABLE IF NOT EXISTS knowledge_bases -( - id UUID PRIMARY KEY DEFAULT gen_random_uuid(), -- 知识库 ID - user_id UUID NOT NULL, -- 所属用户 ID - name VARCHAR(128) NOT NULL, -- 知识库名称 - category VARCHAR(64) NOT NULL DEFAULT '', -- 知识库分类 - description TEXT DEFAULT '', -- 知识库描述 - source_type VARCHAR(32) NOT NULL DEFAULT 'local', -- 知识库来源类型,local 自建,sync 同步,web_search 联网搜索 - source_platform VARCHAR(32) NOT NULL DEFAULT '', -- 同步来源平台 - document_count INT NOT NULL DEFAULT 0, -- 文档数量 - storage_bytes BIGINT NOT NULL DEFAULT 0, -- 已占用存储字节数 - status INT NOT NULL DEFAULT 1, -- 知识库状态,1 正常,2 已删除 - created_at TIMESTAMPTZ NOT NULL DEFAULT now(), -- 创建时间 - updated_at TIMESTAMPTZ NOT NULL DEFAULT now(), -- 更新时间 - deleted_at TIMESTAMPTZ, -- 删除时间 - delete_expired_at TIMESTAMPTZ -- 删除保留到期时间 -); - -COMMENT ON TABLE knowledge_bases IS '知识库主表,所有知识库按用户隔离'; -COMMENT ON COLUMN knowledge_bases.id IS '知识库 ID'; -COMMENT ON COLUMN knowledge_bases.user_id IS '所属用户 ID'; -COMMENT ON COLUMN knowledge_bases.name IS '知识库名称'; -COMMENT ON COLUMN knowledge_bases.category IS '知识库分类'; -COMMENT ON COLUMN knowledge_bases.description IS '知识库描述'; -COMMENT ON COLUMN knowledge_bases.source_type IS '知识库来源类型,local 自建,sync 同步,web_search 联网搜索'; -COMMENT ON COLUMN knowledge_bases.source_platform IS '同步来源平台'; -COMMENT ON COLUMN knowledge_bases.document_count IS '文档数量'; -COMMENT ON COLUMN knowledge_bases.storage_bytes IS '已占用存储字节数'; -COMMENT ON COLUMN knowledge_bases.status IS '知识库状态,1 正常,2 已删除'; -COMMENT ON COLUMN knowledge_bases.created_at IS '创建时间'; -COMMENT ON COLUMN knowledge_bases.updated_at IS '更新时间'; -COMMENT ON COLUMN knowledge_bases.deleted_at IS '删除时间'; -COMMENT ON COLUMN knowledge_bases.delete_expired_at IS '删除保留到期时间'; - --- 文档表记录上传文件和处理状态 -CREATE TABLE IF NOT EXISTS documents -( - id UUID PRIMARY KEY DEFAULT gen_random_uuid(), -- 文档 ID - user_id UUID NOT NULL, -- 所属用户 ID - knowledge_base_id UUID NOT NULL, -- 所属知识库 ID - title VARCHAR(255) NOT NULL, -- 文档标题 - file_name VARCHAR(255) NOT NULL, -- 原始文件名 - file_type VARCHAR(32) NOT NULL DEFAULT '', -- 文件类型 - file_size BIGINT NOT NULL DEFAULT 0, -- 文件大小字节数 - storage_path TEXT NOT NULL DEFAULT '', -- 文件存储路径 - file_hash VARCHAR(128) NOT NULL DEFAULT '', -- 原始文件内容指纹 - source_type VARCHAR(32) NOT NULL DEFAULT 'upload', -- 文档来源类型,upload 上传,edit 编辑,sync 同步,web_search 联网搜索 - external_id VARCHAR(255) NOT NULL DEFAULT '', -- 外部平台文档 ID - external_url TEXT NOT NULL DEFAULT '', -- 外部平台文档链接 - source_updated_at TIMESTAMPTZ, -- 外部平台更新时间 - status INT NOT NULL DEFAULT 1, -- 文档状态,1 已上传,2 处理中,3 已就绪,4 处理失败,5 已删除 - error_message TEXT NOT NULL DEFAULT '', -- 处理失败原因 - ready_at TIMESTAMPTZ, -- 文档就绪时间 - created_at TIMESTAMPTZ NOT NULL DEFAULT now(), -- 创建时间 - updated_at TIMESTAMPTZ NOT NULL DEFAULT now(), -- 更新时间 - deleted_at TIMESTAMPTZ, -- 删除时间 - delete_expired_at TIMESTAMPTZ -- 删除保留到期时间 - -); - -COMMENT ON TABLE documents IS '文档主表,记录知识库下的文件和处理状态'; -COMMENT ON COLUMN documents.id IS '文档 ID'; -COMMENT ON COLUMN documents.user_id IS '所属用户 ID'; -COMMENT ON COLUMN documents.knowledge_base_id IS '所属知识库 ID'; -COMMENT ON COLUMN documents.title IS '文档标题'; -COMMENT ON COLUMN documents.file_name IS '原始文件名'; -COMMENT ON COLUMN documents.file_type IS '文件类型'; -COMMENT ON COLUMN documents.file_size IS '文件大小字节数'; -COMMENT ON COLUMN documents.storage_path IS '文件存储路径'; -COMMENT ON COLUMN documents.file_hash IS '原始文件内容指纹'; -COMMENT ON COLUMN documents.source_type IS '文档来源类型,upload 上传,edit 编辑,sync 同步,web_search 联网搜索'; -COMMENT ON COLUMN documents.external_id IS '外部平台文档 ID'; -COMMENT ON COLUMN documents.external_url IS '外部平台文档链接'; -COMMENT ON COLUMN documents.source_updated_at IS '外部平台更新时间'; -COMMENT ON COLUMN documents.status IS '文档状态,1 已上传,2 处理中,3 已就绪,4 处理失败,5 已删除'; -COMMENT ON COLUMN documents.error_message IS '处理失败原因'; -COMMENT ON COLUMN documents.ready_at IS '文档就绪时间'; -COMMENT ON COLUMN documents.deleted_at IS '删除时间'; -COMMENT ON COLUMN documents.delete_expired_at IS '删除保留到期时间'; -COMMENT ON COLUMN documents.created_at IS '创建时间'; -COMMENT ON COLUMN documents.updated_at IS '更新时间'; - -ALTER TABLE documents - ADD COLUMN IF NOT EXISTS external_id VARCHAR(255) NOT NULL DEFAULT '', - ADD COLUMN IF NOT EXISTS external_url TEXT NOT NULL DEFAULT '', - ADD COLUMN IF NOT EXISTS source_updated_at TIMESTAMPTZ; - --- 文档版本表支持在线编辑和重新向量化 -CREATE TABLE IF NOT EXISTS document_versions -( - id UUID PRIMARY KEY DEFAULT gen_random_uuid(), -- 版本 ID - user_id UUID NOT NULL, -- 所属用户 ID - document_id UUID NOT NULL, -- 所属文档 ID - version_no INT NOT NULL DEFAULT 1, -- 版本号 - content TEXT NOT NULL DEFAULT '', -- 版本正文内容 - content_hash VARCHAR(128) NOT NULL DEFAULT '', -- 版本内容哈希 - change_summary TEXT NOT NULL DEFAULT '', -- 变更摘要 - created_at TIMESTAMPTZ NOT NULL DEFAULT now(), -- 创建时间 - - CONSTRAINT document_versions_document_version_unique UNIQUE (document_id, version_no) -); - -COMMENT ON TABLE document_versions IS '文档版本表,用于记录原始解析内容和在线编辑历史'; -COMMENT ON COLUMN document_versions.id IS '版本 ID'; -COMMENT ON COLUMN document_versions.user_id IS '所属用户 ID'; -COMMENT ON COLUMN document_versions.document_id IS '所属文档 ID'; -COMMENT ON COLUMN document_versions.version_no IS '版本号'; -COMMENT ON COLUMN document_versions.content IS '版本正文内容'; -COMMENT ON COLUMN document_versions.content_hash IS '版本内容哈希'; -COMMENT ON COLUMN document_versions.change_summary IS '变更摘要'; -COMMENT ON COLUMN document_versions.created_at IS '创建时间'; - --- 文档分块表保存文本分块和向量数据 -CREATE TABLE IF NOT EXISTS document_chunks -( - id UUID PRIMARY KEY DEFAULT gen_random_uuid(), -- 分块 ID - user_id UUID NOT NULL, -- 所属用户 ID - knowledge_base_id UUID NOT NULL, -- 所属知识库 ID - document_id UUID NOT NULL, -- 所属文档 ID - version_id UUID NOT NULL, -- 所属文档版本 ID - chunk_index INT NOT NULL, -- 分块序号 - section_title TEXT NOT NULL DEFAULT '', -- 分块所属章节标题 - content TEXT NOT NULL, -- 分块文本内容 - token_count INT NOT NULL DEFAULT 0, -- 分块 token 数 - page_number INT, -- 来源页码 - embedding_model VARCHAR(128) NOT NULL DEFAULT '', -- 向量模型名称 - embedding vector(1024), -- 分块向量数据 - keywords TEXT[] NOT NULL DEFAULT '{}'::text[], -- 分块关键词 - metadata JSONB NOT NULL DEFAULT '{}'::jsonb, -- 分块扩展元数据 - created_at TIMESTAMPTZ NOT NULL DEFAULT now(), -- 创建时间 - - CONSTRAINT document_chunks_version_index_unique UNIQUE (version_id, chunk_index) -); - -COMMENT ON TABLE document_chunks IS '文档分块表,保存可检索文本片段和 embedding 向量'; -COMMENT ON COLUMN document_chunks.id IS '分块 ID'; -COMMENT ON COLUMN document_chunks.user_id IS '所属用户 ID'; -COMMENT ON COLUMN document_chunks.knowledge_base_id IS '所属知识库 ID'; -COMMENT ON COLUMN document_chunks.document_id IS '所属文档 ID'; -COMMENT ON COLUMN document_chunks.version_id IS '所属文档版本 ID'; -COMMENT ON COLUMN document_chunks.chunk_index IS '分块序号'; -COMMENT ON COLUMN document_chunks.section_title IS '分块所属章节标题'; -COMMENT ON COLUMN document_chunks.content IS '分块文本内容'; -COMMENT ON COLUMN document_chunks.token_count IS '分块 token 数'; -COMMENT ON COLUMN document_chunks.page_number IS '来源页码'; -COMMENT ON COLUMN document_chunks.embedding_model IS '向量模型名称'; -COMMENT ON COLUMN document_chunks.embedding IS '分块向量数据'; -COMMENT ON COLUMN document_chunks.keywords IS '分块关键词'; -COMMENT ON COLUMN document_chunks.metadata IS '分块扩展元数据'; -COMMENT ON COLUMN document_chunks.created_at IS '创建时间'; - --- 文档处理任务表记录解析、分块、向量化和重建索引状态 -CREATE TABLE IF NOT EXISTS document_processing_jobs -( - id UUID PRIMARY KEY DEFAULT gen_random_uuid(), -- 任务 ID - user_id UUID NOT NULL, -- 所属用户 ID - document_id UUID NOT NULL, -- 所属文档 ID - job_type VARCHAR(32) NOT NULL, -- 任务类型,parse 解析,chunk 分块,embed 向量化,reindex 重建索引 - status INT NOT NULL DEFAULT 1, -- 任务状态,1 待处理,2 运行中,3 成功,4 失败 - error_message TEXT NOT NULL DEFAULT '', -- 任务失败原因 - started_at TIMESTAMPTZ, -- 开始时间 - finished_at TIMESTAMPTZ, -- 完成时间 - created_at TIMESTAMPTZ NOT NULL DEFAULT now(), -- 创建时间 - updated_at TIMESTAMPTZ NOT NULL DEFAULT now() -- 更新时间 - -); - -COMMENT ON TABLE document_processing_jobs IS '文档处理任务表,记录解析、分块、向量化和重建索引状态'; -COMMENT ON COLUMN document_processing_jobs.id IS '任务 ID'; -COMMENT ON COLUMN document_processing_jobs.user_id IS '所属用户 ID'; -COMMENT ON COLUMN document_processing_jobs.document_id IS '所属文档 ID'; -COMMENT ON COLUMN document_processing_jobs.job_type IS '任务类型,parse 解析,chunk 分块,embed 向量化,reindex 重建索引'; -COMMENT ON COLUMN document_processing_jobs.status IS '任务状态,1 待处理,2 运行中,3 成功,4 失败'; -COMMENT ON COLUMN document_processing_jobs.error_message IS '任务失败原因'; -COMMENT ON COLUMN document_processing_jobs.started_at IS '开始时间'; -COMMENT ON COLUMN document_processing_jobs.finished_at IS '完成时间'; -COMMENT ON COLUMN document_processing_jobs.created_at IS '创建时间'; -COMMENT ON COLUMN document_processing_jobs.updated_at IS '更新时间'; - --- 同步源表记录钉钉等外部平台同步配置 -CREATE TABLE IF NOT EXISTS sync_sources -( - id UUID PRIMARY KEY DEFAULT gen_random_uuid(), -- 同步源 ID - user_id UUID NOT NULL, -- 所属用户 ID - knowledge_base_id UUID NOT NULL, -- 绑定知识库 ID - name VARCHAR(128) NOT NULL, -- 同步源名称 - platform VARCHAR(32) NOT NULL, -- 同步平台,当前支持 dingtalk - source_config JSONB NOT NULL DEFAULT '{}'::jsonb, -- 非敏感同步配置 - status INT NOT NULL DEFAULT 1, -- 同步源状态,1 正常,2 禁用,3 已删除 - last_sync_at TIMESTAMPTZ, -- 最近同步时间 - last_error_message TEXT NOT NULL DEFAULT '', -- 最近同步失败原因 - created_at TIMESTAMPTZ NOT NULL DEFAULT now(), -- 创建时间 - updated_at TIMESTAMPTZ NOT NULL DEFAULT now(), -- 更新时间 - deleted_at TIMESTAMPTZ -- 删除时间 -); - -COMMENT ON TABLE sync_sources IS '同步源配置表,记录钉钉等外部平台同步入口'; -COMMENT ON COLUMN sync_sources.id IS '同步源 ID'; -COMMENT ON COLUMN sync_sources.user_id IS '所属用户 ID'; -COMMENT ON COLUMN sync_sources.knowledge_base_id IS '绑定知识库 ID'; -COMMENT ON COLUMN sync_sources.name IS '同步源名称'; -COMMENT ON COLUMN sync_sources.platform IS '同步平台,当前支持 dingtalk'; -COMMENT ON COLUMN sync_sources.source_config IS '非敏感同步配置'; -COMMENT ON COLUMN sync_sources.status IS '同步源状态,1 正常,2 禁用,3 已删除'; -COMMENT ON COLUMN sync_sources.last_sync_at IS '最近同步时间'; -COMMENT ON COLUMN sync_sources.last_error_message IS '最近同步失败原因'; -COMMENT ON COLUMN sync_sources.created_at IS '创建时间'; -COMMENT ON COLUMN sync_sources.updated_at IS '更新时间'; -COMMENT ON COLUMN sync_sources.deleted_at IS '删除时间'; - --- 同步任务表记录每次手动触发同步的执行状态 -CREATE TABLE IF NOT EXISTS sync_jobs -( - id UUID PRIMARY KEY DEFAULT gen_random_uuid(), -- 同步任务 ID - user_id UUID NOT NULL, -- 所属用户 ID - sync_source_id UUID NOT NULL, -- 同步源 ID - knowledge_base_id UUID NOT NULL, -- 绑定知识库 ID - job_type VARCHAR(32) NOT NULL, -- 任务类型,manual 手动同步 - status INT NOT NULL DEFAULT 1, -- 任务状态,1 待同步,2 同步中,3 成功,4 失败 - total_count INT NOT NULL DEFAULT 0, -- 同步总数 - success_count INT NOT NULL DEFAULT 0, -- 同步成功数 - failed_count INT NOT NULL DEFAULT 0, -- 同步失败数 - error_message TEXT NOT NULL DEFAULT '', -- 任务失败原因 - started_at TIMESTAMPTZ, -- 开始时间 - finished_at TIMESTAMPTZ, -- 完成时间 - created_at TIMESTAMPTZ NOT NULL DEFAULT now(), -- 创建时间 - updated_at TIMESTAMPTZ NOT NULL DEFAULT now() -- 更新时间 -); - -COMMENT ON TABLE sync_jobs IS '同步任务表,记录外部平台同步执行状态'; -COMMENT ON COLUMN sync_jobs.id IS '同步任务 ID'; -COMMENT ON COLUMN sync_jobs.user_id IS '所属用户 ID'; -COMMENT ON COLUMN sync_jobs.sync_source_id IS '同步源 ID'; -COMMENT ON COLUMN sync_jobs.knowledge_base_id IS '绑定知识库 ID'; -COMMENT ON COLUMN sync_jobs.job_type IS '任务类型,manual 手动同步'; -COMMENT ON COLUMN sync_jobs.status IS '任务状态,1 待同步,2 同步中,3 成功,4 失败'; -COMMENT ON COLUMN sync_jobs.total_count IS '同步总数'; -COMMENT ON COLUMN sync_jobs.success_count IS '同步成功数'; -COMMENT ON COLUMN sync_jobs.failed_count IS '同步失败数'; -COMMENT ON COLUMN sync_jobs.error_message IS '任务失败原因'; -COMMENT ON COLUMN sync_jobs.started_at IS '开始时间'; -COMMENT ON COLUMN sync_jobs.finished_at IS '完成时间'; -COMMENT ON COLUMN sync_jobs.created_at IS '创建时间'; -COMMENT ON COLUMN sync_jobs.updated_at IS '更新时间'; - --- 同步目录项表记录外部知识库中的目录和文件元数据 -CREATE TABLE IF NOT EXISTS sync_items -( - id UUID PRIMARY KEY DEFAULT gen_random_uuid(), -- 同步目录项 ID - user_id UUID NOT NULL, -- 所属用户 ID - sync_source_id UUID NOT NULL, -- 同步源 ID - knowledge_base_id UUID NOT NULL, -- 绑定知识库 ID - external_id VARCHAR(255) NOT NULL, -- 外部节点 ID - parent_external_id VARCHAR(255) NOT NULL DEFAULT '', -- 外部父节点 ID - name VARCHAR(255) NOT NULL, -- 节点名称 - item_type VARCHAR(32) NOT NULL, -- 节点类型,FILE 或 FOLDER - category VARCHAR(64) NOT NULL DEFAULT '', -- 钉钉节点分类 - extension VARCHAR(32) NOT NULL DEFAULT '', -- 文件扩展名 - external_url TEXT NOT NULL DEFAULT '', -- 外部原文链接 - file_size BIGINT NOT NULL DEFAULT 0, -- 文件大小 - has_children BOOLEAN NOT NULL DEFAULT FALSE, -- 是否有子节点 - source_updated_at TIMESTAMPTZ, -- 外部更新时间 - local_document_id UUID, -- 已导入本地文档 ID - import_status INT NOT NULL DEFAULT 1, -- 导入状态,1 未导入,2 导入中,3 已导入,4 导入失败 - error_message TEXT NOT NULL DEFAULT '', -- 导入失败原因 - created_at TIMESTAMPTZ NOT NULL DEFAULT now(), -- 创建时间 - updated_at TIMESTAMPTZ NOT NULL DEFAULT now(), -- 更新时间 - - CONSTRAINT sync_items_user_source_external_unique UNIQUE (user_id, sync_source_id, external_id) -); - -COMMENT ON TABLE sync_items IS '同步目录项表,记录外部知识库中的目录和文件元数据'; -COMMENT ON COLUMN sync_items.id IS '同步目录项 ID'; -COMMENT ON COLUMN sync_items.user_id IS '所属用户 ID'; -COMMENT ON COLUMN sync_items.sync_source_id IS '同步源 ID'; -COMMENT ON COLUMN sync_items.knowledge_base_id IS '绑定知识库 ID'; -COMMENT ON COLUMN sync_items.external_id IS '外部节点 ID'; -COMMENT ON COLUMN sync_items.parent_external_id IS '外部父节点 ID'; -COMMENT ON COLUMN sync_items.name IS '节点名称'; -COMMENT ON COLUMN sync_items.item_type IS '节点类型,FILE 或 FOLDER'; -COMMENT ON COLUMN sync_items.category IS '钉钉节点分类'; -COMMENT ON COLUMN sync_items.extension IS '文件扩展名'; -COMMENT ON COLUMN sync_items.external_url IS '外部原文链接'; -COMMENT ON COLUMN sync_items.file_size IS '文件大小'; -COMMENT ON COLUMN sync_items.has_children IS '是否有子节点'; -COMMENT ON COLUMN sync_items.source_updated_at IS '外部更新时间'; -COMMENT ON COLUMN sync_items.local_document_id IS '已导入本地文档 ID'; -COMMENT ON COLUMN sync_items.import_status IS '导入状态,1 未导入,2 导入中,3 已导入,4 导入失败'; -COMMENT ON COLUMN sync_items.error_message IS '导入失败原因'; -COMMENT ON COLUMN sync_items.created_at IS '创建时间'; -COMMENT ON COLUMN sync_items.updated_at IS '更新时间'; - --- 钉钉用户绑定表保存系统用户与钉钉身份的对应关系 -CREATE TABLE IF NOT EXISTS dingtalk_user_bindings -( - id UUID PRIMARY KEY DEFAULT gen_random_uuid(), -- 绑定 ID - user_id UUID NOT NULL, -- 系统用户 ID - ding_open_id VARCHAR(128) NOT NULL DEFAULT '', -- 钉钉 openid - ding_union_id VARCHAR(128) NOT NULL, -- 钉钉 unionId - corp_id VARCHAR(128) NOT NULL DEFAULT '', -- 钉钉企业 ID - nickname VARCHAR(128) NOT NULL DEFAULT '', -- 钉钉用户昵称 - avatar TEXT NOT NULL DEFAULT '', -- 钉钉用户头像 - created_at TIMESTAMPTZ NOT NULL DEFAULT now(), -- 创建时间 - updated_at TIMESTAMPTZ NOT NULL DEFAULT now(), -- 更新时间 - - CONSTRAINT dingtalk_user_bindings_user_unique UNIQUE (user_id), - CONSTRAINT dingtalk_user_bindings_union_unique UNIQUE (ding_union_id) -); - -COMMENT ON TABLE dingtalk_user_bindings IS '钉钉用户绑定表,用于保存系统用户与钉钉身份的对应关系'; -COMMENT ON COLUMN dingtalk_user_bindings.id IS '绑定 ID'; -COMMENT ON COLUMN dingtalk_user_bindings.user_id IS '系统用户 ID'; -COMMENT ON COLUMN dingtalk_user_bindings.ding_open_id IS '钉钉 openid'; -COMMENT ON COLUMN dingtalk_user_bindings.ding_union_id IS '钉钉 unionId'; -COMMENT ON COLUMN dingtalk_user_bindings.corp_id IS '钉钉企业 ID'; -COMMENT ON COLUMN dingtalk_user_bindings.nickname IS '钉钉用户昵称'; -COMMENT ON COLUMN dingtalk_user_bindings.avatar IS '钉钉用户头像'; -COMMENT ON COLUMN dingtalk_user_bindings.created_at IS '创建时间'; -COMMENT ON COLUMN dingtalk_user_bindings.updated_at IS '更新时间'; - --- 存储配额表记录用户存储上限和已用容量 -CREATE TABLE IF NOT EXISTS storage_quotas -( - id UUID PRIMARY KEY DEFAULT gen_random_uuid(), -- 配额记录 ID - user_id UUID NOT NULL, -- 所属用户 ID - max_storage_bytes BIGINT NOT NULL DEFAULT 10737418240, -- 最大可用存储字节数 - used_storage_bytes BIGINT NOT NULL DEFAULT 0, -- 已用存储字节数 - created_at TIMESTAMPTZ NOT NULL DEFAULT now(), -- 创建时间 - updated_at TIMESTAMPTZ NOT NULL DEFAULT now(), -- 更新时间 - - CONSTRAINT storage_quotas_user_unique UNIQUE (user_id) -); - -COMMENT ON TABLE storage_quotas IS '用户存储配额表,用于限制单用户知识库容量'; -COMMENT ON COLUMN storage_quotas.id IS '配额记录 ID'; -COMMENT ON COLUMN storage_quotas.user_id IS '所属用户 ID'; -COMMENT ON COLUMN storage_quotas.max_storage_bytes IS '最大可用存储字节数'; -COMMENT ON COLUMN storage_quotas.used_storage_bytes IS '已用存储字节数'; -COMMENT ON COLUMN storage_quotas.created_at IS '创建时间'; -COMMENT ON COLUMN storage_quotas.updated_at IS '更新时间'; - --- 常用查询索引 -CREATE INDEX IF NOT EXISTS idx_knowledge_bases_user_id - ON knowledge_bases(user_id); - -CREATE UNIQUE INDEX IF NOT EXISTS knowledge_bases_user_name_normal_unique - ON knowledge_bases(user_id, name) - WHERE status = 1; - -CREATE INDEX IF NOT EXISTS idx_documents_user_kb - ON documents(user_id, knowledge_base_id); - -CREATE INDEX IF NOT EXISTS idx_documents_user_external - ON documents(user_id, source_type, external_id) - WHERE external_id <> ''; - -CREATE INDEX IF NOT EXISTS idx_document_versions_document_id - ON document_versions(document_id); - -CREATE INDEX IF NOT EXISTS idx_document_chunks_user_kb - ON document_chunks(user_id, knowledge_base_id); - -CREATE INDEX IF NOT EXISTS idx_document_chunks_document_id - ON document_chunks(document_id); - -CREATE INDEX IF NOT EXISTS idx_document_processing_jobs_document_id - ON document_processing_jobs(document_id); - -CREATE INDEX IF NOT EXISTS idx_sync_sources_user_kb - ON sync_sources(user_id, knowledge_base_id); - -CREATE INDEX IF NOT EXISTS idx_sync_jobs_source_id - ON sync_jobs(sync_source_id); - -CREATE INDEX IF NOT EXISTS idx_sync_jobs_user_kb - ON sync_jobs(user_id, knowledge_base_id); - -CREATE INDEX IF NOT EXISTS idx_sync_items_source_parent - ON sync_items(sync_source_id, parent_external_id); - -CREATE INDEX IF NOT EXISTS idx_sync_items_user_kb - ON sync_items(user_id, knowledge_base_id); - -CREATE INDEX IF NOT EXISTS idx_dingtalk_user_bindings_user_id - ON dingtalk_user_bindings(user_id); - -CREATE INDEX IF NOT EXISTS idx_dingtalk_user_bindings_corp_id - ON dingtalk_user_bindings(corp_id); - -CREATE INDEX IF NOT EXISTS idx_document_chunks_embedding - ON document_chunks - USING ivfflat (embedding vector_cosine_ops) - WITH (lists = 100) - WHERE embedding IS NOT NULL; From 5b745f035132385c73c6fe0761e8ab74bf80b873 Mon Sep 17 00:00:00 2001 From: st <2663600842@qq.com> Date: Mon, 10 Aug 2026 17:09:27 +0800 Subject: [PATCH 10/10] =?UTF-8?q?revert:=20=E6=81=A2=E5=A4=8D=20client=5Ft?= =?UTF-8?q?est.go=20=E5=92=8C=20init=5Fknowledge=5Fschema.sql?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 恢复 internal/integration/dingtalk/client_test.go - 恢复 scripts/init_knowledge_schema.sql --- internal/integration/dingtalk/client_test.go | 235 ++++++++++ scripts/init_knowledge_schema.sql | 424 +++++++++++++++++++ 2 files changed, 659 insertions(+) create mode 100644 internal/integration/dingtalk/client_test.go create mode 100644 scripts/init_knowledge_schema.sql diff --git a/internal/integration/dingtalk/client_test.go b/internal/integration/dingtalk/client_test.go new file mode 100644 index 0000000..6543614 --- /dev/null +++ b/internal/integration/dingtalk/client_test.go @@ -0,0 +1,235 @@ +package dingtalk + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "solvify-agent/pkg/config" +) + +// TestNodeUnmarshalModifiedTimeFormats 验证节点更新时间兼容时间戳和分钟精度时间 +func TestNodeUnmarshalModifiedTimeFormats(t *testing.T) { + tests := []struct { + name string + value string + expected int64 + }{ + {name: "毫秒时间戳", value: `"1719999999000"`, expected: 1719999999000}, + {name: "分钟精度时间", value: `"2026-07-01T19:22Z"`, expected: time.Date(2026, 7, 1, 19, 22, 0, 0, time.UTC).Unix()}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var node Node + if err := json.Unmarshal([]byte(`{"modifiedTime":`+tt.value+`}`), &node); err != nil { + t.Fatalf("解析节点更新时间失败: %v", err) + } + if node.ModifiedAt != tt.expected { + t.Fatalf("节点更新时间不符合预期: got=%d want=%d", node.ModifiedAt, tt.expected) + } + }) + } +} + +// TestClientListNodesUsesHeaderToken 验证节点列表使用 Header 鉴权和分页参数 +func TestClientListNodesUsesHeaderToken(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/v1.0/oauth2/accessToken": + if r.Method != http.MethodPost { + t.Fatalf("accessToken 请求方法错误: %s", r.Method) + } + if !strings.Contains(r.Header.Get("Content-Type"), "application/json") { + t.Fatalf("accessToken 请求体类型错误") + } + _, _ = w.Write([]byte(`{"accessToken":"token-1","expireIn":7200}`)) + case "/v2.0/wiki/nodes": + if r.Header.Get("x-acs-dingtalk-access-token") != "token-1" { + t.Fatalf("未使用钉钉 Header 鉴权") + } + if r.URL.Query().Get("parentNodeId") != "root-1" || r.URL.Query().Get("nextToken") != "next-1" { + t.Fatalf("节点列表分页参数错误: %s", r.URL.RawQuery) + } + _, _ = w.Write([]byte(`{"nodes":[{"nodeId":"node-1","workspaceId":"ws-1","name":"a.md","size":"12","type":"FILE","modifiedTime":"1719999999000"}],"nextToken":"next-2"}`)) + default: + t.Fatalf("未预期的请求路径: %s", r.URL.Path) + } + })) + defer server.Close() + + client := NewClient(config.DingTalkConfig{AppKey: "app-key", AppSecret: "app-secret"}) + client.httpClient = server.Client() + client.accessTokenURL = server.URL + "/v1.0/oauth2/accessToken" + client.apiBaseURL = server.URL + + nodes, nextToken, err := client.ListNodes(context.Background(), "union-1", "root-1", "next-1", 50) + if err != nil { + t.Fatalf("获取节点列表失败: %v", err) + } + if len(nodes) != 1 || nodes[0].NodeID != "node-1" || nodes[0].Size != 12 || nodes[0].ModifiedAt != 1719999999000 || nextToken != "next-2" { + t.Fatalf("节点列表响应解析错误: nodes=%v next=%s", nodes, nextToken) + } +} + +// TestClientQueryDentryIDEscapesPath 验证 dentryUuid 路径参数会转义 +func TestClientQueryDentryIDEscapesPath(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/v1.0/oauth2/accessToken": + _, _ = w.Write([]byte(`{"accessToken":"token-1","expireIn":7200}`)) + case "/v2.0/doc/dentries/abc/def/queryDentryId": + if !strings.Contains(r.URL.RawPath, "abc%2Fdef") && !strings.Contains(r.RequestURI, "abc%2Fdef") { + t.Fatalf("dentryUuid 未正确转义: %s", r.RequestURI) + } + _, _ = w.Write([]byte(`{"dentryUuid":"abc/def","dentryId":"d-1","spaceId":"s-1"}`)) + default: + t.Fatalf("未预期的请求路径: %s", r.URL.Path) + } + })) + defer server.Close() + + client := NewClient(config.DingTalkConfig{AppKey: "app-key", AppSecret: "app-secret"}) + client.httpClient = server.Client() + client.accessTokenURL = server.URL + "/v1.0/oauth2/accessToken" + client.apiBaseURL = server.URL + + output, err := client.QueryDentryID(context.Background(), "union-1", "abc/def") + if err != nil { + t.Fatalf("查询 dentryId 失败: %v", err) + } + if output.SpaceID != "s-1" || output.DentryID != "d-1" { + t.Fatalf("dentryId 响应解析错误: %+v", output) + } +} + +// TestClientDownloadFileUsesReturnedHeaders 验证下载文件使用钉钉返回的签名 Header +func TestClientDownloadFileUsesReturnedHeaders(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/v1.0/oauth2/accessToken": + _, _ = w.Write([]byte(`{"accessToken":"token-1","expireIn":7200}`)) + case "/v1.0/storage/spaces/s-1/dentries/d-1/downloadInfos/query": + if r.Header.Get("x-acs-dingtalk-access-token") != "token-1" { + t.Fatalf("下载信息未使用钉钉 Header 鉴权") + } + _, _ = w.Write([]byte(`{"protocol":"HEADER_SIGNATURE","headerSignatureInfo":{"resourceUrls":["` + serverURL(r) + `/download"],"headers":{"X-Sign":"ok"}}}`)) + case "/download": + if r.Header.Get("X-Sign") != "ok" { + t.Fatalf("文件下载未携带签名 Header") + } + _, _ = w.Write([]byte("hello")) + default: + t.Fatalf("未预期的请求路径: %s", r.URL.Path) + } + })) + defer server.Close() + + client := NewClient(config.DingTalkConfig{AppKey: "app-key", AppSecret: "app-secret"}) + client.httpClient = server.Client() + client.accessTokenURL = server.URL + "/v1.0/oauth2/accessToken" + client.apiBaseURL = server.URL + + data, hash, err := client.DownloadFile(context.Background(), "union-1", "s-1", "d-1") + if err != nil { + t.Fatalf("下载文件失败: %v", err) + } + if string(data) != "hello" || hash == "" { + t.Fatalf("下载内容解析错误: data=%q hash=%s", string(data), hash) + } +} + +// TestClientQueryDocumentBlocks 验证在线文档块元素查询参数和响应 +func TestClientQueryDocumentBlocks(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/v1.0/oauth2/accessToken": + _, _ = w.Write([]byte(`{"accessToken":"token-1","expireIn":7200}`)) + case "/v1.0/doc/suites/documents/doc-1/blocks": + if r.Method != http.MethodGet || r.URL.Query().Get("operatorId") != "union-1" { + t.Fatalf("块元素查询参数错误: %s %s", r.Method, r.URL.RawQuery) + } + _, _ = w.Write([]byte(`{"result":{"data":[{"blockType":"paragraph","paragraph":{"text":"正文"}}]},"success":true}`)) + default: + t.Fatalf("未预期的请求路径: %s", r.URL.Path) + } + })) + defer server.Close() + + client := NewClient(config.DingTalkConfig{AppKey: "app-key", AppSecret: "app-secret"}) + client.httpClient = server.Client() + client.accessTokenURL = server.URL + "/v1.0/oauth2/accessToken" + client.apiBaseURL = server.URL + + blocks, err := client.QueryDocumentBlocks(context.Background(), "union-1", "doc-1") + if err != nil { + t.Fatalf("查询在线文档块元素失败: %v", err) + } + if len(blocks) != 1 || blocks[0]["blockType"] != "paragraph" { + t.Fatalf("块元素响应解析错误: %+v", blocks) + } +} + +// serverURL 从测试请求还原服务地址 +func serverURL(r *http.Request) string { + return "http://" + r.Host +} + +// TestClientExchangeUserAccessToken 验证扫码授权码使用新版 OAuth JSON 请求兑换用户 token +func TestClientExchangeUserAccessToken(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/v1.0/oauth2/userAccessToken" { + t.Fatalf("未预期的请求路径: %s", r.URL.Path) + } + if r.Method != http.MethodPost { + t.Fatalf("用户 token 请求方法错误: %s", r.Method) + } + if !strings.Contains(r.Header.Get("Content-Type"), "application/json") { + t.Fatalf("用户 token 请求体类型错误") + } + _, _ = w.Write([]byte(`{"accessToken":"user-token","refreshToken":"refresh-token","expireIn":7200,"corpId":"corp-1"}`)) + })) + defer server.Close() + + client := NewClient(config.DingTalkConfig{AppKey: "app-key", AppSecret: "app-secret"}) + client.httpClient = server.Client() + client.apiBaseURL = server.URL + + output, err := client.ExchangeUserAccessToken(context.Background(), "auth-code") + if err != nil { + t.Fatalf("兑换用户 token 失败: %v", err) + } + if output.AccessToken != "user-token" || output.CorpID != "corp-1" { + t.Fatalf("用户 token 响应解析错误: %+v", output) + } +} + +// TestClientGetCurrentUserInfoUsesUserToken 验证用户信息接口使用个人 token Header +func TestClientGetCurrentUserInfoUsesUserToken(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/v1.0/contact/users/me" { + t.Fatalf("未预期的请求路径: %s", r.URL.Path) + } + if r.Header.Get("x-acs-dingtalk-access-token") != "user-token" { + t.Fatalf("用户信息接口未使用个人 token Header") + } + _, _ = w.Write([]byte(`{"nick":"张三","avatarUrl":"https://example.com/a.png","openId":"open-1","unionId":"union-1","email":"a@example.com"}`)) + })) + defer server.Close() + + client := NewClient(config.DingTalkConfig{AppKey: "app-key", AppSecret: "app-secret"}) + client.httpClient = server.Client() + client.apiBaseURL = server.URL + + output, err := client.GetCurrentUserInfo(context.Background(), "user-token") + if err != nil { + t.Fatalf("获取用户信息失败: %v", err) + } + if output.UnionID != "union-1" || output.OpenID != "open-1" { + t.Fatalf("用户信息响应解析错误: %+v", output) + } +} diff --git a/scripts/init_knowledge_schema.sql b/scripts/init_knowledge_schema.sql new file mode 100644 index 0000000..c730294 --- /dev/null +++ b/scripts/init_knowledge_schema.sql @@ -0,0 +1,424 @@ +CREATE +EXTENSION IF NOT EXISTS pgcrypto; +CREATE +EXTENSION IF NOT EXISTS vector; + +-- 用户表作为当前阶段的数据隔离边界 +CREATE TABLE IF NOT EXISTS users +( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), -- 用户 ID + username VARCHAR(64) NOT NULL, -- 用户名 + email VARCHAR(128) NOT NULL, -- 邮箱 + password TEXT NOT NULL DEFAULT '', -- 密码哈希 + status INT NOT NULL DEFAULT 1, -- 用户状态,1 正常,2 禁用,3 注销,4 待验证 + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), -- 创建时间 + updated_at TIMESTAMPTZ NOT NULL DEFAULT now(), -- 更新时间 + + CONSTRAINT users_username_unique UNIQUE (username), + CONSTRAINT users_email_unique UNIQUE (email) +); + +COMMENT ON TABLE users IS '用户基础表,用于隔离每个用户自己的知识库和文档'; +COMMENT ON COLUMN users.id IS '用户 ID'; +COMMENT ON COLUMN users.username IS '用户名'; +COMMENT ON COLUMN users.email IS '邮箱'; +COMMENT ON COLUMN users.password IS '密码哈希'; +COMMENT ON COLUMN users.status IS '用户状态,1 正常,2 禁用,3 注销,4 待验证'; +COMMENT ON COLUMN users.created_at IS '创建时间'; +COMMENT ON COLUMN users.updated_at IS '更新时间'; + +-- 知识库表保存用户自建、同步和联网搜索知识库 +CREATE TABLE IF NOT EXISTS knowledge_bases +( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), -- 知识库 ID + user_id UUID NOT NULL, -- 所属用户 ID + name VARCHAR(128) NOT NULL, -- 知识库名称 + category VARCHAR(64) NOT NULL DEFAULT '', -- 知识库分类 + description TEXT DEFAULT '', -- 知识库描述 + source_type VARCHAR(32) NOT NULL DEFAULT 'local', -- 知识库来源类型,local 自建,sync 同步,web_search 联网搜索 + source_platform VARCHAR(32) NOT NULL DEFAULT '', -- 同步来源平台 + document_count INT NOT NULL DEFAULT 0, -- 文档数量 + storage_bytes BIGINT NOT NULL DEFAULT 0, -- 已占用存储字节数 + status INT NOT NULL DEFAULT 1, -- 知识库状态,1 正常,2 已删除 + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), -- 创建时间 + updated_at TIMESTAMPTZ NOT NULL DEFAULT now(), -- 更新时间 + deleted_at TIMESTAMPTZ, -- 删除时间 + delete_expired_at TIMESTAMPTZ -- 删除保留到期时间 +); + +COMMENT ON TABLE knowledge_bases IS '知识库主表,所有知识库按用户隔离'; +COMMENT ON COLUMN knowledge_bases.id IS '知识库 ID'; +COMMENT ON COLUMN knowledge_bases.user_id IS '所属用户 ID'; +COMMENT ON COLUMN knowledge_bases.name IS '知识库名称'; +COMMENT ON COLUMN knowledge_bases.category IS '知识库分类'; +COMMENT ON COLUMN knowledge_bases.description IS '知识库描述'; +COMMENT ON COLUMN knowledge_bases.source_type IS '知识库来源类型,local 自建,sync 同步,web_search 联网搜索'; +COMMENT ON COLUMN knowledge_bases.source_platform IS '同步来源平台'; +COMMENT ON COLUMN knowledge_bases.document_count IS '文档数量'; +COMMENT ON COLUMN knowledge_bases.storage_bytes IS '已占用存储字节数'; +COMMENT ON COLUMN knowledge_bases.status IS '知识库状态,1 正常,2 已删除'; +COMMENT ON COLUMN knowledge_bases.created_at IS '创建时间'; +COMMENT ON COLUMN knowledge_bases.updated_at IS '更新时间'; +COMMENT ON COLUMN knowledge_bases.deleted_at IS '删除时间'; +COMMENT ON COLUMN knowledge_bases.delete_expired_at IS '删除保留到期时间'; + +-- 文档表记录上传文件和处理状态 +CREATE TABLE IF NOT EXISTS documents +( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), -- 文档 ID + user_id UUID NOT NULL, -- 所属用户 ID + knowledge_base_id UUID NOT NULL, -- 所属知识库 ID + title VARCHAR(255) NOT NULL, -- 文档标题 + file_name VARCHAR(255) NOT NULL, -- 原始文件名 + file_type VARCHAR(32) NOT NULL DEFAULT '', -- 文件类型 + file_size BIGINT NOT NULL DEFAULT 0, -- 文件大小字节数 + storage_path TEXT NOT NULL DEFAULT '', -- 文件存储路径 + file_hash VARCHAR(128) NOT NULL DEFAULT '', -- 原始文件内容指纹 + source_type VARCHAR(32) NOT NULL DEFAULT 'upload', -- 文档来源类型,upload 上传,edit 编辑,sync 同步,web_search 联网搜索 + external_id VARCHAR(255) NOT NULL DEFAULT '', -- 外部平台文档 ID + external_url TEXT NOT NULL DEFAULT '', -- 外部平台文档链接 + source_updated_at TIMESTAMPTZ, -- 外部平台更新时间 + status INT NOT NULL DEFAULT 1, -- 文档状态,1 已上传,2 处理中,3 已就绪,4 处理失败,5 已删除 + error_message TEXT NOT NULL DEFAULT '', -- 处理失败原因 + ready_at TIMESTAMPTZ, -- 文档就绪时间 + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), -- 创建时间 + updated_at TIMESTAMPTZ NOT NULL DEFAULT now(), -- 更新时间 + deleted_at TIMESTAMPTZ, -- 删除时间 + delete_expired_at TIMESTAMPTZ -- 删除保留到期时间 + +); + +COMMENT ON TABLE documents IS '文档主表,记录知识库下的文件和处理状态'; +COMMENT ON COLUMN documents.id IS '文档 ID'; +COMMENT ON COLUMN documents.user_id IS '所属用户 ID'; +COMMENT ON COLUMN documents.knowledge_base_id IS '所属知识库 ID'; +COMMENT ON COLUMN documents.title IS '文档标题'; +COMMENT ON COLUMN documents.file_name IS '原始文件名'; +COMMENT ON COLUMN documents.file_type IS '文件类型'; +COMMENT ON COLUMN documents.file_size IS '文件大小字节数'; +COMMENT ON COLUMN documents.storage_path IS '文件存储路径'; +COMMENT ON COLUMN documents.file_hash IS '原始文件内容指纹'; +COMMENT ON COLUMN documents.source_type IS '文档来源类型,upload 上传,edit 编辑,sync 同步,web_search 联网搜索'; +COMMENT ON COLUMN documents.external_id IS '外部平台文档 ID'; +COMMENT ON COLUMN documents.external_url IS '外部平台文档链接'; +COMMENT ON COLUMN documents.source_updated_at IS '外部平台更新时间'; +COMMENT ON COLUMN documents.status IS '文档状态,1 已上传,2 处理中,3 已就绪,4 处理失败,5 已删除'; +COMMENT ON COLUMN documents.error_message IS '处理失败原因'; +COMMENT ON COLUMN documents.ready_at IS '文档就绪时间'; +COMMENT ON COLUMN documents.deleted_at IS '删除时间'; +COMMENT ON COLUMN documents.delete_expired_at IS '删除保留到期时间'; +COMMENT ON COLUMN documents.created_at IS '创建时间'; +COMMENT ON COLUMN documents.updated_at IS '更新时间'; + +ALTER TABLE documents + ADD COLUMN IF NOT EXISTS external_id VARCHAR(255) NOT NULL DEFAULT '', + ADD COLUMN IF NOT EXISTS external_url TEXT NOT NULL DEFAULT '', + ADD COLUMN IF NOT EXISTS source_updated_at TIMESTAMPTZ; + +-- 文档版本表支持在线编辑和重新向量化 +CREATE TABLE IF NOT EXISTS document_versions +( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), -- 版本 ID + user_id UUID NOT NULL, -- 所属用户 ID + document_id UUID NOT NULL, -- 所属文档 ID + version_no INT NOT NULL DEFAULT 1, -- 版本号 + content TEXT NOT NULL DEFAULT '', -- 版本正文内容 + content_hash VARCHAR(128) NOT NULL DEFAULT '', -- 版本内容哈希 + change_summary TEXT NOT NULL DEFAULT '', -- 变更摘要 + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), -- 创建时间 + + CONSTRAINT document_versions_document_version_unique UNIQUE (document_id, version_no) +); + +COMMENT ON TABLE document_versions IS '文档版本表,用于记录原始解析内容和在线编辑历史'; +COMMENT ON COLUMN document_versions.id IS '版本 ID'; +COMMENT ON COLUMN document_versions.user_id IS '所属用户 ID'; +COMMENT ON COLUMN document_versions.document_id IS '所属文档 ID'; +COMMENT ON COLUMN document_versions.version_no IS '版本号'; +COMMENT ON COLUMN document_versions.content IS '版本正文内容'; +COMMENT ON COLUMN document_versions.content_hash IS '版本内容哈希'; +COMMENT ON COLUMN document_versions.change_summary IS '变更摘要'; +COMMENT ON COLUMN document_versions.created_at IS '创建时间'; + +-- 文档分块表保存文本分块和向量数据 +CREATE TABLE IF NOT EXISTS document_chunks +( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), -- 分块 ID + user_id UUID NOT NULL, -- 所属用户 ID + knowledge_base_id UUID NOT NULL, -- 所属知识库 ID + document_id UUID NOT NULL, -- 所属文档 ID + version_id UUID NOT NULL, -- 所属文档版本 ID + chunk_index INT NOT NULL, -- 分块序号 + section_title TEXT NOT NULL DEFAULT '', -- 分块所属章节标题 + content TEXT NOT NULL, -- 分块文本内容 + token_count INT NOT NULL DEFAULT 0, -- 分块 token 数 + page_number INT, -- 来源页码 + embedding_model VARCHAR(128) NOT NULL DEFAULT '', -- 向量模型名称 + embedding vector(1024), -- 分块向量数据 + keywords TEXT[] NOT NULL DEFAULT '{}'::text[], -- 分块关键词 + metadata JSONB NOT NULL DEFAULT '{}'::jsonb, -- 分块扩展元数据 + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), -- 创建时间 + + CONSTRAINT document_chunks_version_index_unique UNIQUE (version_id, chunk_index) +); + +COMMENT ON TABLE document_chunks IS '文档分块表,保存可检索文本片段和 embedding 向量'; +COMMENT ON COLUMN document_chunks.id IS '分块 ID'; +COMMENT ON COLUMN document_chunks.user_id IS '所属用户 ID'; +COMMENT ON COLUMN document_chunks.knowledge_base_id IS '所属知识库 ID'; +COMMENT ON COLUMN document_chunks.document_id IS '所属文档 ID'; +COMMENT ON COLUMN document_chunks.version_id IS '所属文档版本 ID'; +COMMENT ON COLUMN document_chunks.chunk_index IS '分块序号'; +COMMENT ON COLUMN document_chunks.section_title IS '分块所属章节标题'; +COMMENT ON COLUMN document_chunks.content IS '分块文本内容'; +COMMENT ON COLUMN document_chunks.token_count IS '分块 token 数'; +COMMENT ON COLUMN document_chunks.page_number IS '来源页码'; +COMMENT ON COLUMN document_chunks.embedding_model IS '向量模型名称'; +COMMENT ON COLUMN document_chunks.embedding IS '分块向量数据'; +COMMENT ON COLUMN document_chunks.keywords IS '分块关键词'; +COMMENT ON COLUMN document_chunks.metadata IS '分块扩展元数据'; +COMMENT ON COLUMN document_chunks.created_at IS '创建时间'; + +-- 文档处理任务表记录解析、分块、向量化和重建索引状态 +CREATE TABLE IF NOT EXISTS document_processing_jobs +( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), -- 任务 ID + user_id UUID NOT NULL, -- 所属用户 ID + document_id UUID NOT NULL, -- 所属文档 ID + job_type VARCHAR(32) NOT NULL, -- 任务类型,parse 解析,chunk 分块,embed 向量化,reindex 重建索引 + status INT NOT NULL DEFAULT 1, -- 任务状态,1 待处理,2 运行中,3 成功,4 失败 + error_message TEXT NOT NULL DEFAULT '', -- 任务失败原因 + started_at TIMESTAMPTZ, -- 开始时间 + finished_at TIMESTAMPTZ, -- 完成时间 + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), -- 创建时间 + updated_at TIMESTAMPTZ NOT NULL DEFAULT now() -- 更新时间 + +); + +COMMENT ON TABLE document_processing_jobs IS '文档处理任务表,记录解析、分块、向量化和重建索引状态'; +COMMENT ON COLUMN document_processing_jobs.id IS '任务 ID'; +COMMENT ON COLUMN document_processing_jobs.user_id IS '所属用户 ID'; +COMMENT ON COLUMN document_processing_jobs.document_id IS '所属文档 ID'; +COMMENT ON COLUMN document_processing_jobs.job_type IS '任务类型,parse 解析,chunk 分块,embed 向量化,reindex 重建索引'; +COMMENT ON COLUMN document_processing_jobs.status IS '任务状态,1 待处理,2 运行中,3 成功,4 失败'; +COMMENT ON COLUMN document_processing_jobs.error_message IS '任务失败原因'; +COMMENT ON COLUMN document_processing_jobs.started_at IS '开始时间'; +COMMENT ON COLUMN document_processing_jobs.finished_at IS '完成时间'; +COMMENT ON COLUMN document_processing_jobs.created_at IS '创建时间'; +COMMENT ON COLUMN document_processing_jobs.updated_at IS '更新时间'; + +-- 同步源表记录钉钉等外部平台同步配置 +CREATE TABLE IF NOT EXISTS sync_sources +( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), -- 同步源 ID + user_id UUID NOT NULL, -- 所属用户 ID + knowledge_base_id UUID NOT NULL, -- 绑定知识库 ID + name VARCHAR(128) NOT NULL, -- 同步源名称 + platform VARCHAR(32) NOT NULL, -- 同步平台,当前支持 dingtalk + source_config JSONB NOT NULL DEFAULT '{}'::jsonb, -- 非敏感同步配置 + status INT NOT NULL DEFAULT 1, -- 同步源状态,1 正常,2 禁用,3 已删除 + last_sync_at TIMESTAMPTZ, -- 最近同步时间 + last_error_message TEXT NOT NULL DEFAULT '', -- 最近同步失败原因 + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), -- 创建时间 + updated_at TIMESTAMPTZ NOT NULL DEFAULT now(), -- 更新时间 + deleted_at TIMESTAMPTZ -- 删除时间 +); + +COMMENT ON TABLE sync_sources IS '同步源配置表,记录钉钉等外部平台同步入口'; +COMMENT ON COLUMN sync_sources.id IS '同步源 ID'; +COMMENT ON COLUMN sync_sources.user_id IS '所属用户 ID'; +COMMENT ON COLUMN sync_sources.knowledge_base_id IS '绑定知识库 ID'; +COMMENT ON COLUMN sync_sources.name IS '同步源名称'; +COMMENT ON COLUMN sync_sources.platform IS '同步平台,当前支持 dingtalk'; +COMMENT ON COLUMN sync_sources.source_config IS '非敏感同步配置'; +COMMENT ON COLUMN sync_sources.status IS '同步源状态,1 正常,2 禁用,3 已删除'; +COMMENT ON COLUMN sync_sources.last_sync_at IS '最近同步时间'; +COMMENT ON COLUMN sync_sources.last_error_message IS '最近同步失败原因'; +COMMENT ON COLUMN sync_sources.created_at IS '创建时间'; +COMMENT ON COLUMN sync_sources.updated_at IS '更新时间'; +COMMENT ON COLUMN sync_sources.deleted_at IS '删除时间'; + +-- 同步任务表记录每次手动触发同步的执行状态 +CREATE TABLE IF NOT EXISTS sync_jobs +( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), -- 同步任务 ID + user_id UUID NOT NULL, -- 所属用户 ID + sync_source_id UUID NOT NULL, -- 同步源 ID + knowledge_base_id UUID NOT NULL, -- 绑定知识库 ID + job_type VARCHAR(32) NOT NULL, -- 任务类型,manual 手动同步 + status INT NOT NULL DEFAULT 1, -- 任务状态,1 待同步,2 同步中,3 成功,4 失败 + total_count INT NOT NULL DEFAULT 0, -- 同步总数 + success_count INT NOT NULL DEFAULT 0, -- 同步成功数 + failed_count INT NOT NULL DEFAULT 0, -- 同步失败数 + error_message TEXT NOT NULL DEFAULT '', -- 任务失败原因 + started_at TIMESTAMPTZ, -- 开始时间 + finished_at TIMESTAMPTZ, -- 完成时间 + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), -- 创建时间 + updated_at TIMESTAMPTZ NOT NULL DEFAULT now() -- 更新时间 +); + +COMMENT ON TABLE sync_jobs IS '同步任务表,记录外部平台同步执行状态'; +COMMENT ON COLUMN sync_jobs.id IS '同步任务 ID'; +COMMENT ON COLUMN sync_jobs.user_id IS '所属用户 ID'; +COMMENT ON COLUMN sync_jobs.sync_source_id IS '同步源 ID'; +COMMENT ON COLUMN sync_jobs.knowledge_base_id IS '绑定知识库 ID'; +COMMENT ON COLUMN sync_jobs.job_type IS '任务类型,manual 手动同步'; +COMMENT ON COLUMN sync_jobs.status IS '任务状态,1 待同步,2 同步中,3 成功,4 失败'; +COMMENT ON COLUMN sync_jobs.total_count IS '同步总数'; +COMMENT ON COLUMN sync_jobs.success_count IS '同步成功数'; +COMMENT ON COLUMN sync_jobs.failed_count IS '同步失败数'; +COMMENT ON COLUMN sync_jobs.error_message IS '任务失败原因'; +COMMENT ON COLUMN sync_jobs.started_at IS '开始时间'; +COMMENT ON COLUMN sync_jobs.finished_at IS '完成时间'; +COMMENT ON COLUMN sync_jobs.created_at IS '创建时间'; +COMMENT ON COLUMN sync_jobs.updated_at IS '更新时间'; + +-- 同步目录项表记录外部知识库中的目录和文件元数据 +CREATE TABLE IF NOT EXISTS sync_items +( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), -- 同步目录项 ID + user_id UUID NOT NULL, -- 所属用户 ID + sync_source_id UUID NOT NULL, -- 同步源 ID + knowledge_base_id UUID NOT NULL, -- 绑定知识库 ID + external_id VARCHAR(255) NOT NULL, -- 外部节点 ID + parent_external_id VARCHAR(255) NOT NULL DEFAULT '', -- 外部父节点 ID + name VARCHAR(255) NOT NULL, -- 节点名称 + item_type VARCHAR(32) NOT NULL, -- 节点类型,FILE 或 FOLDER + category VARCHAR(64) NOT NULL DEFAULT '', -- 钉钉节点分类 + extension VARCHAR(32) NOT NULL DEFAULT '', -- 文件扩展名 + external_url TEXT NOT NULL DEFAULT '', -- 外部原文链接 + file_size BIGINT NOT NULL DEFAULT 0, -- 文件大小 + has_children BOOLEAN NOT NULL DEFAULT FALSE, -- 是否有子节点 + source_updated_at TIMESTAMPTZ, -- 外部更新时间 + local_document_id UUID, -- 已导入本地文档 ID + import_status INT NOT NULL DEFAULT 1, -- 导入状态,1 未导入,2 导入中,3 已导入,4 导入失败 + error_message TEXT NOT NULL DEFAULT '', -- 导入失败原因 + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), -- 创建时间 + updated_at TIMESTAMPTZ NOT NULL DEFAULT now(), -- 更新时间 + + CONSTRAINT sync_items_user_source_external_unique UNIQUE (user_id, sync_source_id, external_id) +); + +COMMENT ON TABLE sync_items IS '同步目录项表,记录外部知识库中的目录和文件元数据'; +COMMENT ON COLUMN sync_items.id IS '同步目录项 ID'; +COMMENT ON COLUMN sync_items.user_id IS '所属用户 ID'; +COMMENT ON COLUMN sync_items.sync_source_id IS '同步源 ID'; +COMMENT ON COLUMN sync_items.knowledge_base_id IS '绑定知识库 ID'; +COMMENT ON COLUMN sync_items.external_id IS '外部节点 ID'; +COMMENT ON COLUMN sync_items.parent_external_id IS '外部父节点 ID'; +COMMENT ON COLUMN sync_items.name IS '节点名称'; +COMMENT ON COLUMN sync_items.item_type IS '节点类型,FILE 或 FOLDER'; +COMMENT ON COLUMN sync_items.category IS '钉钉节点分类'; +COMMENT ON COLUMN sync_items.extension IS '文件扩展名'; +COMMENT ON COLUMN sync_items.external_url IS '外部原文链接'; +COMMENT ON COLUMN sync_items.file_size IS '文件大小'; +COMMENT ON COLUMN sync_items.has_children IS '是否有子节点'; +COMMENT ON COLUMN sync_items.source_updated_at IS '外部更新时间'; +COMMENT ON COLUMN sync_items.local_document_id IS '已导入本地文档 ID'; +COMMENT ON COLUMN sync_items.import_status IS '导入状态,1 未导入,2 导入中,3 已导入,4 导入失败'; +COMMENT ON COLUMN sync_items.error_message IS '导入失败原因'; +COMMENT ON COLUMN sync_items.created_at IS '创建时间'; +COMMENT ON COLUMN sync_items.updated_at IS '更新时间'; + +-- 钉钉用户绑定表保存系统用户与钉钉身份的对应关系 +CREATE TABLE IF NOT EXISTS dingtalk_user_bindings +( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), -- 绑定 ID + user_id UUID NOT NULL, -- 系统用户 ID + ding_open_id VARCHAR(128) NOT NULL DEFAULT '', -- 钉钉 openid + ding_union_id VARCHAR(128) NOT NULL, -- 钉钉 unionId + corp_id VARCHAR(128) NOT NULL DEFAULT '', -- 钉钉企业 ID + nickname VARCHAR(128) NOT NULL DEFAULT '', -- 钉钉用户昵称 + avatar TEXT NOT NULL DEFAULT '', -- 钉钉用户头像 + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), -- 创建时间 + updated_at TIMESTAMPTZ NOT NULL DEFAULT now(), -- 更新时间 + + CONSTRAINT dingtalk_user_bindings_user_unique UNIQUE (user_id), + CONSTRAINT dingtalk_user_bindings_union_unique UNIQUE (ding_union_id) +); + +COMMENT ON TABLE dingtalk_user_bindings IS '钉钉用户绑定表,用于保存系统用户与钉钉身份的对应关系'; +COMMENT ON COLUMN dingtalk_user_bindings.id IS '绑定 ID'; +COMMENT ON COLUMN dingtalk_user_bindings.user_id IS '系统用户 ID'; +COMMENT ON COLUMN dingtalk_user_bindings.ding_open_id IS '钉钉 openid'; +COMMENT ON COLUMN dingtalk_user_bindings.ding_union_id IS '钉钉 unionId'; +COMMENT ON COLUMN dingtalk_user_bindings.corp_id IS '钉钉企业 ID'; +COMMENT ON COLUMN dingtalk_user_bindings.nickname IS '钉钉用户昵称'; +COMMENT ON COLUMN dingtalk_user_bindings.avatar IS '钉钉用户头像'; +COMMENT ON COLUMN dingtalk_user_bindings.created_at IS '创建时间'; +COMMENT ON COLUMN dingtalk_user_bindings.updated_at IS '更新时间'; + +-- 存储配额表记录用户存储上限和已用容量 +CREATE TABLE IF NOT EXISTS storage_quotas +( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), -- 配额记录 ID + user_id UUID NOT NULL, -- 所属用户 ID + max_storage_bytes BIGINT NOT NULL DEFAULT 10737418240, -- 最大可用存储字节数 + used_storage_bytes BIGINT NOT NULL DEFAULT 0, -- 已用存储字节数 + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), -- 创建时间 + updated_at TIMESTAMPTZ NOT NULL DEFAULT now(), -- 更新时间 + + CONSTRAINT storage_quotas_user_unique UNIQUE (user_id) +); + +COMMENT ON TABLE storage_quotas IS '用户存储配额表,用于限制单用户知识库容量'; +COMMENT ON COLUMN storage_quotas.id IS '配额记录 ID'; +COMMENT ON COLUMN storage_quotas.user_id IS '所属用户 ID'; +COMMENT ON COLUMN storage_quotas.max_storage_bytes IS '最大可用存储字节数'; +COMMENT ON COLUMN storage_quotas.used_storage_bytes IS '已用存储字节数'; +COMMENT ON COLUMN storage_quotas.created_at IS '创建时间'; +COMMENT ON COLUMN storage_quotas.updated_at IS '更新时间'; + +-- 常用查询索引 +CREATE INDEX IF NOT EXISTS idx_knowledge_bases_user_id + ON knowledge_bases(user_id); + +CREATE UNIQUE INDEX IF NOT EXISTS knowledge_bases_user_name_normal_unique + ON knowledge_bases(user_id, name) + WHERE status = 1; + +CREATE INDEX IF NOT EXISTS idx_documents_user_kb + ON documents(user_id, knowledge_base_id); + +CREATE INDEX IF NOT EXISTS idx_documents_user_external + ON documents(user_id, source_type, external_id) + WHERE external_id <> ''; + +CREATE INDEX IF NOT EXISTS idx_document_versions_document_id + ON document_versions(document_id); + +CREATE INDEX IF NOT EXISTS idx_document_chunks_user_kb + ON document_chunks(user_id, knowledge_base_id); + +CREATE INDEX IF NOT EXISTS idx_document_chunks_document_id + ON document_chunks(document_id); + +CREATE INDEX IF NOT EXISTS idx_document_processing_jobs_document_id + ON document_processing_jobs(document_id); + +CREATE INDEX IF NOT EXISTS idx_sync_sources_user_kb + ON sync_sources(user_id, knowledge_base_id); + +CREATE INDEX IF NOT EXISTS idx_sync_jobs_source_id + ON sync_jobs(sync_source_id); + +CREATE INDEX IF NOT EXISTS idx_sync_jobs_user_kb + ON sync_jobs(user_id, knowledge_base_id); + +CREATE INDEX IF NOT EXISTS idx_sync_items_source_parent + ON sync_items(sync_source_id, parent_external_id); + +CREATE INDEX IF NOT EXISTS idx_sync_items_user_kb + ON sync_items(user_id, knowledge_base_id); + +CREATE INDEX IF NOT EXISTS idx_dingtalk_user_bindings_user_id + ON dingtalk_user_bindings(user_id); + +CREATE INDEX IF NOT EXISTS idx_dingtalk_user_bindings_corp_id + ON dingtalk_user_bindings(corp_id); + +CREATE INDEX IF NOT EXISTS idx_document_chunks_embedding + ON document_chunks + USING ivfflat (embedding vector_cosine_ops) + WITH (lists = 100) + WHERE embedding IS NOT NULL;
{{ pendingApproval.detail }}