From 9ee3f3abbdba9adf55c073bb22e0b6980cce871c Mon Sep 17 00:00:00 2001 From: mlogclub Date: Sun, 12 Apr 2026 10:45:11 +0800 Subject: [PATCH] feat: enhance knowledge retrieval by introducing RetrieveContextByOptions method and streamline context handling --- .../ai/runtime/internal/engine/service.go | 14 ++--- .../impl/retrievers/knowledge_retriever.go | 56 ++++++++++++++++++- 2 files changed, 57 insertions(+), 13 deletions(-) diff --git a/internal/ai/runtime/internal/engine/service.go b/internal/ai/runtime/internal/engine/service.go index d671cb1..06f2190 100644 --- a/internal/ai/runtime/internal/engine/service.go +++ b/internal/ai/runtime/internal/engine/service.go @@ -159,18 +159,12 @@ func (s *Service) Run(ctx context.Context, req Request) (*Summary, error) { messages = append(messages, history.Messages...) retriever := retrievers.NewKnowledgeRetriever(req.AIAgent) - if retrieveResult, retrieveErr := retriever.RetrieveContext(ctx, strings.TrimSpace(req.UserMessage.Content)); retrieveErr == nil && retrieveResult != nil { + if retrieveResult, retrieveErr := retriever.RetrieveContextByOptions(ctx, retrievers.KnowledgeRetrieveOptions{ + 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) - for _, item := range retrieveResult.Hits { - collector.Data.Retriever.Items = append(collector.Data.Retriever.Items, callbacks.RetrieverTraceItem{ - Query: preview(req.UserMessage.Content, 120), - KnowledgeBaseID: item.KnowledgeBaseID, - DocumentID: item.DocumentID, - DocumentTitle: item.DocumentTitle, - Score: float64(item.Score), - }) - } + 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/retrievers/knowledge_retriever.go b/internal/ai/runtime/internal/impl/retrievers/knowledge_retriever.go index 92d5dc1..bc9869b 100644 --- a/internal/ai/runtime/internal/impl/retrievers/knowledge_retriever.go +++ b/internal/ai/runtime/internal/impl/retrievers/knowledge_retriever.go @@ -5,6 +5,7 @@ import ( "strings" "cs-agent/internal/ai/rag" + "cs-agent/internal/ai/runtime/internal/impl/callbacks" "cs-agent/internal/models" "cs-agent/internal/pkg/utils" ) @@ -15,6 +16,13 @@ type KnowledgeRetriever struct { AIAgent *models.AIAgent } +type KnowledgeRetrieveOptions struct { + ContextMaxTokens int + TopK int + ScoreThreshold float64 + QueryPreview string +} + type KnowledgeRetrieveResult struct { KnowledgeBaseIDs []int64 Query string @@ -22,6 +30,7 @@ type KnowledgeRetrieveResult struct { ContextResults []rag.RetrieveResult ContextText string Trace *rag.RetrieveTrace + TraceItems []callbacks.RetrieverTraceItem } func NewKnowledgeRetriever(aiAgent *models.AIAgent) *KnowledgeRetriever { @@ -36,16 +45,34 @@ func (r *KnowledgeRetriever) KnowledgeBaseIDs() []int64 { } func (r *KnowledgeRetriever) Retrieve(ctx context.Context, query string) ([]rag.RetrieveResult, *rag.RetrieveTrace, error) { + return r.RetrieveByOptions(ctx, KnowledgeRetrieveOptions{}, query) +} + +func (r *KnowledgeRetriever) RetrieveByOptions(ctx context.Context, opts KnowledgeRetrieveOptions, query string) ([]rag.RetrieveResult, *rag.RetrieveTrace, error) { ids := r.KnowledgeBaseIDs() return rag.Retrieve.RetrieveWithTrace(ctx, rag.RetrieveRequest{ Query: query, KnowledgeBaseIDs: ids, + TopK: opts.TopK, + ScoreThreshold: opts.ScoreThreshold, }) } func (r *KnowledgeRetriever) RetrieveContext(ctx context.Context, query string) (*KnowledgeRetrieveResult, error) { + return r.RetrieveContextByOptions(ctx, KnowledgeRetrieveOptions{}, query) +} + +func (r *KnowledgeRetriever) RetrieveContextByOptions(ctx context.Context, opts KnowledgeRetrieveOptions, query string) (*KnowledgeRetrieveResult, error) { query = strings.TrimSpace(query) knowledgeBaseIDs := r.KnowledgeBaseIDs() + contextMaxTokens := opts.ContextMaxTokens + if contextMaxTokens <= 0 { + contextMaxTokens = defaultRuntimeKnowledgeContextTokens + } + queryPreview := strings.TrimSpace(opts.QueryPreview) + if queryPreview == "" { + queryPreview = query + } ret := &KnowledgeRetrieveResult{ KnowledgeBaseIDs: append([]int64(nil), knowledgeBaseIDs...), Query: query, @@ -53,13 +80,36 @@ func (r *KnowledgeRetriever) RetrieveContext(ctx context.Context, query string) if query == "" || len(knowledgeBaseIDs) == 0 { return ret, nil } - results, trace, err := r.Retrieve(ctx, query) + results, trace, err := r.RetrieveByOptions(ctx, opts, query) if err != nil { return nil, err } ret.Hits = append([]rag.RetrieveResult(nil), results...) ret.Trace = trace - ret.ContextResults = rag.Retrieve.SelectContextResults(results, defaultRuntimeKnowledgeContextTokens) - ret.ContextText = strings.TrimSpace(rag.Retrieve.BuildContext(ctx, results, defaultRuntimeKnowledgeContextTokens)) + ret.ContextResults = rag.Retrieve.SelectContextResults(results, contextMaxTokens) + ret.ContextText = strings.TrimSpace(rag.Retrieve.BuildContext(ctx, results, contextMaxTokens)) + ret.TraceItems = buildRetrieverTraceItems(queryPreview, results, trace) return ret, nil } + +func buildRetrieverTraceItems(queryPreview string, results []rag.RetrieveResult, trace *rag.RetrieveTrace) []callbacks.RetrieverTraceItem { + if len(results) == 0 { + return nil + } + latencyMs := int64(0) + if trace != nil { + latencyMs = trace.EmbeddingMs + trace.VectorSearchMs + trace.HydrateMs + } + ret := make([]callbacks.RetrieverTraceItem, 0, len(results)) + for _, item := range results { + ret = append(ret, callbacks.RetrieverTraceItem{ + Query: queryPreview, + KnowledgeBaseID: item.KnowledgeBaseID, + DocumentID: item.DocumentID, + DocumentTitle: item.DocumentTitle, + Score: float64(item.Score), + LatencyMs: latencyMs, + }) + } + return ret +}