From cf8469342149ac6ffb606a704219c8684c21f1ff Mon Sep 17 00:00:00 2001 From: mlogclub Date: Sun, 12 Apr 2026 10:34:24 +0800 Subject: [PATCH] feat: refactor knowledge retrieval to introduce RetrieveContext method and streamline context handling --- .../ai/runtime/internal/engine/service.go | 41 +++---------------- .../impl/retrievers/knowledge_retriever.go | 33 +++++++++++++++ 2 files changed, 39 insertions(+), 35 deletions(-) diff --git a/internal/ai/runtime/internal/engine/service.go b/internal/ai/runtime/internal/engine/service.go index 55d6df6..d671cb1 100644 --- a/internal/ai/runtime/internal/engine/service.go +++ b/internal/ai/runtime/internal/engine/service.go @@ -6,7 +6,6 @@ import ( "fmt" "strings" - "cs-agent/internal/ai/rag" "cs-agent/internal/ai/runtime/internal/impl/adapter" "cs-agent/internal/ai/runtime/internal/impl/callbacks" "cs-agent/internal/ai/runtime/internal/impl/factory" @@ -160,10 +159,10 @@ func (s *Service) Run(ctx context.Context, req Request) (*Summary, error) { messages = append(messages, history.Messages...) retriever := retrievers.NewKnowledgeRetriever(req.AIAgent) - if results, _, retrieveErr := retriever.Retrieve(ctx, strings.TrimSpace(req.UserMessage.Content)); retrieveErr == nil { - summary.RetrieverCount = len(results) - collector.Data.Retriever.Count = len(results) - for _, item := range results { + if retrieveResult, retrieveErr := retriever.RetrieveContext(ctx, 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, @@ -172,8 +171,8 @@ func (s *Service) Run(ctx context.Context, req Request) (*Summary, error) { Score: float64(item.Score), }) } - if knowledgeContext := buildKnowledgeContext(results); knowledgeContext != "" { - messages = append(messages, schema.SystemMessage(knowledgeContext)) + if strings.TrimSpace(retrieveResult.ContextText) != "" { + messages = append(messages, schema.SystemMessage(retrieveResult.ContextText)) } } @@ -525,34 +524,6 @@ func preview(value string, limit int) string { return string(runes[:limit]) + "..." } -func buildKnowledgeContext(items []rag.RetrieveResult) string { - if len(items) == 0 { - return "" - } - var builder strings.Builder - builder.WriteString("以下是可供参考的知识库内容,请优先基于这些内容回答;如果仍不确定,请明确说明并向用户澄清。\n\n") - for i, item := range items { - if i >= 5 { - break - } - builder.WriteString("[知识片段") - builder.WriteString(fmt.Sprintf("%d", i+1)) - builder.WriteString("]\n") - if strings.TrimSpace(item.DocumentTitle) != "" { - builder.WriteString("标题: ") - builder.WriteString(strings.TrimSpace(item.DocumentTitle)) - builder.WriteString("\n") - } - if strings.TrimSpace(item.Content) != "" { - builder.WriteString("内容: ") - builder.WriteString(strings.TrimSpace(item.Content)) - builder.WriteString("\n") - } - builder.WriteString("\n") - } - return strings.TrimSpace(builder.String()) -} - func toolSetStaticTools(toolSet *registry.ToolSet) []einotool.BaseTool { if toolSet == nil { return nil diff --git a/internal/ai/runtime/internal/impl/retrievers/knowledge_retriever.go b/internal/ai/runtime/internal/impl/retrievers/knowledge_retriever.go index e21fd1f..92d5dc1 100644 --- a/internal/ai/runtime/internal/impl/retrievers/knowledge_retriever.go +++ b/internal/ai/runtime/internal/impl/retrievers/knowledge_retriever.go @@ -2,16 +2,28 @@ package retrievers import ( "context" + "strings" "cs-agent/internal/ai/rag" "cs-agent/internal/models" "cs-agent/internal/pkg/utils" ) +const defaultRuntimeKnowledgeContextTokens = 4000 + type KnowledgeRetriever struct { AIAgent *models.AIAgent } +type KnowledgeRetrieveResult struct { + KnowledgeBaseIDs []int64 + Query string + Hits []rag.RetrieveResult + ContextResults []rag.RetrieveResult + ContextText string + Trace *rag.RetrieveTrace +} + func NewKnowledgeRetriever(aiAgent *models.AIAgent) *KnowledgeRetriever { return &KnowledgeRetriever{AIAgent: aiAgent} } @@ -30,3 +42,24 @@ func (r *KnowledgeRetriever) Retrieve(ctx context.Context, query string) ([]rag. KnowledgeBaseIDs: ids, }) } + +func (r *KnowledgeRetriever) RetrieveContext(ctx context.Context, query string) (*KnowledgeRetrieveResult, error) { + query = strings.TrimSpace(query) + knowledgeBaseIDs := r.KnowledgeBaseIDs() + ret := &KnowledgeRetrieveResult{ + KnowledgeBaseIDs: append([]int64(nil), knowledgeBaseIDs...), + Query: query, + } + if query == "" || len(knowledgeBaseIDs) == 0 { + return ret, nil + } + results, trace, err := r.Retrieve(ctx, 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)) + return ret, nil +}