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/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 ── 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 4f7329a..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,240 +123,113 @@ func (e *Engine) runAgent(ctx context.Context, req Request, chatModel model.Tool } var systemPromptFinal string if req.SystemPrompt != "" { - systemPromptFinal = baseSystemPrompt + "\n\n" + req.SystemPrompt + enhanced := strings.TrimLeft(req.SystemPrompt, "\n") + systemPromptFinal = baseSystemPrompt + "\n\n" + enhanced } else { - systemPromptFinal = buildEnhancedSystemPromptForAgent(baseSystemPrompt, req.Summary, req.Memories, req.UserCtx) + 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 - } + inputMessages := buildInputMessages(req.Query, req.History) - e.processStream(ctx, stream, ksToolForStream, eventCh) + maxStep := e.cfg.MaxIterations + if maxStep <= 0 { + maxStep = 5 } -} -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] - } - 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 + // ── 创建 adk.ChatModelAgent ── + toolsNodeConfig := compose.ToolsNodeConfig{ + Tools: allTools, - for { - msg, err := stream.Recv() - if err == io.EOF { - break - } - if err != nil { - if ctx.Err() != nil { - logger.Infof("Agent 流被用户中断,已收集 %d 字符", len(fullAnswer)) - break + // 兜底: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 } - logger.Errorf("Agent 流读取失败: %v", err) - eventCh <- Event{ - Type: EventError, - Title: "推理过程中断", - Detail: "深度推理过程中断,请重试", - Error: err.Error(), - Status: "error", - Retryable: true, - Done: true, + 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 } - return + return arguments, nil + }, + + 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 + }, + } + // 有危险工具时注入审批中间件 + if dangerousNames := e.dangerousToolNames(); len(dangerousNames) > 0 { + toolsNodeConfig.ToolCallMiddlewares = []compose.ToolMiddleware{ + {Invokable: buildDangerousToolMiddleware(dangerousNames)}, } - if msg == nil { - continue - } - - if msg.Role != schema.Assistant { - continue + logger.Infof("[Agent] 已注入危险工具审批中间件,工具列表=%v", dangerousNames) + } + + 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) } - - 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", - } - } - continue - } - - if msg.Content != "" { - // ── 最终答案轮次(没有下一步 ToolCalls,真正面向用户的正文)── - fullAnswer += msg.Content - eventCh <- Event{Type: EventAnswer, Content: msg.Content} - } - } - - // ── 兜底:极端情况(每一轮都有 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)) - } - } - sb.WriteString("\n如需进一步分析请补充问题细节,或切换到快速模式获取更直接的回答。") - fullAnswer = sb.String() - eventCh <- Event{Type: EventAnswer, Content: fullAnswer} - } - - 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, - } + eventCh <- Event{ + Type: EventError, + Title: "深度模式启动失败", + Detail: "请尝试切换到快速模式,或稍后重试", + Error: err.Error(), + Status: "error", + Retryable: true, + Done: true, } - 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, - }) + return } - if strings.TrimSpace(fullAnswer) != "" { - eventCh <- Event{Type: EventThinking, Title: "正在生成答案", Status: "success"} + // ── 创建 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, + }) - if len(sources) > 0 { - eventCh <- Event{Type: EventSources, Sources: sources} - } + // ── 执行:首次 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) @@ -376,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"` @@ -386,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{ @@ -414,111 +284,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/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/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/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_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/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/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/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_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/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 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, } } diff --git a/internal/service/chat_service.go b/internal/service/chat_service.go index 24d7410..6d7dd3a 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 创建聊天业务服务 @@ -158,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 { @@ -628,9 +648,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/chat_service_graph_quick.go b/internal/service/chat_service_graph_quick.go index 6e72cc9..bfc2dae 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" @@ -18,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" @@ -42,15 +44,84 @@ 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 读写。 type quickGraphState struct { Input *quickGraphInput - RewrittenQuery string + 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 } +// 查询改写意图类型 +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"` + NeedClarify bool `json:"need_clarify"` + ClarifyQuestion string `json:"clarify_question,omitempty"` + ClarifyOptions []string `json:"clarify_options,omitempty"` +} + +// rewriteMaxHistoryRounds 改写时拼入历史的最大轮数(每轮=user+assistant) +const rewriteMaxHistoryRounds = 3 + +// rewriteSystemPrompt 改写专用的 System Prompt +const rewriteSystemPrompt = `你是一个查询改写助手。根据用户的原始问题和对话历史,对问题进行改写并识别意图。 + +## 改写规则 +1. 消解指代:把"这个"、"那个方案"、"它"、"之前说的"等代词替换为对话历史中的具体名词 +2. 扩展关键词:补充与问题相关的同义词、上下位词,方便知识库检索 +3. 拆分复合问题:如果原问题包含多个子问题,改写为一个完整句子即可(不要拆成多行) +4. 保持原意:改写后的问题必须和原问题核心意图一致,不要引入新主题 + +## 意图识别 +- greeting: 问候语(你好、hi、在吗、早上好) +- chitchat: 闲聊(今天天气怎么样、讲个笑话、随便聊聊) +- question: 知识查询(业务问题、技术问题、需要从知识库找答案) +- identity: 身份确认(你是谁、你能做什么、介绍一下你自己) +- meta: 元问题(我的历史记录、你刚才说了什么、回顾对话) + +## 澄清追问判断 +当用户问题过于模糊、存在多种理解且无法从历史对话推断真实意图时,设置 need_clarify=true: +- 没有历史上下文时,单个指代性问题(如"那个方案"、"它")且知识库依赖强 → 追问 +- 问题包含可能冲突的关键概念(如"怎么导出数据"未指明导出格式/导出范围)→ 追问 +- 用户同时提及多个实体且未指明主体 → 追问 +以下情况**不要**追问: +- 打招呼、闲聊、身份类意图(greeting/chitchat/identity)→ 直接返回原问题 +- 有历史对话可以消解歧义 → 直接改写,need_clarify=false +- 即使问题有些宽泛,但可以给一个通用回答 → 直接回答,need_clarify=false + +## 输出格式 +严格使用 JSON,不要输出任何多余文字或 Markdown 代码块: +{"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" graphQuickNodeRetrieve = "retrieve" @@ -94,7 +165,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 +173,10 @@ func addQuickRewriteNode(g *einoCompose.Graph[*quickGraphInput, *schema.StreamRe ) } -// quickRewriteFn 节点 1 实现:暂存输入到 State 并原样返回查询 +// 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 { return "", apperrors.NewDefault(apperrors.CodeInvalidParam) @@ -113,22 +187,166 @@ func quickRewriteFn(ctx context.Context, input *quickGraphInput) (string, error) }); err != nil { return "", err } + startAt := time.Now() - rewritten := input.OriginalQuery + 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 }) + 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), + "need_clarify": fmt.Sprintf("%v", needClarify), + "rewrite_ms": durMs, + "model_id": input.ModelName, }) return rewritten, nil } +// doRewriteWithLLM 调 LLM 做改写,失败时 fallback 原始 query。 +// 返回 (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, false, "", nil + } + + // 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, false, "", nil + } + + // 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, false, "", nil + } + + // 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 + + // 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 返回的意图是否在合法枚举内 +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 +369,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 +395,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] @@ -472,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_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 } 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) { 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 + } } 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]) - } - } -}