diff --git a/internal/ai/runtime/internal/engine/service.go b/internal/ai/runtime/internal/engine/service.go index 06f2190..1a1d825 100644 --- a/internal/ai/runtime/internal/engine/service.go +++ b/internal/ai/runtime/internal/engine/service.go @@ -163,7 +163,7 @@ func (s *Service) Run(ctx context.Context, req Request) (*Summary, error) { QueryPreview: preview(req.UserMessage.Content, 120), }, strings.TrimSpace(req.UserMessage.Content)); retrieveErr == nil && retrieveResult != nil { summary.RetrieverCount = len(retrieveResult.Hits) - collector.Data.Retriever.Count = len(retrieveResult.Hits) + collector.SetRetrieverSummary(retrieveResult.TraceSummary) collector.Data.Retriever.Items = append(collector.Data.Retriever.Items, retrieveResult.TraceItems...) if strings.TrimSpace(retrieveResult.ContextText) != "" { messages = append(messages, schema.SystemMessage(retrieveResult.ContextText)) diff --git a/internal/ai/runtime/internal/impl/callbacks/runlog_callback.go b/internal/ai/runtime/internal/impl/callbacks/runlog_callback.go index 9fb1090..e1211dd 100644 --- a/internal/ai/runtime/internal/impl/callbacks/runlog_callback.go +++ b/internal/ai/runtime/internal/impl/callbacks/runlog_callback.go @@ -65,6 +65,22 @@ func (c *RuntimeTraceCollector) SetSkillMiddleware(enabled bool, toolName string c.Data.Skill.MiddlewareToolName = toolName } +func (c *RuntimeTraceCollector) SetRetrieverSummary(summary RetrieverTraceSummary) { + if c == nil { + return + } + c.mu.Lock() + defer c.mu.Unlock() + c.Data.Retriever.TopK = summary.TopK + c.Data.Retriever.ScoreThreshold = summary.ScoreThreshold + c.Data.Retriever.ContextMaxTokens = summary.ContextMaxTokens + c.Data.Retriever.Count = summary.HitCount + c.Data.Retriever.ContextCount = summary.ContextCount + c.Data.Retriever.EmbeddingMs = summary.EmbeddingMs + c.Data.Retriever.VectorSearchMs = summary.VectorSearchMs + c.Data.Retriever.HydrateMs = summary.HydrateMs +} + func (c *RuntimeTraceCollector) AddToolItem(item ToolTraceItem) { if c == nil { return diff --git a/internal/ai/runtime/internal/impl/callbacks/trace_callback.go b/internal/ai/runtime/internal/impl/callbacks/trace_callback.go index df80b62..5e45c3f 100644 --- a/internal/ai/runtime/internal/impl/callbacks/trace_callback.go +++ b/internal/ai/runtime/internal/impl/callbacks/trace_callback.go @@ -41,6 +41,17 @@ type RetrieverTraceItem struct { LatencyMs int64 `json:"latencyMs,omitempty"` } +type RetrieverTraceSummary struct { + TopK int + ScoreThreshold float64 + ContextMaxTokens int + HitCount int + ContextCount int + EmbeddingMs int64 + VectorSearchMs int64 + HydrateMs int64 +} + type InstructionTraceSummary struct { SectionTitles []string HasProjectRule bool @@ -81,8 +92,15 @@ type RuntimeTraceData struct { CurrentUserMessagePreview string `json:"currentUserMessagePreview,omitempty"` } `json:"input"` Retriever struct { - Count int `json:"count,omitempty"` - Items []RetrieverTraceItem `json:"items,omitempty"` + Count int `json:"count,omitempty"` + TopK int `json:"topK,omitempty"` + ScoreThreshold float64 `json:"scoreThreshold,omitempty"` + ContextMaxTokens int `json:"contextMaxTokens,omitempty"` + ContextCount int `json:"contextCount,omitempty"` + EmbeddingMs int64 `json:"embeddingMs,omitempty"` + VectorSearchMs int64 `json:"vectorSearchMs,omitempty"` + HydrateMs int64 `json:"hydrateMs,omitempty"` + Items []RetrieverTraceItem `json:"items,omitempty"` } `json:"retriever"` Tools struct { Count int `json:"count,omitempty"` diff --git a/internal/ai/runtime/internal/impl/retrievers/knowledge_retriever.go b/internal/ai/runtime/internal/impl/retrievers/knowledge_retriever.go index bc9869b..901c70a 100644 --- a/internal/ai/runtime/internal/impl/retrievers/knowledge_retriever.go +++ b/internal/ai/runtime/internal/impl/retrievers/knowledge_retriever.go @@ -26,11 +26,13 @@ type KnowledgeRetrieveOptions struct { type KnowledgeRetrieveResult struct { KnowledgeBaseIDs []int64 Query string + Options KnowledgeRetrieveOptions Hits []rag.RetrieveResult ContextResults []rag.RetrieveResult ContextText string Trace *rag.RetrieveTrace TraceItems []callbacks.RetrieverTraceItem + TraceSummary callbacks.RetrieverTraceSummary } func NewKnowledgeRetriever(aiAgent *models.AIAgent) *KnowledgeRetriever { @@ -76,6 +78,12 @@ func (r *KnowledgeRetriever) RetrieveContextByOptions(ctx context.Context, opts ret := &KnowledgeRetrieveResult{ KnowledgeBaseIDs: append([]int64(nil), knowledgeBaseIDs...), Query: query, + Options: KnowledgeRetrieveOptions{ + ContextMaxTokens: contextMaxTokens, + TopK: opts.TopK, + ScoreThreshold: opts.ScoreThreshold, + QueryPreview: queryPreview, + }, } if query == "" || len(knowledgeBaseIDs) == 0 { return ret, nil @@ -89,6 +97,7 @@ func (r *KnowledgeRetriever) RetrieveContextByOptions(ctx context.Context, opts ret.ContextResults = rag.Retrieve.SelectContextResults(results, contextMaxTokens) ret.ContextText = strings.TrimSpace(rag.Retrieve.BuildContext(ctx, results, contextMaxTokens)) ret.TraceItems = buildRetrieverTraceItems(queryPreview, results, trace) + ret.TraceSummary = buildRetrieverTraceSummary(ret.Options, ret.ContextResults, results, trace) return ret, nil } @@ -113,3 +122,19 @@ func buildRetrieverTraceItems(queryPreview string, results []rag.RetrieveResult, } return ret } + +func buildRetrieverTraceSummary(opts KnowledgeRetrieveOptions, contextResults []rag.RetrieveResult, results []rag.RetrieveResult, trace *rag.RetrieveTrace) callbacks.RetrieverTraceSummary { + ret := callbacks.RetrieverTraceSummary{ + TopK: opts.TopK, + ScoreThreshold: opts.ScoreThreshold, + ContextMaxTokens: opts.ContextMaxTokens, + HitCount: len(results), + ContextCount: len(contextResults), + } + if trace != nil { + ret.EmbeddingMs = trace.EmbeddingMs + ret.VectorSearchMs = trace.VectorSearchMs + ret.HydrateMs = trace.HydrateMs + } + return ret +}