diff --git a/cmd/migrate/main.go b/cmd/migrate/main.go deleted file mode 100644 index acdf51c..0000000 --- a/cmd/migrate/main.go +++ /dev/null @@ -1,45 +0,0 @@ -// Command migrate 执行数据库 schema 迁移(用 GORM AutoMigrate)。 -// 用法: go run ./cmd/migrate -package main - -import ( - "fmt" - "log" - - "solvify-agent/internal/model/entity" - "solvify-agent/pkg/config" - "solvify-agent/pkg/database" - "solvify-agent/pkg/logger" -) - -func main() { - _ = logger.InitDefault() - - cfg, err := config.Load("configs/config.yaml") - if err != nil { - log.Fatalf("加载配置失败: %v", err) - } - - db, err := database.OpenPostgreSQL(&cfg.Database.Postgres) - if err != nil { - log.Fatalf("连接 PostgreSQL 失败: %v", err) - } - sqlDB, _ := db.DB() - defer sqlDB.Close() - - fmt.Println("开始迁移...") - - // 迁移 ChatSession(自动补 pending_clarify / pending_checkpoint 列) - if err := db.AutoMigrate(&entity.ChatSession{}); err != nil { - log.Fatalf("迁移 ChatSession 失败: %v", err) - } - fmt.Println("✓ chat_sessions 已就绪") - - // 创建 agent_checkpoints 表 - if err := db.AutoMigrate(&entity.AgentCheckpoint{}); err != nil { - log.Fatalf("迁移 AgentCheckpoint 失败: %v", err) - } - fmt.Println("✓ agent_checkpoints 已就绪") - - fmt.Println("迁移完成 ✅") -} diff --git a/internal/agent/execute.go b/internal/agent/execute.go index d9f2e55..cf673c3 100644 --- a/internal/agent/execute.go +++ b/internal/agent/execute.go @@ -214,7 +214,7 @@ func (e *Engine) runAgent(ctx context.Context, req Request, chatModel model.Tool store := e.buildCheckpointStore(req.SessionID) runner := adk.NewRunner(ctx, adk.RunnerConfig{ Agent: agent, - EnableStreaming: true, + EnableStreaming: true, CheckPointStore: store, }) diff --git a/internal/agent/runner_adapter.go b/internal/agent/runner_adapter.go index aff5db4..922528d 100644 --- a/internal/agent/runner_adapter.go +++ b/internal/agent/runner_adapter.go @@ -96,60 +96,60 @@ func (e *Engine) runWithRunner( } // ── Interrupt 处理 ── - if agentEvent.Action != nil && agentEvent.Action.Interrupted != nil { - if !interruptSent { - interruptSent = true - ii := agentEvent.Action.Interrupted - interruptCtx := ii.InterruptContexts - var interruptID string - var infoStr string - if len(interruptCtx) > 0 { - interruptID = interruptCtx[0].ID - if s, ok := interruptCtx[0].Info.(string); ok { - infoStr = s - } + if agentEvent.Action != nil && agentEvent.Action.Interrupted != nil { + if !interruptSent { + interruptSent = true + ii := agentEvent.Action.Interrupted + interruptCtx := ii.InterruptContexts + var interruptID string + var infoStr string + if len(interruptCtx) > 0 { + interruptID = interruptCtx[0].ID + if s, ok := interruptCtx[0].Info.(string); ok { + infoStr = s } - logger.Infof("[Agent] 执行中断: checkpointID=%s, interruptID=%s, info=%s", checkpointID, interruptID, truncateStr(infoStr, 200)) + } + logger.Infof("[Agent] 执行中断: checkpointID=%s, interruptID=%s, info=%s", checkpointID, interruptID, truncateStr(infoStr, 200)) - infoType, infoData := parseInterruptInfo(infoStr) + infoType, infoData := parseInterruptInfo(infoStr) - if infoType == "clarify" { - eventCh <- Event{ - Type: EventInterrupt, - Title: "需要澄清", - Detail: getString(infoData, "question"), - Status: "clarify", - Error: interruptID, - CheckpointID: checkpointID, - InterruptID: interruptID, - IsClarify: true, - ClarifyQuestion: getString(infoData, "question"), - ClarifyOptions: getStringSlice(infoData, "options"), - ClarifyContext: getString(infoData, "context"), - Done: true, - } - } else { - // danger 或未知类型 → 按审批处理 - message := getString(infoData, "message") - if message == "" { - message = formatInterruptInfo(infoStr) - } - eventCh <- Event{ - Type: EventInterrupt, - Title: "需要人工确认", - Detail: truncateStr(message, 256), - Status: "interrupt", - Error: interruptID, - CheckpointID: checkpointID, - InterruptID: interruptID, - InterruptInfo: infoData, - Done: true, - } + if infoType == "clarify" { + eventCh <- Event{ + Type: EventInterrupt, + Title: "需要澄清", + Detail: getString(infoData, "question"), + Status: "clarify", + Error: interruptID, + CheckpointID: checkpointID, + InterruptID: interruptID, + IsClarify: true, + ClarifyQuestion: getString(infoData, "question"), + ClarifyOptions: getStringSlice(infoData, "options"), + ClarifyContext: getString(infoData, "context"), + Done: true, + } + } else { + // danger 或未知类型 → 按审批处理 + message := getString(infoData, "message") + if message == "" { + message = formatInterruptInfo(infoStr) + } + eventCh <- Event{ + Type: EventInterrupt, + Title: "需要人工确认", + Detail: truncateStr(message, 256), + Status: "interrupt", + Error: interruptID, + CheckpointID: checkpointID, + InterruptID: interruptID, + InterruptInfo: infoData, + Done: true, } - return } - continue + return } + continue + } if agentEvent.Output == nil || agentEvent.Output.MessageOutput == nil { continue @@ -449,8 +449,6 @@ func buildFallbackAnswer(sources []tool.SourceDocument) string { return sb.String() } - - func getString(m map[string]any, key string) string { if v, ok := m[key]; ok { if s, ok := v.(string); ok { diff --git a/internal/agent/types.go b/internal/agent/types.go index 9653cd6..47bc2ce 100644 --- a/internal/agent/types.go +++ b/internal/agent/types.go @@ -59,12 +59,12 @@ type Event struct { // ToolResult 工具调用结果(完整内容,供前端展示) ToolResult string `json:"tool_result,omitempty"` // interrupt 事件字段 - CheckpointID string `json:"checkpoint_id,omitempty"` - InterruptID string `json:"interrupt_id,omitempty"` + CheckpointID string `json:"checkpoint_id,omitempty"` + InterruptID string `json:"interrupt_id,omitempty"` InterruptInfo map[string]any `json:"interrupt_info,omitempty"` // clarify 事件字段(ask_clarify 触发的中断) - IsClarify bool `json:"is_clarify,omitempty"` - ClarifyQuestion string `json:"clarify_question,omitempty"` + IsClarify bool `json:"is_clarify,omitempty"` + ClarifyQuestion string `json:"clarify_question,omitempty"` ClarifyOptions []string `json:"clarify_options,omitempty"` ClarifyContext string `json:"clarify_context,omitempty"` } diff --git a/internal/observability/eino_callback.go b/internal/observability/eino_callback.go index 314db80..2944ffd 100644 --- a/internal/observability/eino_callback.go +++ b/internal/observability/eino_callback.go @@ -258,6 +258,11 @@ func einoOnEnd(rec Recorder) func(ctx context.Context, info *callbacks.RunInfo, } endSpanIfPresent(rec, ctx, state, attrs, SpanStatusOK, nil) + // End 后把 ctx 的当前 span 恢复为 parent:eino 会把 OnEnd 返回的 ctx 继续传给 + // 下一个兄弟节点,不恢复的话兄弟会错误地挂到这个已 End 的组件下面(链状嵌套)。 + if state.span != nil && state.span.parent != nil { + ctx = context.WithValue(ctx, currentSpanKey{}, state.span.parent) + } return ctx } } diff --git a/internal/observability/recorder.go b/internal/observability/recorder.go index 1ae7183..cb7e791 100644 --- a/internal/observability/recorder.go +++ b/internal/observability/recorder.go @@ -122,20 +122,26 @@ func RecorderFromContext(ctx context.Context) Recorder { return r } +// currentSpanKey 用于 ctx 直接携带当前自研 *Span 引用。 +// StartSpan 写入返回的 ctx,子 span 挂树和 CurrentSpanFromContext 都优先读它, +// 不依赖 OTel span 的 IsRecording 状态(eino 流式组件的 OnEnd 会提前 End 父 span)。 +type currentSpanKey struct{} + // CurrentSpanFromContext 定位当前正在运行的 span。 // 返回项目自研 *Span(包含 otelSpan 字段),找不到时返回 nil,调用方静默降级。 func CurrentSpanFromContext(ctx context.Context) *Span { + // 优先走 currentSpanKey:与 End 状态解耦,流式场景父 span 可能已被回调提前 End + if s, ok := ctx.Value(currentSpanKey{}).(*Span); ok && s != nil { + return s + } + // 兜底:OTel span 指针反查(不经 StartSpan 返回 ctx 的旧调用路径) otelSpan := trace.SpanFromContext(ctx) if otelSpan == nil { return nil } - // SpanFromContext 返回的是 noopSpan 时,没有项目自研 Span 对应 - // 这里通过 otelSpan.IsRecording() 判断是否是真实 span if !otelSpan.IsRecording() { return nil } - // 项目自研 Span 通过 tracer.Start 时存放在 otelSpan 的私有字段里, - // 但 OTel 接口不暴露这个,所以用 sync.Map 按 span 指针关联。 if s, ok := spanByOtel.Load(otelSpan); ok { return s.(*Span) } @@ -269,13 +275,14 @@ func (r *defaultRecorder) StartSpan(ctx context.Context, name string, component ctx = context.WithValue(ctx, traceIDKey, traceID) } - // 落库 parent-child:先从入参 ctx 找 parent 自研 Span(必须在 tracer.Start 之前, - // 因为 tracer.Start 返回的 ctxWithSpan 会把当前 span 设为 current,再找就是自己了)。 + // 落库 parent-child:优先从入参 ctx 的 currentSpanKey 取 parent 自研 Span 引用。 + // 不能用 trace.SpanFromContext(ctx).IsRecording() 判断:eino 流式组件(如 adk Agent)的 + // OnEnd 会在输出流刚返回时就 End 掉 span,而该 span 仍作为 ctx 的 current 传给后续子组件, + // IsRecording()=false 会让整棵子树找不到 parent 变成孤儿(深度模式 trace 断裂的根因)。 + // ctx 引用与 End 状态解耦后,已 End 的 span 仍是合法 parent。 var parentSpan *Span - if parentOtelSpan := trace.SpanFromContext(ctx); parentOtelSpan != nil && parentOtelSpan.IsRecording() { - if v, ok := spanByOtel.Load(parentOtelSpan); ok { - parentSpan, _ = v.(*Span) - } + if ps, ok := ctx.Value(currentSpanKey{}).(*Span); ok && ps != nil { + parentSpan = ps } // 用 OTel tracer.Start 创建运行时 span,OTel 自动管理 parent-child 关系。 @@ -300,22 +307,24 @@ func (r *defaultRecorder) StartSpan(ctx context.Context, name string, component s.ParentID = parentSpan.SpanID } - // 关联 OTel span 到自研 Span,CurrentSpanFromContext 用 + // 关联 OTel span 到自研 Span,CurrentSpanFromContext 兜底路径用 spanByOtel.Store(otelSpan, s) - // 登记 traceState 的根 span + // ctx 携带自研 Span 引用:子 span 的 parent 查找与 CurrentSpanFromContext 走这里, + // 与 span 的 End 状态解耦(见上方 parent 查找注释)。 + ctxWithSpan = context.WithValue(ctxWithSpan, currentSpanKey{}, s) + + // 登记 traceState 的根 span。用 LoadOrStore 防止后到的孤儿 span 覆盖已登记的树。 isChatRoot := ctx.Value(rootAttrsKey) != nil - if parentSpan == nil && !isChatRoot { - st := &traceState{Trace: &Trace{ID: traceID, Root: s, SampleRate: r.cfg.SamplingRate}} - r.traceStates.Store(traceID, st) - } else if parentSpan == nil && isChatRoot { - // chat 场景:存入 traceStates 供 publishTrace 合并 children,但标记为中间 root 不触发 finalizeTrace - st := &traceState{Trace: &Trace{ID: traceID, Root: s, SampleRate: r.cfg.SamplingRate}} - r.traceStates.Store(traceID, st) - if s.Attrs == nil { - s.Attrs = Attrs{} + if parentSpan == nil { + if isChatRoot { + // chat 场景:Root 是 chat.deep/chat.quick 等中间根,publishTrace 时合并到合成 chat.request 下 + if s.Attrs == nil { + s.Attrs = Attrs{} + } + s.Attrs["__chat_intermediate_root"] = true } - s.Attrs["__chat_intermediate_root"] = true + r.traceStates.LoadOrStore(traceID, &traceState{Trace: &Trace{ID: traceID, Root: s, SampleRate: r.cfg.SamplingRate}}) } r.metrics.obsSpanStartTotal.WithLabelValues(string(component)).Inc() diff --git a/internal/rag/eino_adapter.go b/internal/rag/eino_adapter.go index 19c3ece..d3744fa 100644 --- a/internal/rag/eino_adapter.go +++ b/internal/rag/eino_adapter.go @@ -60,7 +60,7 @@ func WithUserID(uid string) retriever.Option { // 等内部检索逻辑完全不变,仅做输入/输出格式对齐,使上游(eino Graph、Agent、 // 可观测性 callback)能按 eino 统一组件标准接入。 type EinoRetrieverAdapter struct { - inner Retriever + inner Retriever defaultTopK int } diff --git a/internal/service/chat_budget.go b/internal/service/chat_budget.go new file mode 100644 index 0000000..eb46097 --- /dev/null +++ b/internal/service/chat_budget.go @@ -0,0 +1,193 @@ +package service + +import ( + "solvify-agent/internal/model/entity" + "solvify-agent/pkg/tokenutil" +) + +// calculateContextBudgets 根据模型最大上下文窗口 + 工具定义占用,分配历史、检索、记忆的 token 预算。 +// +// P0-④ 关键修复:toolsTokens (深度模式/多工具场景的工具 JSON Schema 真 token 数) 必须先从总窗口 +// 扣除,再分配回复预留和固定预留,否则多工具时直接把历史 + 检索预算挤成负数或零。 +// 同时所有预算最终都按 0.95*maxCtx 的安全顶封顶,给偶发的角色名/special token 留余量。 +func calculateContextBudgets(maxContextLength int, toolsTokens ...int) (historyBudget, retrievalBudget, memoryBudget int) { + toolReserve := 0 + if len(toolsTokens) > 0 && toolsTokens[0] > 0 { + toolReserve = toolsTokens[0] + } + if maxContextLength <= 0 { + maxContextLength = 8192 + } + // 0.95 的安全顶:角色标记、特殊 token、工具结果 JSON 序列化扩展,都容易让"算刚好"爆。 + safeCap := int(float64(maxContextLength) * 0.95) + if toolReserve >= safeCap { + // 工具定义已经吃掉整个窗口(极端异常配置):给历史留 200 保底,其他归零 + return 200, 0, 0 + } + remaining := safeCap - toolReserve + + // 1. 回复预留:不超过 4096 或 safeCap 的 1/4 + completionReserved := 4096 + if remaining/4 < completionReserved { + completionReserved = remaining / 4 + } + if completionReserved < 200 { + completionReserved = 200 + } + remaining -= completionReserved + + // 2. 固定预留:System Prompt 基础骨架 + 当前 user question 包装 + 安全边距 + fixedReserved := 1500 + if remaining-fixedReserved < 500 { + fixedReserved = remaining / 4 + if fixedReserved < 300 { + fixedReserved = 300 + } + } + remaining -= fixedReserved + + // 3. 检索上下文块(RAG context)优先保证至少 500,最多取 min(3000, remaining/3) + retrievalBudget = 3000 + if remaining-retrievalBudget < 1000 { + retrievalBudget = remaining / 3 + } + if retrievalBudget < 500 { + retrievalBudget = 500 + } + if retrievalBudget > remaining { + retrievalBudget = max(remaining, 0) + } + remaining -= retrievalBudget + + // 4. 记忆预算:与模型窗口成正比,但封顶。8k 及以下不给记忆,省出空间给历史。 + memoryBudget = 800 + switch { + case maxContextLength >= 32000: + memoryBudget = 1200 + case maxContextLength >= 16000: + memoryBudget = 1000 + case maxContextLength <= 8192: + memoryBudget = 400 + } + if memoryBudget > remaining { + memoryBudget = max(remaining/2, 0) + } + remaining -= memoryBudget + + // 5. 历史消息预算:剩下的全给历史,保底 500,封顶 6000(防止过大的上下文拖慢模型推理) + historyBudget = remaining + historyBudget = max(historyBudget, 500) + if historyBudget > 6000 { + historyBudget = 6000 + } + return +} + +func max(a, b int) int { + if a > b { + return a + } + return b +} + +// truncateHistoryByTokens 按轮对(user + assistant 配对)从尾部保留历史消息, +// 保证最后一条 user 问题不被截断,且 assistant 不会孤立存在。 +func truncateHistoryByTokens(messages []entity.ChatMessage, maxTokens int, modelName string) []entity.ChatMessage { + if maxTokens <= 0 { + return nil + } + n := len(messages) + if n == 0 { + return nil + } + + tailIdx := n - 1 + tailReserved := 0 + tailCutMsg := (*entity.ChatMessage)(nil) + if messages[tailIdx].Role == "user" { + t := tokenutil.CountTokens(messages[tailIdx].Content, modelName) + if t > maxTokens { + cut, actual := tokenutil.TruncateByTokens(messages[tailIdx].Content, modelName, max(maxTokens-50, 50)) + if actual > 0 { + m := messages[tailIdx] + m.Content = cut + tailCutMsg = &m + tailReserved = actual + } + } else { + tailReserved = t + } + } + + pairs := make([][]entity.ChatMessage, 0, 4) + total := tailReserved + i := n - 1 + if messages[tailIdx].Role == "user" { + i-- + } + for i >= 0 { + if messages[i].Role != "assistant" { + i-- + continue + } + a := i + u := -1 + for j := i - 1; j >= 0; j-- { + if messages[j].Role == "user" { + u = j + break + } + } + if u < 0 { + break + } + pairTokens := 0 + for k := u; k <= a; k++ { + pairTokens += tokenutil.CountTokens(messages[k].Content, modelName) + } + if total+pairTokens > maxTokens { + remain := maxTokens - total + if remain >= 120 { + m := messages[u] + cut, actual := truncateContentHeadByTokens(m.Content, modelName, remain) + if actual > 0 { + m.Content = cut + "\n\n(内容过长,已截断)" + pairs = append(pairs, []entity.ChatMessage{m}) + } + } + break + } + total += pairTokens + pair := append([]entity.ChatMessage(nil), messages[u:a+1]...) + pairs = append(pairs, pair) + i = u - 1 + } + + for l, r := 0, len(pairs)-1; l < r; l, r = l+1, r-1 { + pairs[l], pairs[r] = pairs[r], pairs[l] + } + out := make([]entity.ChatMessage, 0, 2*len(pairs)+1) + for _, p := range pairs { + out = append(out, p...) + } + if tailCutMsg != nil { + out = append(out, *tailCutMsg) + } else if messages[tailIdx].Role == "user" && tailReserved > 0 { + out = append(out, messages[tailIdx]) + } + return out +} + +// truncateContentHeadByTokens 从"头"按真 BPE 截断到至多 maxTokens。 +// 与 tokenutil.TruncateByTokens 的区别:后者默认从左往右,这里再包一层 +// 统一返回(截断后文本,实际 token)。 +func truncateContentHeadByTokens(content, modelName string, maxTokens int) (string, int) { + return tokenutil.TruncateByTokens(content, modelName, maxTokens) +} + +// truncateContentByTokens 保留旧签名给现有调用方,内部转成新接口。 +// 新代码优先用 tokenutil.TruncateByTokens,可拿到实际用了多少 token。 +func truncateContentByTokens(content string, maxTokens int) string { + out, _ := tokenutil.TruncateByTokens(content, "", maxTokens) + return out +} diff --git a/internal/service/chat_helpers.go b/internal/service/chat_helpers.go new file mode 100644 index 0000000..888ae03 --- /dev/null +++ b/internal/service/chat_helpers.go @@ -0,0 +1,29 @@ +package service + +import ( + "context" + + "github.com/google/uuid" +) + +func requestIDFromCtx(ctx context.Context) string { + type iKey string + const key iKey = "request_id" + if v := ctx.Value(key); v != nil { + if s, ok := v.(string); ok { + return s + } + } + return uuid.New().String() +} + +func mergeStrMap(a, b map[string]string) map[string]string { + out := make(map[string]string, len(a)+len(b)) + for k, v := range a { + out[k] = v + } + for k, v := range b { + out[k] = v + } + return out +} diff --git a/internal/service/chat_interface.go b/internal/service/chat_interface.go index a6aab85..d195c2f 100644 --- a/internal/service/chat_interface.go +++ b/internal/service/chat_interface.go @@ -10,10 +10,10 @@ import ( // FeedbackRequest 反馈提交请求 type FeedbackRequest struct { - Rating int `json:"rating"` - Reasons []string `json:"reasons"` - Comment string `json:"comment"` - IsQuick bool `json:"is_quick_reply"` + Rating int `json:"rating"` + Reasons []string `json:"reasons"` + Comment string `json:"comment"` + IsQuick bool `json:"is_quick_reply"` } // FeedbackListResponse 反馈列表响应 @@ -24,22 +24,22 @@ type FeedbackListResponse struct { // TraceAgentTaskResponse Agent 任务追踪详情 type TraceAgentTaskResponse struct { - ID string `json:"id"` - TraceID string `json:"trace_id,omitempty"` - SessionID string `json:"session_id,omitempty"` - UserID string `json:"user_id,omitempty"` - ModelID string `json:"model_id,omitempty"` - SearchMode string `json:"search_mode,omitempty"` - StartedAt string `json:"started_at"` - EndedAt string `json:"ended_at,omitempty"` - TotalSteps int `json:"total_steps,omitempty"` - ToolCalls int `json:"tool_calls,omitempty"` - Status string `json:"status,omitempty"` - AbortReason string `json:"abort_reason,omitempty"` - TokensPrompt int `json:"tokens_prompt,omitempty"` - TokensCompletion int `json:"tokens_completion,omitempty"` - TotalCost float64 `json:"total_cost,omitempty"` - ErrorSummary string `json:"error_summary,omitempty"` + ID string `json:"id"` + TraceID string `json:"trace_id,omitempty"` + SessionID string `json:"session_id,omitempty"` + UserID string `json:"user_id,omitempty"` + ModelID string `json:"model_id,omitempty"` + SearchMode string `json:"search_mode,omitempty"` + StartedAt string `json:"started_at"` + EndedAt string `json:"ended_at,omitempty"` + TotalSteps int `json:"total_steps,omitempty"` + ToolCalls int `json:"tool_calls,omitempty"` + Status string `json:"status,omitempty"` + AbortReason string `json:"abort_reason,omitempty"` + TokensPrompt int `json:"tokens_prompt,omitempty"` + TokensCompletion int `json:"tokens_completion,omitempty"` + TotalCost float64 `json:"total_cost,omitempty"` + ErrorSummary string `json:"error_summary,omitempty"` } // TraceAgentStepResponse Agent 单步追踪详情 @@ -58,22 +58,22 @@ type TraceAgentStepResponse struct { // TraceResponse 单次追踪详情响应 type TraceResponse struct { - ID string `json:"id"` - RequestID string `json:"request_id,omitempty"` - UserID string `json:"user_id,omitempty"` - SessionID string `json:"session_id,omitempty"` - SearchMode string `json:"search_mode,omitempty"` - SampleRate float64 `json:"sample_rate,omitempty"` - Sampled bool `json:"sampled"` - DurationMs int64 `json:"duration_ms,omitempty"` - Status string `json:"status,omitempty"` - Error string `json:"error,omitempty"` - Attrs any `json:"attrs,omitempty"` - AttrsDisplay any `json:"attrs_display,omitempty"` - SpanTree any `json:"span_tree,omitempty"` - AgentTask *TraceAgentTaskResponse `json:"agent_task,omitempty"` - AgentSteps []TraceAgentStepResponse `json:"agent_steps,omitempty"` - CreatedAt string `json:"created_at"` + ID string `json:"id"` + RequestID string `json:"request_id,omitempty"` + UserID string `json:"user_id,omitempty"` + SessionID string `json:"session_id,omitempty"` + SearchMode string `json:"search_mode,omitempty"` + SampleRate float64 `json:"sample_rate,omitempty"` + Sampled bool `json:"sampled"` + DurationMs int64 `json:"duration_ms,omitempty"` + Status string `json:"status,omitempty"` + Error string `json:"error,omitempty"` + Attrs any `json:"attrs,omitempty"` + AttrsDisplay any `json:"attrs_display,omitempty"` + SpanTree any `json:"span_tree,omitempty"` + AgentTask *TraceAgentTaskResponse `json:"agent_task,omitempty"` + AgentSteps []TraceAgentStepResponse `json:"agent_steps,omitempty"` + CreatedAt string `json:"created_at"` } // TraceListResponse 追踪列表响应 diff --git a/internal/service/chat_mode_registry.go b/internal/service/chat_mode_registry.go new file mode 100644 index 0000000..6447890 --- /dev/null +++ b/internal/service/chat_mode_registry.go @@ -0,0 +1,59 @@ +package service + +import ( + "context" + + requestdto "solvify-agent/internal/model/dto/request" + dto "solvify-agent/internal/model/dto/response" +) + +const ( + chatModeQuick = "quick" + chatModeDeep = "smart-reasoning" +) + +// chatMode 对话模式处理器接口。 +// 新增对话模式只需:1) 实现 chatMode 接口;2) 在 modeRegistry 注册一行。 +type chatMode interface { + Name() string + Handle(ctx context.Context, s *chatService, userID, sessionID, userMsgID string, + req requestdto.SendMessageRequest, eventCh chan<- dto.StreamEvent) +} + +// modeRegistry 对话模式注册表:搜索模式字符串 → 处理器实例。 +// 默认值 "" 和未知值都回退到快速检索模式。 +var modeRegistry = map[string]chatMode{ + chatModeQuick: &quickModeHandler{}, + "": &quickModeHandler{}, + chatModeDeep: &deepModeHandler{}, +} + +// getModeHandler 根据 searchMode 查找处理器,找不到回退到快速模式。 +func getModeHandler(searchMode string) chatMode { + if h, ok := modeRegistry[searchMode]; ok { + return h + } + return modeRegistry[chatModeQuick] +} + +// ─── 快速检索模式 ────────────────────────────────────────── + +type quickModeHandler struct{} + +func (h *quickModeHandler) Name() string { return chatModeQuick } + +func (h *quickModeHandler) Handle(ctx context.Context, s *chatService, userID, sessionID, userMsgID string, + req requestdto.SendMessageRequest, eventCh chan<- dto.StreamEvent) { + s.processMessageGraphQuick(ctx, userID, sessionID, userMsgID, req, eventCh) +} + +// ─── 深度思考模式 ────────────────────────────────────────── + +type deepModeHandler struct{} + +func (h *deepModeHandler) Name() string { return chatModeDeep } + +func (h *deepModeHandler) Handle(ctx context.Context, s *chatService, userID, sessionID, userMsgID string, + req requestdto.SendMessageRequest, eventCh chan<- dto.StreamEvent) { + s.processDeepMode(ctx, userID, sessionID, userMsgID, req, eventCh) +} diff --git a/internal/service/chat_prompt_builder.go b/internal/service/chat_prompt_builder.go index c4a1a46..2e42d1e 100644 --- a/internal/service/chat_prompt_builder.go +++ b/internal/service/chat_prompt_builder.go @@ -248,11 +248,6 @@ func (b *PromptBuilder) BuildHistory(history []entity.ChatMessage) []*schema.Mes return msgs } -// BuildHistoryForAgent agent.Request.History 深度模式专用(复用 BuildHistory,语义更清晰) -func (b *PromptBuilder) BuildHistoryForAgent(history []entity.ChatMessage) []entity.ChatMessage { - return history -} - // BuildAgentRequestFields 深度模式:把 builder 中的摘要 / 记忆 / 用户上下文填充到 agent.Request 对应字段 // System Prompt 由 PromptBuilder.BuildSystem() 统一产出后塞到 agent.Request.SystemPrompt, // agent.runAgent 只负责在前面拼接 ReAct 规则,保证快速/深度两模式的摘要/记忆/偏好注入完全一致。 @@ -291,26 +286,6 @@ func (b *PromptBuilder) toAgentPromptUserContext() agentpkg.PromptUserContext { } } -// takeStringOrProfile 从 Profile entity 取出字段值(或空) -func takeStringOrProfile(u *entity.User, field string) string { - if u == nil { - return "" - } - switch field { - case "Department": - return u.Department - case "Position": - return u.Position - case "Expertise": - return u.Expertise - case "PreferredLanguage": - return u.PreferredLanguage - case "Timezone": - return u.Timezone - } - return "" -} - // UserContext 注入到 System Prompt 的用户上下文信息 type UserContext struct { ID string @@ -358,27 +333,3 @@ func (u UserContext) WithPreference(p *entity.UserPreference) UserContext { u.CitationStyle = p.CitationStyle return u } - -// takeAnswerStyle 取用户回答风格,空对象返回空串 -func takeAnswerStyle(p *entity.UserPreference) string { - if p == nil { - return "" - } - return p.AnswerStyle -} - -// takeTableFirst 取是否优先表格呈现,空对象默认 true -func takeTableFirst(p *entity.UserPreference) bool { - if p == nil { - return true - } - return p.UseMarkdownTable -} - -// takeCitationStyle 取引用格式,空对象默认 section_title -func takeCitationStyle(p *entity.UserPreference) string { - if p == nil { - return "section_title" - } - return p.CitationStyle -} diff --git a/internal/service/chat_service.go b/internal/service/chat_service.go index 934cf45..d1b421f 100644 --- a/internal/service/chat_service.go +++ b/internal/service/chat_service.go @@ -2,7 +2,6 @@ package service import ( "context" - "encoding/json" "fmt" "strings" "time" @@ -26,7 +25,7 @@ import ( ) const ( -// sessionStatusActive 表示会话处于活跃状态 + // sessionStatusActive 表示会话处于活跃状态 sessionStatusActive = "active" ) @@ -159,7 +158,7 @@ func (s *chatService) SendMessage(ctx context.Context, userID, sessionID string, abortReason = "runtime_panic" } }() - + // 澄清追问恢复:检查 session 是否有待处理的澄清 // 未超时 → 清掉 pending,历史自然串成 [user→assistant追问→user回答],正常跑流程 // 超时 → 清掉 pending,正常跑(用户已遗忘之前的追问,新消息按新问题处理) @@ -177,12 +176,8 @@ func (s *chatService) SendMessage(ctx context.Context, userID, sessionID string, } } } - - if searchMode == "smart-reasoning" { - s.processDeepMode(ctx, userID, sessionID, userMsgID, req, eventCh) - } else { - s.processMessage(ctx, userID, sessionID, userMsgID, req, eventCh) - } + + getModeHandler(searchMode).Handle(ctx, s, userID, sessionID, userMsgID, req, eventCh) if s.obs != nil { s.obs.FlushTrace(ctx, userID, sessionID, userMsgID) } @@ -196,17 +191,6 @@ func (s *chatService) SendMessage(ctx context.Context, userID, sessionID string, return eventCh, nil } -func requestIDFromCtx(ctx context.Context) string { - type iKey string - const key iKey = "request_id" - if v := ctx.Value(key); v != nil { - if s, ok := v.(string); ok { - return s - } - } - return uuid.New().String() -} - // updateUserLastModel 更新用户上次使用的模型(缓存比对策略) func (s *chatService) updateUserLastModel(ctx context.Context, userID, modelID string) { cacheKey := "user:model:" + userID @@ -514,18 +498,6 @@ func (s *chatService) initContext(ctx context.Context, userID, sessionID, modelI return client, enhancedCtx, nil } -// mergeStrMap 返回 base+extra 合并后的 map(不修改原始 map,避免 label 串场) -func mergeStrMap(a, b map[string]string) map[string]string { - out := make(map[string]string, len(a)+len(b)) - for k, v := range a { - out[k] = v - } - for k, v := range b { - out[k] = v - } - return out -} - // loadUserContext 加载用户基本信息,失败时返回空上下文(不阻断主流程) func (s *chatService) loadUserContext(ctx context.Context, userID string) UserContext { if s.userRepo == nil || userID == "" { @@ -539,91 +511,6 @@ func (s *chatService) loadUserContext(ctx context.Context, userID string) UserCo return NewUserContext(*user) } -// calculateContextBudgets 根据模型最大上下文窗口 + 工具定义占用,分配历史、检索、记忆的 token 预算。 -// -// P0-④ 关键修复:toolsTokens (深度模式/多工具场景的工具 JSON Schema 真 token 数) 必须先从总窗口 -// 扣除,再分配回复预留和固定预留,否则多工具时直接把历史 + 检索预算挤成负数或零。 -// 同时所有预算最终都按 0.95*maxCtx 的安全顶封顶,给偶发的角色名/special token 留余量。 -func calculateContextBudgets(maxContextLength int, toolsTokens ...int) (historyBudget, retrievalBudget, memoryBudget int) { - toolReserve := 0 - if len(toolsTokens) > 0 && toolsTokens[0] > 0 { - toolReserve = toolsTokens[0] - } - if maxContextLength <= 0 { - maxContextLength = 8192 - } - // 0.95 的安全顶:角色标记、特殊 token、工具结果 JSON 序列化扩展,都容易让"算刚好"爆。 - safeCap := int(float64(maxContextLength) * 0.95) - if toolReserve >= safeCap { - // 工具定义已经吃掉整个窗口(极端异常配置):给历史留 200 保底,其他归零 - return 200, 0, 0 - } - remaining := safeCap - toolReserve - - // 1. 回复预留:不超过 4096 或 safeCap 的 1/4 - completionReserved := 4096 - if remaining/4 < completionReserved { - completionReserved = remaining / 4 - } - if completionReserved < 200 { - completionReserved = 200 - } - remaining -= completionReserved - - // 2. 固定预留:System Prompt 基础骨架 + 当前 user question 包装 + 安全边距 - fixedReserved := 1500 - if remaining-fixedReserved < 500 { - fixedReserved = remaining / 4 - if fixedReserved < 300 { - fixedReserved = 300 - } - } - remaining -= fixedReserved - - // 3. 检索上下文块(RAG context)优先保证至少 500,最多取 min(3000, remaining/3) - retrievalBudget = 3000 - if remaining-retrievalBudget < 1000 { - retrievalBudget = remaining / 3 - } - if retrievalBudget < 500 { - retrievalBudget = 500 - } - if retrievalBudget > remaining { - retrievalBudget = max(remaining, 0) - } - remaining -= retrievalBudget - - // 4. 记忆预算:与模型窗口成正比,但封顶。8k 及以下不给记忆,省出空间给历史。 - memoryBudget = 800 - switch { - case maxContextLength >= 32000: - memoryBudget = 1200 - case maxContextLength >= 16000: - memoryBudget = 1000 - case maxContextLength <= 8192: - memoryBudget = 400 - } - if memoryBudget > remaining { - memoryBudget = max(remaining/2, 0) - } - remaining -= memoryBudget - - // 5. 历史消息预算:剩下的全给历史,保底 500,封顶 6000(防止过大的上下文拖慢模型推理) - historyBudget = remaining - historyBudget = max(historyBudget, 500) - if historyBudget > 6000 { - historyBudget = 6000 - } - return -} - -func max(a, b int) int { - if a > b { - return a - } - return b -} - // resolveClient 根据模型配置解析 LLM 客户端 func (s *chatService) resolveClient(ctx context.Context, userID, modelID, modelType string) (*llm.OpenAIClient, error) { var cfg llm.ModelConfig @@ -731,6 +618,7 @@ func (s *chatService) bgComputeEmbedding(ctx context.Context, messageID, content } }() } + // truncateHistoryByTokens 按真 BPE token 预算从尾部向前保留"完整轮对"。 // // 关键修复(P0-②):旧代码"从尾部逐个 append 头插"实际是按时间正序塞,但因为 @@ -738,579 +626,3 @@ func (s *chatService) bgComputeEmbedding(ctx context.Context, messageID, content // 最新的用户问题反而被截断。现在按轮对(user + assistant 配对)从尾部保留, // 并且保证最后一条必须是本轮 user(最后一条如果是 assistant 会在调用方补 user query, // 所以我们只确保"若最后一条恰好是 user 就不能截掉")。 -func truncateHistoryByTokens(messages []entity.ChatMessage, maxTokens int, modelName string) []entity.ChatMessage { - if maxTokens <= 0 { - return nil - } - n := len(messages) - if n == 0 { - return nil - } - - // 1. 对齐轮对边界:若最后一条是 user,尾部刚好半个轮对 → 预占它的 token, - // 保证无论如何都能保留。 - tailIdx := n - 1 - tailReserved := 0 - tailCutMsg := (*entity.ChatMessage)(nil) - if messages[tailIdx].Role == "user" { - t := tokenutil.CountTokens(messages[tailIdx].Content, modelName) - if t > maxTokens { - cut, actual := tokenutil.TruncateByTokens(messages[tailIdx].Content, modelName, max(maxTokens-50, 50)) - if actual > 0 { - m := messages[tailIdx] - m.Content = cut - tailCutMsg = &m - tailReserved = actual - } - } else { - tailReserved = t - } - } - - // 2. 轮对从尾部向前保留,遇到 user 开始收集一对;到 budget 不够就丢弃该整轮 - // (不保留半截 assistant,否则会出现"assistant 的问题没人问") - pairs := make([][]entity.ChatMessage, 0, 4) - total := tailReserved - i := n - 1 - if messages[tailIdx].Role == "user" { - i-- - } - for i >= 0 { - // 找一对 user i0..i:先从 i 往前走到最近的 user,再回退一条 assistant - if messages[i].Role != "assistant" { - // 异常孤立消息(中间多了一条 user),直接跳过这一条保持轮对完整 - i-- - continue - } - a := i - u := -1 - for j := i - 1; j >= 0; j-- { - if messages[j].Role == "user" { - u = j - break - } - } - if u < 0 { - break - } - pairTokens := 0 - for k := u; k <= a; k++ { - pairTokens += tokenutil.CountTokens(messages[k].Content, modelName) - } - if total+pairTokens > maxTokens { - // 还有至少 120 token 空间 → 截断 user 头部保留主题,丢 assistant - remain := maxTokens - total - if remain >= 120 { - m := messages[u] - cut, actual := truncateContentHeadByTokens(m.Content, modelName, remain) - if actual > 0 { - m.Content = cut + "\n\n(内容过长,已截断)" - pairs = append(pairs, []entity.ChatMessage{m}) - } - } - break - } - total += pairTokens - pair := append([]entity.ChatMessage(nil), messages[u:a+1]...) - pairs = append(pairs, pair) - i = u - 1 - } - - // 3. 组装:轮对顺序是"先收集的靠后",所以要 reverse - for l, r := 0, len(pairs)-1; l < r; l, r = l+1, r-1 { - pairs[l], pairs[r] = pairs[r], pairs[l] - } - out := make([]entity.ChatMessage, 0, 2*len(pairs)+1) - for _, p := range pairs { - out = append(out, p...) - } - if tailCutMsg != nil { - out = append(out, *tailCutMsg) - } else if messages[tailIdx].Role == "user" && tailReserved > 0 { - out = append(out, messages[tailIdx]) - } - return out -} - -// truncateContentHeadByTokens 从"头"按真 BPE 截断到至多 maxTokens。 -// 与 tokenutil.TruncateByTokens 的区别:后者默认从左往右,这里再包一层 -// 统一返回(截断后文本,实际 token)。 -func truncateContentHeadByTokens(content, modelName string, maxTokens int) (string, int) { - return tokenutil.TruncateByTokens(content, modelName, maxTokens) -} - -// truncateContentByTokens 保留旧签名给现有调用方,内部转成新接口。 -// 新代码优先用 tokenutil.TruncateByTokens,可拿到实际用了多少 token。 -func truncateContentByTokens(content string, maxTokens int) string { - out, _ := tokenutil.TruncateByTokens(content, "", maxTokens) - return out -} - -// ─── 可观测:反馈 / Trace / Metrics 查询接口 ─────────────────────────────────── - -// SubmitFeedback 提交消息反馈 -func (s *chatService) SubmitFeedback(ctx context.Context, userID, messageID string, req FeedbackRequest) error { - if req.Rating != 1 && req.Rating != -1 { - return fmt.Errorf("rating 必须为 1 或 -1") - } - if messageID == "" || userID == "" { - return fmt.Errorf("message_id / user_id 不能为空") - } - msg, err := s.messageRepo.FindByID(ctx, messageID) - if err != nil { - return fmt.Errorf("查询消息失败: %w", err) - } - if msg == nil { - return fmt.Errorf("消息不存在或无权限") - } - if msg.SessionID != "" { - if vErr := s.validateSession(ctx, userID, msg.SessionID); vErr != nil { - return fmt.Errorf("消息不存在或无权限") - } - } - var traceID string - if raw := msg.Metadata; len(raw) > 0 { - if meta := metadataAsMap(raw); meta != nil { - if v, ok := meta["trace_id"].(string); ok { - traceID = v - } - } - } - primaryTag := "" - if len(req.Reasons) > 0 { - primaryTag = req.Reasons[0] - } - fb := &entity.MessageFeedback{ - ID: uuid.New().String(), - MessageID: messageID, - UserID: userID, - SessionID: msg.SessionID, - Rating: req.Rating, - ReasonTag: primaryTag, - Comment: req.Comment, - IsQuick: req.IsQuick, - TraceID: traceID, - } - fb.SetReasons(req.Reasons) - if s.obsRepo != nil { - if e := s.obsRepo.CreateFeedback(ctx, fb); e != nil { - return fmt.Errorf("保存反馈失败: %w", e) - } - } - if s.obs != nil { - s.obs.Incr(ctx, "chat_feedback_total", map[string]string{ - "rating": ratingLabel(req.Rating), - "reason_tag": reasonTagOrDefault(primaryTag), - "has_comment": boolLabel(req.Comment != ""), - }, 1) - } - if s.obs != nil { - s.obs.RecordFeedback(&observability.Feedback{ - MessageID: fb.MessageID, - UserID: fb.UserID, - SessionID: fb.SessionID, - Rating: fb.Rating, - Reasons: fb.Reasons(), - Comment: fb.Comment, - TraceID: fb.TraceID, - CreatedAt: fb.CreatedAt, - }) - } - return nil -} - -// ListFeedbacks 分页查询用户反馈列表 -func (s *chatService) ListFeedbacks(ctx context.Context, userID string, offset, limit int) (FeedbackListResponse, error) { - if limit <= 0 { - limit = 20 - } - if limit > 100 { - limit = 100 - } - if offset < 0 { - offset = 0 - } - if s.obsRepo == nil { - return FeedbackListResponse{Total: 0, Feedbacks: []any{}}, nil - } - list, total, err := s.obsRepo.ListByUser(ctx, userID, offset, limit) - if err != nil { - return FeedbackListResponse{}, err - } - type out struct { - entity.MessageFeedback - Reasons []string `json:"reasons"` - } - items := make([]any, 0, len(list)) - for _, f := range list { - items = append(items, out{MessageFeedback: f, Reasons: f.Reasons()}) - } - return FeedbackListResponse{Total: total, Feedbacks: items}, nil -} - -// buildTraceResponse 构建追踪详情响应 -func (s *chatService) buildTraceResponse(t *entity.ChatTrace, includeAgentDetail bool) TraceResponse { - if t == nil { - return TraceResponse{} - } - resp := TraceResponse{ - ID: t.ID, - RequestID: t.RequestID, - UserID: t.UserID, - SessionID: t.SessionID, - SearchMode: extractSearchMode(t.Attrs), - SampleRate: t.SampleRate, - Sampled: t.Sampled, - DurationMs: t.DurationMs, - Status: t.Status, - Error: t.Error, - Attrs: t.Attrs, - SpanTree: t.SpanTree, - CreatedAt: t.CreatedAt.Format("2006-01-02 15:04:05"), - } - if !includeAgentDetail || s.obsRepo == nil { - return resp - } - task, steps, _ := s.obsRepo.FindByTraceID(context.Background(), t.ID) - resp.AgentTask = chatAgentTaskEntityToResponse(task) - resp.AgentSteps = chatAgentStepEntityToResponse(steps) - // 有 AgentStep 信息时,把 TotalSteps / ToolCalls 反填回 AgentTask(如果之前 MarkEnded 没填充好) - if resp.AgentTask != nil && len(resp.AgentSteps) > 0 { - toolCalls := 0 - for _, st := range resp.AgentSteps { - if st.ToolName != "" && st.ToolName != "llm.reasoning" { - toolCalls++ - } - } - if resp.AgentTask.TotalSteps <= 0 { - resp.AgentTask.TotalSteps = len(resp.AgentSteps) - } - if resp.AgentTask.ToolCalls <= 0 { - resp.AgentTask.ToolCalls = toolCalls - } - } - return resp -} - -// GetTrace 根据追踪 ID 查询追踪详情 -func (s *chatService) GetTrace(ctx context.Context, userID, traceID string, isAdmin bool) (*TraceResponse, error) { - if s.obsRepo == nil || traceID == "" { - return nil, fmt.Errorf("trace 存储未启用") - } - t, err := s.obsRepo.FindByID(ctx, traceID) - if err != nil { - return nil, fmt.Errorf("trace 不存在: %w", err) - } - if !isAdmin && t.UserID != userID { - return nil, fmt.Errorf("无权限访问该 trace") - } - resp := s.buildTraceResponse(t, true) - return &resp, nil -} - -// ListSessionTraces 分页查询会话维度的追踪列表 -func (s *chatService) ListSessionTraces(ctx context.Context, userID, sessionID string, isAdmin bool, offset, limit int) (TraceListResponse, error) { - if limit <= 0 { - limit = 20 - } - if limit > 200 { - limit = 200 - } - if offset < 0 { - offset = 0 - } - if s.obsRepo == nil { - return TraceListResponse{Total: 0, Traces: []any{}}, nil - } - if !isAdmin { - if err := s.validateSession(ctx, userID, sessionID); err != nil { - return TraceListResponse{}, err - } - } - list, total, err := s.obsRepo.ListBySession(ctx, sessionID, userID, offset, limit) - if isAdmin { - list, total, err = s.obsRepo.ListAll(ctx, sessionID, "", offset, limit) - } - if err != nil { - return TraceListResponse{}, err - } - items := make([]any, 0, len(list)) - for i := range list { - items = append(items, s.buildTraceResponse(&list[i], false)) - } - return TraceListResponse{Total: total, Traces: items}, nil -} - -// AdminListTraces 管理员分页查询全量追踪列表 -func (s *chatService) AdminListTraces(ctx context.Context, sessionID string, rating int, status string, offset, limit int) (TraceListResponse, error) { - if limit <= 0 { - limit = 50 - } - if limit > 500 { - limit = 500 - } - if offset < 0 { - offset = 0 - } - if s.obsRepo == nil { - return TraceListResponse{Total: 0, Traces: []any{}}, nil - } - list, total, err := s.obsRepo.ListAll(ctx, sessionID, status, offset, limit) - if err != nil { - return TraceListResponse{}, err - } - items := make([]any, 0, len(list)) - for i := range list { - items = append(items, s.buildTraceResponse(&list[i], false)) - } - return TraceListResponse{Total: total, Traces: items}, nil -} - -// GetMetricsSnapshot 获取可观测性指标快照 -func (s *chatService) GetMetricsSnapshot() (map[string]any, error) { - if s.obs == nil { - return nil, fmt.Errorf("observability 未启用") - } - raw, err := s.obs.MetricsSnapshot() - if err != nil { - return nil, err - } - rawCounters, _ := raw["counters"].([]any) - rawGauges, _ := raw["gauges"].([]any) - rawHistos, _ := raw["histograms"].([]any) - labelDropped, _ := raw["label_cardinality_dropped_total"].(int64) - var generatedTs string - if ts, ok := raw["generated_at_seconds"].(int64); ok { - generatedTs = time.Unix(ts, 0).Format(time.RFC3339) - } - - samplingRate := 0.0 - labelCardLimit := 0 - bufferDropped := int64(0) - piiMasked := int64(0) - if ss, ok := raw["sink_stats"].(map[string]any); ok { - if v, ok := ss["dropped_records_total"].(int64); ok { - bufferDropped = v - } - } - if c, ok := s.obs.(interface{ SamplingRate() float64 }); ok { - samplingRate = c.SamplingRate() - } else if cfg := s.cfgObservability(); cfg != nil { - samplingRate = cfg.SamplingRate - } - - labelsToMap := func(raw []any) map[string]string { - out := map[string]string{} - for _, r := range raw { - if m, ok := r.(map[string]any); ok { - k, _ := m["name"].(string) - v, _ := m["value"].(string) - if k != "" { - out[k] = v - } - } - } - return out - } - type namedSamples struct { - Name string - Help string - Samples []any - } - groupByMetric := func(rows []any) []namedSamples { - groups := map[string]*namedSamples{} - order := []string{} - for _, r := range rows { - m, ok := r.(map[string]any) - if !ok { - continue - } - name, _ := m["name"].(string) - if name == "" { - continue - } - if _, seen := groups[name]; !seen { - groups[name] = &namedSamples{Name: name} - order = append(order, name) - } - sample := map[string]any{} - labelsRaw, _ := m["labels"].([]any) - labelsM := labelsToMap(labelsRaw) - if len(labelsM) > 0 { - sample["labels"] = labelsM - } - switch { - case m["value"] != nil: - if v, ok := m["value"].(float64); ok { - sample["value"] = int64(v) - } else { - sample["value"] = m["value"] - } - case m["count"] != nil: - if v, ok := m["count"].(int64); ok { - sample["count"] = v - } else { - sample["count"] = m["count"] - } - if sum, ok := m["sum"].(float64); ok { - sample["sum"] = sum - } - if buckets, ok := m["buckets"].([]any); ok { - outB := make([]any, 0, len(buckets)) - for _, b := range buckets { - bm, ok := b.(map[string]any) - if !ok { - continue - } - le := bm["le"] - // JS Number.MAX_SAFE_INTEGER = 9007199254740991,+Inf 语义替换为该值,保证前端 TS number 类型一致 - if s, _ := le.(string); s == "+Inf" { - le = float64(9007199254740991) - } else if le == "+inf" || le == "Inf" { - le = float64(9007199254740991) - } - cnt, _ := bm["delta_count"].(int64) - outB = append(outB, map[string]any{"le": le, "count": cnt}) - } - sample["buckets"] = outB - } - } - groups[name].Samples = append(groups[name].Samples, sample) - } - out := make([]namedSamples, 0, len(order)) - for _, n := range order { - out = append(out, *groups[n]) - } - return out - } - cGroups := groupByMetric(rawCounters) - gGroups := groupByMetric(rawGauges) - hGroups := groupByMetric(rawHistos) - counters := make([]any, 0, len(cGroups)) - for _, c := range cGroups { - counters = append(counters, map[string]any{"name": c.Name, "help": "", "samples": c.Samples}) - } - gauges := make([]any, 0, len(gGroups)) - for _, g := range gGroups { - gauges = append(gauges, map[string]any{"name": g.Name, "help": "", "samples": g.Samples}) - } - histos := make([]any, 0, len(hGroups)) - for _, h := range hGroups { - histos = append(histos, map[string]any{"name": h.Name, "help": "", "samples": h.Samples}) - } - return map[string]any{ - "ts": generatedTs, - "counters": counters, - "gauges": gauges, - "histograms": histos, - "sampling_rate": samplingRate, - "label_cardinality_limit": labelCardLimit, - "buffer_dropped_total": bufferDropped, - "pii_masked_total": piiMasked, - "label_cardinality_dropped_total": labelDropped, - }, nil -} - -func (s *chatService) cfgObservability() *config.ObservabilityConfig { - if s.obs == nil { - return nil - } - type cfgProvider interface{ Config() config.ObservabilityConfig } - if p, ok := s.obs.(cfgProvider); ok { - c := p.Config() - return &c - } - return nil -} - -func ratingLabel(r int) string { - switch r { - case 1: - return "up" - case -1: - return "down" - default: - return "unknown" - } -} - -func boolLabel(b bool) string { - if b { - return "true" - } - return "false" -} - -func reasonTagOrDefault(tag string) string { - if tag == "" { - return "none" - } - return tag -} - -func metadataAsMap(raw datatypes.JSON) map[string]any { - if len(raw) == 0 { - return nil - } - var m map[string]any - if err := json.Unmarshal(raw, &m); err != nil { - return nil - } - return m -} - -// extractSearchMode 从 Attrs JSON / SpanTree JSON 里取 search_mode 作为 TraceResponse 顶层字段 -// 兼容 datatypes.JSON(GORM)、map[string]any(内存对象)、[]byte 三种来源; -// 找不到时再回退硬解析 ChatTrace.SpanTree root 的 attrs.search_mode,避免新老数据过渡时为空 -func extractSearchMode(attrs any, spanTreeHint ...datatypes.JSON) string { - if s := extractSearchModeFromAny(attrs); s != "" { - return s - } - for _, st := range spanTreeHint { - if len(st) == 0 { - continue - } - var root struct { - Attrs map[string]any `json:"attrs"` - } - if err := json.Unmarshal(st, &root); err == nil { - if s, ok := root.Attrs["search_mode"].(string); ok && s != "" { - return s - } - } - } - return "" -} - -func extractSearchModeFromAny(attrs any) string { - if attrs == nil { - return "" - } - switch v := attrs.(type) { - case map[string]any: - if s, ok := v["search_mode"].(string); ok { - return s - } - case datatypes.JSON: - if len(v) == 0 { - return "" - } - var m map[string]any - if err := json.Unmarshal(v, &m); err == nil { - if s, ok := m["search_mode"].(string); ok { - return s - } - } - case []byte: - if len(v) == 0 { - return "" - } - var m map[string]any - if err := json.Unmarshal(v, &m); err == nil { - if s, ok := m["search_mode"].(string); ok { - return s - } - } - } - return "" -} diff --git a/internal/service/chat_service_graph_quick.go b/internal/service/chat_service_graph_quick.go index ce9d76e..93688e3 100644 --- a/internal/service/chat_service_graph_quick.go +++ b/internal/service/chat_service_graph_quick.go @@ -58,19 +58,15 @@ type quickGraphInput struct { // quickGraphState Graph Local State,通过 ProcessState 读写。 type quickGraphState struct { - Input *quickGraphInput - RewrittenQuery string // 改写后的查询,供 Retrieve / BuildMsgs 使用 - Intent string // greeting / chitchat / question / identity / meta - SkipRetrieve bool // Greeting/Chitchat 跳过知识库检索 - NeedClarify bool // 意图不明确,需要用户澄清 - ClarifyQuestion string // 追问文本 + Input *quickGraphInput + 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 - - // RewriteDone 后台 goroutine 完成 LLM 改写后关闭,Retrieve 阶段可选等待 - RewriteDone chan struct{} - RewriteErr error + Keywords []string // 改写时提取的关键词,可用于日志/调试 + RetrievedDocs []*schema.Document } // 查询改写意图类型 @@ -180,89 +176,71 @@ func addQuickRewriteNode(g *einoCompose.Graph[*quickGraphInput, *schema.StreamRe ) } -// quickRewriteFn 节点 1 实现:触发查询改写 + 意图识别,立即返回原始 query 让 Retrieve 先行。 -// 并行策略:LLM 改写放到后台 goroutine 异步执行,Retrieve 用 OriginalQuery 先跑, -// 改写结果通过 State.RewriteDone channel 通知下游,Retrieve StatePostHandler 可选等待并补检索。 -// 短路优化:如果 graphInput.PreRewrittenQuery 已填(Graph 执行前已算过),直接同步执行不再起 goroutine。 +// quickRewriteFn 节点 1 实现:查询改写 + 意图识别。 +// Graph 启动前 processMessageGraphQuick 已经预执行过 doRewriteWithLLM, +// 所以正常路径走 PreRewrittenQuery 短路,直接把预计算结果写进 state。 +// 当 PreRewrittenQuery 为空时(Graph 被独立调用的防御性路径),同步调 LLM 改写。 func quickRewriteFn(ctx context.Context, input *quickGraphInput) (string, error) { if input == nil { return "", apperrors.NewDefault(apperrors.CodeInvalidParam) } - // 初始化 State:存 Input + 创建 RewriteDone channel if err := einoCompose.ProcessState(ctx, func(_ context.Context, state *quickGraphState) error { state.Input = input - state.RewriteDone = make(chan struct{}) return nil }); err != nil { return "", err } - // 短路:Graph 执行前已算好 → 同步写 state,不走并行 + var rewritten, intent string + var keywords []string + var skipRetrieve, needClarify bool + var clarifyQ string + var clarifyO []string + if input.PreRewrittenQuery != "" { - rewritten, intent, keywords := input.PreRewrittenQuery, input.PreIntent, input.PreKeywords - skipRetrieve, needClarify := input.PreSkipRetrieve, input.PreNeedClarify - clarifyQ, clarifyO := input.PreClarifyQuestion, input.PreClarifyOptions - _ = 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 = clarifyQ - state.ClarifyOptions = clarifyO - close(state.RewriteDone) - return nil - }) - observability.SetSpanAttrs(ctx, observability.Attrs{ - "original_query": input.OriginalQuery, - "rewritten_query": rewritten, - "rewrite_mode": "precomputed", - }) - return rewritten, nil - } - - // 正常路径:后台 goroutine 异步跑 LLM 改写,立即返回 OriginalQuery 让 Retrieve 并行 - go func() { - startAt := time.Now() - rewritten, intent, keywords, skipRetrieve, needClarify, clarifyQ, clarifyO := 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 = clarifyQ - state.ClarifyOptions = clarifyO - close(state.RewriteDone) - return nil - }) - observability.SetSpanAttrs(ctx, observability.Attrs{ - "rewrite_ms": time.Since(startAt).Milliseconds(), - "rewrite_mode": "async", - "original_query": input.OriginalQuery, - "rewritten_query": rewritten, - "intent": intent, - "skip_retrieve": fmt.Sprintf("%v", skipRetrieve), - "need_clarify": fmt.Sprintf("%v", needClarify), - }) - }() + rewritten = input.PreRewrittenQuery + intent = input.PreIntent + keywords = input.PreKeywords + skipRetrieve = input.PreSkipRetrieve + needClarify = input.PreNeedClarify + clarifyQ = input.PreClarifyQuestion + clarifyO = input.PreClarifyOptions + } else { + rewritten, intent, keywords, skipRetrieve, needClarify, clarifyQ, clarifyO = 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 = clarifyQ + state.ClarifyOptions = clarifyO + return nil + }) observability.SetSpanAttrs(ctx, observability.Attrs{ - "original_query": input.OriginalQuery, - "rewrite_mode": "async_started", + "original_query": input.OriginalQuery, + "rewritten_query": rewritten, + "intent": intent, + "skip_retrieve": fmt.Sprintf("%v", skipRetrieve), + "need_clarify": fmt.Sprintf("%v", needClarify), }) - return input.OriginalQuery, nil + + return rewritten, nil } // matchLocalIntent 本地快速意图匹配(纯正则 + 关键词,0ms)。 // 返回 (intent, matched) —— matched=false 表示交给 LLM 判定。 // // 覆盖四类场景: -// greeting: 你好 / hi / 早上好 / 在吗 -// identity: 你是谁 / 你能做什么 / 介绍一下自己 -// chitchat: 今天星期几 / 讲个笑话 / 随便聊聊(含"今天/现在+时间查询") -// meta: 我的历史 / 刚才说了什么 +// +// greeting: 你好 / hi / 早上好 / 在吗 +// identity: 你是谁 / 你能做什么 / 介绍一下自己 +// chitchat: 今天星期几 / 讲个笑话 / 随便聊聊(含"今天/现在+时间查询") +// meta: 我的历史 / 刚才说了什么 // // 不命中时返回 ("", false),交给 LLM 做更精细的意图判定。 func matchLocalIntent(raw string) (string, bool) { @@ -303,11 +281,11 @@ func matchLocalIntent(raw string) (string, bool) { // matchRegex 简单的正则匹配封装,避免每次都 re.Compile var ( - reGreeting = regexp.MustCompile(`^(你好|您好|hi+|hello+|嗨|哈喽|在吗|在不在|早|早上好|下午好|晚上好|晚安|早安|午安|晚安)$`) - reIdentity = regexp.MustCompile(`^(你是谁|你是谁呀|你叫什么|你叫什么名字|你能做什么|你能干什么|你是干什么的|介绍一下你自己|自我介绍|你是什么模型|你是什么)$`) - reTimeInfo = regexp.MustCompile(`(今天|现在|当前|明天|后天)+(星期几|礼拜几|几号|多少号|日期|几号了|几点|几点钟|时间|日期是)`) - reChitchat = regexp.MustCompile(`^(讲个笑话|来个笑话|随便聊聊|聊聊呗|聊聊天|说点什么|有什么好玩的|今天天气怎么样|天气怎么样|心情不好|我心情不好|安慰一下我|夸夸我)$`) - reMeta = regexp.MustCompile(`(我的历史|聊天记录|你刚才说了什么|刚才说的什么|上一个问题|前一个问题|回顾对话|我们聊了什么|你还记得|之前说的)`) + reGreeting = regexp.MustCompile(`^(你好|您好|hi+|hello+|嗨|哈喽|在吗|在不在|早|早上好|下午好|晚上好|晚安|早安|午安|晚安)$`) + reIdentity = regexp.MustCompile(`^(你是谁|你是谁呀|你叫什么|你叫什么名字|你能做什么|你能干什么|你是干什么的|介绍一下你自己|自我介绍|你是什么模型|你是什么)$`) + reTimeInfo = regexp.MustCompile(`(今天|现在|当前|明天|后天)+(星期几|礼拜几|几号|多少号|日期|几号了|几点|几点钟|时间|日期是)`) + reChitchat = regexp.MustCompile(`^(讲个笑话|来个笑话|随便聊聊|聊聊呗|聊聊天|说点什么|有什么好玩的|今天天气怎么样|天气怎么样|心情不好|我心情不好|安慰一下我|夸夸我)$`) + reMeta = regexp.MustCompile(`(我的历史|聊天记录|你刚才说了什么|刚才说的什么|上一个问题|前一个问题|回顾对话|我们聊了什么|你还记得|之前说的)`) ) func matchRegex(pattern string, q string) bool { @@ -455,14 +433,10 @@ func buildRewriteHistory(msgs []*schema.Message, currentUserMsgIdx, maxRounds in return strings.Join(pairs, "\n") } -// rewriteParallelMaxWait Retrieve 阶段等待后台 LLM 改写的最长时间。 -// 原始 query 检索通常 0.3-0.8s,LLM 改写通常 0.5-2s——等 500ms 给快模型一个机会, -// 超时就放弃,不阻塞主路径。 -const rewriteParallelMaxWait = 500 * time.Millisecond - // addQuickRetrieveNode 节点 2:Retrieve -// 改进:用 LambdaNode 替代 AddRetrieverNode,在 Lambda 内部提前检查 SkipRetrieve / NeedClarify, +// 用 LambdaNode 替代 AddRetrieverNode,在 Lambda 内部提前检查 SkipRetrieve / NeedClarify, // 避免 EinoRetrieverAdapter 被实例化后才被 PostHandler 清空——那样知识库查询的开销已经花出去了。 +// QueryRewrite 已经同步完成,Retrieve 直接用改写后的 query(或原始 query)查一次即可,不再做并行改写等待。 func addQuickRetrieveNode(g *einoCompose.Graph[*quickGraphInput, *schema.StreamReader[*schema.Message]], einoRetriever *rag.EinoRetrieverAdapter) error { return g.AddLambdaNode(graphQuickNodeRetrieve, einoCompose.InvokableLambda(func(ctx context.Context, query string) ([]*schema.Document, error) { @@ -483,7 +457,6 @@ func addQuickRetrieveNode(g *einoCompose.Graph[*quickGraphInput, *schema.StreamR // 构造 retriever.Option(KBIDs / UserID / TopK) opts := buildRetrieverOpts(state.Input) - // 用当前 query(Rewrite 返回的原始或改写后 query)先查 docs, err := einoRetriever.Retrieve(ctx, query, opts...) if err != nil { logger.Warnf("quickRetrieveFn: 检索失败,降级为空结果: %v", err) @@ -491,23 +464,6 @@ func addQuickRetrieveNode(g *einoCompose.Graph[*quickGraphInput, *schema.StreamR return nil, nil } state.RetrievedDocs = docs - - // 等后台 LLM 改写(最多 rewriteParallelMaxWait),改写完成后补一次检索 + 合并去重 - if state.RewriteDone != nil { - select { - case <-state.RewriteDone: - if state.RewrittenQuery != "" && state.RewrittenQuery != state.Input.OriginalQuery { - if !state.SkipRetrieve && !state.NeedClarify { - rewrittenDocs, rErr := einoRetriever.Retrieve(ctx, state.RewrittenQuery, opts...) - if rErr == nil && len(rewrittenDocs) > 0 { - docs = mergeDocsByScore(docs, rewrittenDocs) - state.RetrievedDocs = docs - } - } - } - case <-time.After(rewriteParallelMaxWait): - } - } return docs, nil }), einoCompose.WithNodeName("KnowledgeRetrieve"), @@ -532,37 +488,6 @@ func buildRetrieverOpts(input *quickGraphInput) []retriever.Option { return opts } -// mergeDocsByScore 合并两组 docs,按 ID 去重取最高分,最后按分数降序。 -func mergeDocsByScore(a, b []*schema.Document) []*schema.Document { - scoreMap := make(map[string]*schema.Document, len(a)+len(b)) - for _, d := range a { - if d == nil { - continue - } - scoreMap[d.ID] = d - } - for _, d := range b { - if d == nil { - continue - } - if existing, ok := scoreMap[d.ID]; ok { - if d.Score() > existing.Score() { - scoreMap[d.ID] = d - } - } else { - scoreMap[d.ID] = d - } - } - merged := make([]*schema.Document, 0, len(scoreMap)) - for _, d := range scoreMap { - merged = append(merged, d) - } - sort.Slice(merged, func(i, j int) bool { - return merged[i].Score() > merged[j].Score() - }) - return merged -} - // addQuickBuildMsgsNode 节点 3:BuildPromptMessages。 // 从 State 拿 Input,在 userQuestionIndex 前插入检索上下文。 func addQuickBuildMsgsNode(g *einoCompose.Graph[*quickGraphInput, *schema.StreamReader[*schema.Message]]) error { @@ -571,6 +496,7 @@ func addQuickBuildMsgsNode(g *einoCompose.Graph[*quickGraphInput, *schema.Stream einoCompose.WithNodeName("BuildPromptMessages"), ) } + // quickBuildMsgsFn 节点 3 实现:在用户问题前插入检索上下文块 func quickBuildMsgsFn(ctx context.Context, docs []*schema.Document) ([]*schema.Message, error) { @@ -857,6 +783,7 @@ var graphCtxChatModelKey = graphCtxChatModelKeyType{} func withGraphChatModel(ctx context.Context, cm einoModel.BaseChatModel) context.Context { return context.WithValue(ctx, graphCtxChatModelKey, cm) } + // graphChatModelFromContext 从 context 取出 ChatModel func graphChatModelFromContext(ctx context.Context) (einoModel.BaseChatModel, bool) { @@ -940,9 +867,9 @@ func (s *chatService) processMessageGraphQuick( }, Done: true} if obsOk { s.obs.EndSpan(ctx, span, observability.SpanStatusOK, nil, observability.Attrs{ - "need_clarify": "true", + "need_clarify": "true", "clarify_intent": intent, - "clarify_ms": fmt.Sprintf("%d", time.Since(obsNow).Milliseconds()), + "clarify_ms": fmt.Sprintf("%d", time.Since(obsNow).Milliseconds()), }) } return @@ -955,8 +882,7 @@ func (s *chatService) processMessageGraphQuick( graphInput.PreSkipRetrieve = skipRetrieve // 4) 提前创建 graphState:既是 eino stateGenerator 返回值,也是 Invoke 后外部读取 RetrievedDocs 的入口。 - // RewriteDone channel 在这里初始化——rewriteFn 后台 goroutine 完成后 close 它,retrieve StatePostHandler 等它。 - graphState := &quickGraphState{RewriteDone: make(chan struct{})} + graphState := &quickGraphState{} // 5) 构建并编译 compose.Graph(内部已经 push error 事件) sendProgressEvent(eventCh, "正在组装快速检索链路...") diff --git a/internal/service/chat_service_mode.go b/internal/service/chat_service_mode.go index 8233a28..4d7c3cc 100644 --- a/internal/service/chat_service_mode.go +++ b/internal/service/chat_service_mode.go @@ -20,23 +20,15 @@ import ( "solvify-agent/pkg/logger" ) -// ─── 快速检索模式 ─────────────────────────────────────────── - -// processMessage 处理消息的核心流程(快速检索模式) -// 通过 compose.Graph 显式编排:QueryRewrite → Retrieve → BuildPromptMessages → Generate -// 四节点各自独立 Span + Metrics(eino 全局 callback 自动打点)。 -func (s *chatService) processMessage(ctx context.Context, userID, sessionID, userMsgID string, req requestdto.SendMessageRequest, eventCh chan<- dto.StreamEvent) { - s.processMessageGraphQuick(ctx, userID, sessionID, userMsgID, req, eventCh) -} - // ─── 深度思考模式 ─────────────────────────────────────────── // processDeepMode 深度思考模式处理流程 // 使用 eino ReAct Agent,自动管理 Think → Act → Observe 循环 func (s *chatService) processDeepMode(ctx context.Context, userID, sessionID, userMsgID string, req requestdto.SendMessageRequest, eventCh chan<- dto.StreamEvent) { obsOk := s.obs != nil + var span *observability.Span if obsOk { - _, span := s.obs.StartSpan(ctx, "chat.deep", observability.ComponentAgentEngine, observability.Attrs{ + ctx, span = s.obs.StartSpan(ctx, "chat.deep", observability.ComponentAgentEngine, observability.Attrs{ "session_id": sessionID, "user_id": userID, "model_id": req.ModelID, @@ -85,7 +77,6 @@ func (s *chatService) processDeepMode(ctx context.Context, userID, sessionID, us } } } - _ = deepCtx client, enhancedCtx, err := s.initContext(ctx, userID, sessionID, req.ModelID, req.ModelType, req.Content, preToolsTokens) if err != nil { @@ -101,10 +92,10 @@ func (s *chatService) processDeepMode(ctx context.Context, userID, sessionID, us modelName := client.ModelName() if obsOk { s.obs.AddRootAttrs(ctx, observability.Attrs{ - "model_name": modelName, - "tools_tokens": fmt.Sprintf("%d", preToolsTokens), - "history_budget": fmt.Sprintf("%d", enhancedCtx.HistoryBudget), - "retrieval_budget": fmt.Sprintf("%d", enhancedCtx.RetrievalBudget), + "model_name": modelName, + "tools_tokens": fmt.Sprintf("%d", preToolsTokens), + "history_budget": fmt.Sprintf("%d", enhancedCtx.HistoryBudget), + "retrieval_budget": fmt.Sprintf("%d", enhancedCtx.RetrievalBudget), }) s.obs.Observe(ctx, "chat_deep_init_ctx_seconds", map[string]string{"model_id": req.ModelID}, time.Since(t0).Seconds()) } @@ -123,19 +114,19 @@ func (s *chatService) processDeepMode(ctx context.Context, userID, sessionID, us 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) - _ = s.sessionRepo.ClearPendingClarify(ctx, sessionID) + 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) + _ = s.sessionRepo.ClearPendingClarify(ctx, sessionID) } + } } t1 := time.Now() agentEventCh, err := s.agentEngine.Execute(deepCtx, agentReq, chatModel) @@ -170,40 +161,40 @@ func (s *chatService) processDeepMode(ctx context.Context, userID, sessionID, us } // ── interrupt 事件:存 checkpoint 到 session,返回中断事件给前端 ── - if agentEvent.Type == agent.EventInterrupt { - logger.Infof("[ChatService] Agent 中断: checkpointID=%s, interruptID=%s, isClarify=%v", agentEvent.CheckpointID, agentEvent.InterruptID, agentEvent.IsClarify) - if agentEvent.CheckpointID != "" { - pcData := &entity.PendingCheckpointData{ - CheckpointID: agentEvent.CheckpointID, - InterruptID: agentEvent.InterruptID, - Question: agentEvent.Detail, - IsClarify: agentEvent.IsClarify, - Options: agentEvent.ClarifyOptions, - 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) - } + if agentEvent.Type == agent.EventInterrupt { + logger.Infof("[ChatService] Agent 中断: checkpointID=%s, interruptID=%s, isClarify=%v", agentEvent.CheckpointID, agentEvent.InterruptID, agentEvent.IsClarify) + if agentEvent.CheckpointID != "" { + pcData := &entity.PendingCheckpointData{ + CheckpointID: agentEvent.CheckpointID, + InterruptID: agentEvent.InterruptID, + Question: agentEvent.Detail, + IsClarify: agentEvent.IsClarify, + Options: agentEvent.ClarifyOptions, + 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) + } - // clarify 中断额外存到 pending_clarify(兼容旧恢复流程) - if agentEvent.IsClarify { - clarifyData := &entity.PendingClarifyData{ - Question: agentEvent.ClarifyQuestion, - Options: agentEvent.ClarifyOptions, - SetAt: time.Now(), - } - clarifyRaw, _ := json.Marshal(clarifyData) - if sErr := s.sessionRepo.SetPendingClarify(ctx, sessionID, clarifyRaw); sErr != nil { - logger.Errorf("存储 pending clarify 失败: %v", sErr) - } + // clarify 中断额外存到 pending_clarify(兼容旧恢复流程) + if agentEvent.IsClarify { + clarifyData := &entity.PendingClarifyData{ + Question: agentEvent.ClarifyQuestion, + Options: agentEvent.ClarifyOptions, + SetAt: time.Now(), + } + clarifyRaw, _ := json.Marshal(clarifyData) + if sErr := s.sessionRepo.SetPendingClarify(ctx, sessionID, clarifyRaw); sErr != nil { + logger.Errorf("存储 pending clarify 失败: %v", sErr) } } - eventCh <- toStreamEvent(agentEvent) - return } + eventCh <- toStreamEvent(agentEvent) + return + } if agentEvent.Type == agent.EventAnswer { fullContent += agentEvent.Content @@ -247,12 +238,12 @@ func (s *chatService) processDeepMode(ctx context.Context, userID, sessionID, us "tool_calls": fmt.Sprintf("%d", toolCallsN), }, 1) s.obs.AddRootAttrs(ctx, observability.Attrs{ - "tool_calls": toolCallsN, - "tool_errors": toolErrorsN, - "steps_n": len(reasoningSteps), - "rag_docs_n": len(agentSources), - "agent_error": agentErrorSeen, - "tool_used": toolEventSeen, + "tool_calls": toolCallsN, + "tool_errors": toolErrorsN, + "steps_n": len(reasoningSteps), + "rag_docs_n": len(agentSources), + "agent_error": agentErrorSeen, + "tool_used": toolEventSeen, "assistant_chars": len([]rune(fullContent)), }) } diff --git a/internal/service/feedback_service.go b/internal/service/feedback_service.go new file mode 100644 index 0000000..e6381db --- /dev/null +++ b/internal/service/feedback_service.go @@ -0,0 +1,115 @@ +package service + +import ( + "context" + "fmt" + + "github.com/google/uuid" + + "solvify-agent/internal/model/entity" + "solvify-agent/internal/observability" +) + +// ─── 可观测:反馈 / Trace / Metrics 查询接口 ─────────────────────────────────── + +// SubmitFeedback 提交消息反馈 +func (s *chatService) SubmitFeedback(ctx context.Context, userID, messageID string, req FeedbackRequest) error { + if req.Rating != 1 && req.Rating != -1 { + return fmt.Errorf("rating 必须为 1 或 -1") + } + if messageID == "" || userID == "" { + return fmt.Errorf("message_id / user_id 不能为空") + } + msg, err := s.messageRepo.FindByID(ctx, messageID) + if err != nil { + return fmt.Errorf("查询消息失败: %w", err) + } + if msg == nil { + return fmt.Errorf("消息不存在或无权限") + } + if msg.SessionID != "" { + if vErr := s.validateSession(ctx, userID, msg.SessionID); vErr != nil { + return fmt.Errorf("消息不存在或无权限") + } + } + var traceID string + if raw := msg.Metadata; len(raw) > 0 { + if meta := metadataAsMap(raw); meta != nil { + if v, ok := meta["trace_id"].(string); ok { + traceID = v + } + } + } + primaryTag := "" + if len(req.Reasons) > 0 { + primaryTag = req.Reasons[0] + } + fb := &entity.MessageFeedback{ + ID: uuid.New().String(), + MessageID: messageID, + UserID: userID, + SessionID: msg.SessionID, + Rating: req.Rating, + ReasonTag: primaryTag, + Comment: req.Comment, + IsQuick: req.IsQuick, + TraceID: traceID, + } + fb.SetReasons(req.Reasons) + if s.obsRepo != nil { + if e := s.obsRepo.CreateFeedback(ctx, fb); e != nil { + return fmt.Errorf("保存反馈失败: %w", e) + } + } + if s.obs != nil { + s.obs.Incr(ctx, "chat_feedback_total", map[string]string{ + "rating": ratingLabel(req.Rating), + "reason_tag": reasonTagOrDefault(primaryTag), + "has_comment": boolLabel(req.Comment != ""), + }, 1) + } + if s.obs != nil { + s.obs.RecordFeedback(&observability.Feedback{ + MessageID: fb.MessageID, + UserID: fb.UserID, + SessionID: fb.SessionID, + Rating: fb.Rating, + Reasons: fb.Reasons(), + Comment: fb.Comment, + TraceID: fb.TraceID, + CreatedAt: fb.CreatedAt, + }) + } + return nil +} + +// ListFeedbacks 分页查询用户反馈列表 +func (s *chatService) ListFeedbacks(ctx context.Context, userID string, offset, limit int) (FeedbackListResponse, error) { + if limit <= 0 { + limit = 20 + } + if limit > 100 { + limit = 100 + } + if offset < 0 { + offset = 0 + } + if s.obsRepo == nil { + return FeedbackListResponse{Total: 0, Feedbacks: []any{}}, nil + } + list, total, err := s.obsRepo.ListByUser(ctx, userID, offset, limit) + if err != nil { + return FeedbackListResponse{}, err + } + type out struct { + entity.MessageFeedback + Reasons []string `json:"reasons"` + } + items := make([]any, 0, len(list)) + for _, f := range list { + items = append(items, out{MessageFeedback: f, Reasons: f.Reasons()}) + } + return FeedbackListResponse{Total: total, Feedbacks: items}, nil +} + +// buildTraceResponse 构建追踪详情响应 diff --git a/internal/service/trace_service.go b/internal/service/trace_service.go new file mode 100644 index 0000000..cc46198 --- /dev/null +++ b/internal/service/trace_service.go @@ -0,0 +1,381 @@ +package service + +import ( + "context" + "encoding/json" + "fmt" + "time" + + "solvify-agent/internal/model/entity" + "solvify-agent/pkg/config" + + "gorm.io/datatypes" +) + +func (s *chatService) buildTraceResponse(t *entity.ChatTrace, includeAgentDetail bool) TraceResponse { + if t == nil { + return TraceResponse{} + } + resp := TraceResponse{ + ID: t.ID, + RequestID: t.RequestID, + UserID: t.UserID, + SessionID: t.SessionID, + SearchMode: extractSearchMode(t.Attrs), + SampleRate: t.SampleRate, + Sampled: t.Sampled, + DurationMs: t.DurationMs, + Status: t.Status, + Error: t.Error, + Attrs: t.Attrs, + SpanTree: t.SpanTree, + CreatedAt: t.CreatedAt.Format("2006-01-02 15:04:05"), + } + if !includeAgentDetail || s.obsRepo == nil { + return resp + } + task, steps, _ := s.obsRepo.FindByTraceID(context.Background(), t.ID) + resp.AgentTask = chatAgentTaskEntityToResponse(task) + resp.AgentSteps = chatAgentStepEntityToResponse(steps) + // 有 AgentStep 信息时,把 TotalSteps / ToolCalls 反填回 AgentTask(如果之前 MarkEnded 没填充好) + if resp.AgentTask != nil && len(resp.AgentSteps) > 0 { + toolCalls := 0 + for _, st := range resp.AgentSteps { + if st.ToolName != "" && st.ToolName != "llm.reasoning" { + toolCalls++ + } + } + if resp.AgentTask.TotalSteps <= 0 { + resp.AgentTask.TotalSteps = len(resp.AgentSteps) + } + if resp.AgentTask.ToolCalls <= 0 { + resp.AgentTask.ToolCalls = toolCalls + } + } + return resp +} + +// GetTrace 根据追踪 ID 查询追踪详情 +func (s *chatService) GetTrace(ctx context.Context, userID, traceID string, isAdmin bool) (*TraceResponse, error) { + if s.obsRepo == nil || traceID == "" { + return nil, fmt.Errorf("trace 存储未启用") + } + t, err := s.obsRepo.FindByID(ctx, traceID) + if err != nil { + return nil, fmt.Errorf("trace 不存在: %w", err) + } + if !isAdmin && t.UserID != userID { + return nil, fmt.Errorf("无权限访问该 trace") + } + resp := s.buildTraceResponse(t, true) + return &resp, nil +} + +// ListSessionTraces 分页查询会话维度的追踪列表 +func (s *chatService) ListSessionTraces(ctx context.Context, userID, sessionID string, isAdmin bool, offset, limit int) (TraceListResponse, error) { + if limit <= 0 { + limit = 20 + } + if limit > 200 { + limit = 200 + } + if offset < 0 { + offset = 0 + } + if s.obsRepo == nil { + return TraceListResponse{Total: 0, Traces: []any{}}, nil + } + if !isAdmin { + if err := s.validateSession(ctx, userID, sessionID); err != nil { + return TraceListResponse{}, err + } + } + list, total, err := s.obsRepo.ListBySession(ctx, sessionID, userID, offset, limit) + if isAdmin { + list, total, err = s.obsRepo.ListAll(ctx, sessionID, "", offset, limit) + } + if err != nil { + return TraceListResponse{}, err + } + items := make([]any, 0, len(list)) + for i := range list { + items = append(items, s.buildTraceResponse(&list[i], false)) + } + return TraceListResponse{Total: total, Traces: items}, nil +} + +// AdminListTraces 管理员分页查询全量追踪列表 +func (s *chatService) AdminListTraces(ctx context.Context, sessionID string, rating int, status string, offset, limit int) (TraceListResponse, error) { + if limit <= 0 { + limit = 50 + } + if limit > 500 { + limit = 500 + } + if offset < 0 { + offset = 0 + } + if s.obsRepo == nil { + return TraceListResponse{Total: 0, Traces: []any{}}, nil + } + list, total, err := s.obsRepo.ListAll(ctx, sessionID, status, offset, limit) + if err != nil { + return TraceListResponse{}, err + } + items := make([]any, 0, len(list)) + for i := range list { + items = append(items, s.buildTraceResponse(&list[i], false)) + } + return TraceListResponse{Total: total, Traces: items}, nil +} + +// GetMetricsSnapshot 获取可观测性指标快照 +func (s *chatService) GetMetricsSnapshot() (map[string]any, error) { + if s.obs == nil { + return nil, fmt.Errorf("observability 未启用") + } + raw, err := s.obs.MetricsSnapshot() + if err != nil { + return nil, err + } + rawCounters, _ := raw["counters"].([]any) + rawGauges, _ := raw["gauges"].([]any) + rawHistos, _ := raw["histograms"].([]any) + labelDropped, _ := raw["label_cardinality_dropped_total"].(int64) + var generatedTs string + if ts, ok := raw["generated_at_seconds"].(int64); ok { + generatedTs = time.Unix(ts, 0).Format(time.RFC3339) + } + + samplingRate := 0.0 + labelCardLimit := 0 + bufferDropped := int64(0) + piiMasked := int64(0) + if ss, ok := raw["sink_stats"].(map[string]any); ok { + if v, ok := ss["dropped_records_total"].(int64); ok { + bufferDropped = v + } + } + if c, ok := s.obs.(interface{ SamplingRate() float64 }); ok { + samplingRate = c.SamplingRate() + } else if cfg := s.cfgObservability(); cfg != nil { + samplingRate = cfg.SamplingRate + } + + labelsToMap := func(raw []any) map[string]string { + out := map[string]string{} + for _, r := range raw { + if m, ok := r.(map[string]any); ok { + k, _ := m["name"].(string) + v, _ := m["value"].(string) + if k != "" { + out[k] = v + } + } + } + return out + } + type namedSamples struct { + Name string + Help string + Samples []any + } + groupByMetric := func(rows []any) []namedSamples { + groups := map[string]*namedSamples{} + order := []string{} + for _, r := range rows { + m, ok := r.(map[string]any) + if !ok { + continue + } + name, _ := m["name"].(string) + if name == "" { + continue + } + if _, seen := groups[name]; !seen { + groups[name] = &namedSamples{Name: name} + order = append(order, name) + } + sample := map[string]any{} + labelsRaw, _ := m["labels"].([]any) + labelsM := labelsToMap(labelsRaw) + if len(labelsM) > 0 { + sample["labels"] = labelsM + } + switch { + case m["value"] != nil: + if v, ok := m["value"].(float64); ok { + sample["value"] = int64(v) + } else { + sample["value"] = m["value"] + } + case m["count"] != nil: + if v, ok := m["count"].(int64); ok { + sample["count"] = v + } else { + sample["count"] = m["count"] + } + if sum, ok := m["sum"].(float64); ok { + sample["sum"] = sum + } + if buckets, ok := m["buckets"].([]any); ok { + outB := make([]any, 0, len(buckets)) + for _, b := range buckets { + bm, ok := b.(map[string]any) + if !ok { + continue + } + le := bm["le"] + // JS Number.MAX_SAFE_INTEGER = 9007199254740991,+Inf 语义替换为该值,保证前端 TS number 类型一致 + if s, _ := le.(string); s == "+Inf" { + le = float64(9007199254740991) + } else if le == "+inf" || le == "Inf" { + le = float64(9007199254740991) + } + cnt, _ := bm["delta_count"].(int64) + outB = append(outB, map[string]any{"le": le, "count": cnt}) + } + sample["buckets"] = outB + } + } + groups[name].Samples = append(groups[name].Samples, sample) + } + out := make([]namedSamples, 0, len(order)) + for _, n := range order { + out = append(out, *groups[n]) + } + return out + } + cGroups := groupByMetric(rawCounters) + gGroups := groupByMetric(rawGauges) + hGroups := groupByMetric(rawHistos) + counters := make([]any, 0, len(cGroups)) + for _, c := range cGroups { + counters = append(counters, map[string]any{"name": c.Name, "help": "", "samples": c.Samples}) + } + gauges := make([]any, 0, len(gGroups)) + for _, g := range gGroups { + gauges = append(gauges, map[string]any{"name": g.Name, "help": "", "samples": g.Samples}) + } + histos := make([]any, 0, len(hGroups)) + for _, h := range hGroups { + histos = append(histos, map[string]any{"name": h.Name, "help": "", "samples": h.Samples}) + } + return map[string]any{ + "ts": generatedTs, + "counters": counters, + "gauges": gauges, + "histograms": histos, + "sampling_rate": samplingRate, + "label_cardinality_limit": labelCardLimit, + "buffer_dropped_total": bufferDropped, + "pii_masked_total": piiMasked, + "label_cardinality_dropped_total": labelDropped, + }, nil +} + +func (s *chatService) cfgObservability() *config.ObservabilityConfig { + if s.obs == nil { + return nil + } + type cfgProvider interface { + Config() config.ObservabilityConfig + } + if p, ok := s.obs.(cfgProvider); ok { + c := p.Config() + return &c + } + return nil +} + +func ratingLabel(r int) string { + switch r { + case 1: + return "up" + case -1: + return "down" + default: + return "unknown" + } +} + +func boolLabel(b bool) string { + if b { + return "true" + } + return "false" +} + +func reasonTagOrDefault(tag string) string { + if tag == "" { + return "none" + } + return tag +} + +func metadataAsMap(raw datatypes.JSON) map[string]any { + if len(raw) == 0 { + return nil + } + var m map[string]any + if err := json.Unmarshal(raw, &m); err != nil { + return nil + } + return m +} + +// extractSearchMode 从 Attrs JSON / SpanTree JSON 里取 search_mode 作为 TraceResponse 顶层字段 +// 兼容 datatypes.JSON(GORM)、map[string]any(内存对象)、[]byte 三种来源; +// 找不到时再回退硬解析 ChatTrace.SpanTree root 的 attrs.search_mode,避免新老数据过渡时为空 +func extractSearchMode(attrs any, spanTreeHint ...datatypes.JSON) string { + if s := extractSearchModeFromAny(attrs); s != "" { + return s + } + for _, st := range spanTreeHint { + if len(st) == 0 { + continue + } + var root struct { + Attrs map[string]any `json:"attrs"` + } + if err := json.Unmarshal(st, &root); err == nil { + if s, ok := root.Attrs["search_mode"].(string); ok && s != "" { + return s + } + } + } + return "" +} + +func extractSearchModeFromAny(attrs any) string { + if attrs == nil { + return "" + } + switch v := attrs.(type) { + case map[string]any: + if s, ok := v["search_mode"].(string); ok { + return s + } + case datatypes.JSON: + if len(v) == 0 { + return "" + } + var m map[string]any + if err := json.Unmarshal(v, &m); err == nil { + if s, ok := m["search_mode"].(string); ok { + return s + } + } + case []byte: + if len(v) == 0 { + return "" + } + var m map[string]any + if err := json.Unmarshal(v, &m); err == nil { + if s, ok := m["search_mode"].(string); ok { + return s + } + } + } + return "" +}