feat: refactor knowledge retrieval to introduce RetrieveContext method and streamline context handling

This commit is contained in:
mlogclub
2026-04-12 10:34:24 +08:00
parent 22b2a08e13
commit cf84693421
2 changed files with 39 additions and 35 deletions
+6 -35
View File
@@ -6,7 +6,6 @@ import (
"fmt" "fmt"
"strings" "strings"
"cs-agent/internal/ai/rag"
"cs-agent/internal/ai/runtime/internal/impl/adapter" "cs-agent/internal/ai/runtime/internal/impl/adapter"
"cs-agent/internal/ai/runtime/internal/impl/callbacks" "cs-agent/internal/ai/runtime/internal/impl/callbacks"
"cs-agent/internal/ai/runtime/internal/impl/factory" "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...) messages = append(messages, history.Messages...)
retriever := retrievers.NewKnowledgeRetriever(req.AIAgent) retriever := retrievers.NewKnowledgeRetriever(req.AIAgent)
if results, _, retrieveErr := retriever.Retrieve(ctx, strings.TrimSpace(req.UserMessage.Content)); retrieveErr == nil { if retrieveResult, retrieveErr := retriever.RetrieveContext(ctx, strings.TrimSpace(req.UserMessage.Content)); retrieveErr == nil && retrieveResult != nil {
summary.RetrieverCount = len(results) summary.RetrieverCount = len(retrieveResult.Hits)
collector.Data.Retriever.Count = len(results) collector.Data.Retriever.Count = len(retrieveResult.Hits)
for _, item := range results { for _, item := range retrieveResult.Hits {
collector.Data.Retriever.Items = append(collector.Data.Retriever.Items, callbacks.RetrieverTraceItem{ collector.Data.Retriever.Items = append(collector.Data.Retriever.Items, callbacks.RetrieverTraceItem{
Query: preview(req.UserMessage.Content, 120), Query: preview(req.UserMessage.Content, 120),
KnowledgeBaseID: item.KnowledgeBaseID, KnowledgeBaseID: item.KnowledgeBaseID,
@@ -172,8 +171,8 @@ func (s *Service) Run(ctx context.Context, req Request) (*Summary, error) {
Score: float64(item.Score), Score: float64(item.Score),
}) })
} }
if knowledgeContext := buildKnowledgeContext(results); knowledgeContext != "" { if strings.TrimSpace(retrieveResult.ContextText) != "" {
messages = append(messages, schema.SystemMessage(knowledgeContext)) messages = append(messages, schema.SystemMessage(retrieveResult.ContextText))
} }
} }
@@ -525,34 +524,6 @@ func preview(value string, limit int) string {
return string(runes[:limit]) + "..." 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 { func toolSetStaticTools(toolSet *registry.ToolSet) []einotool.BaseTool {
if toolSet == nil { if toolSet == nil {
return nil return nil
@@ -2,16 +2,28 @@ package retrievers
import ( import (
"context" "context"
"strings"
"cs-agent/internal/ai/rag" "cs-agent/internal/ai/rag"
"cs-agent/internal/models" "cs-agent/internal/models"
"cs-agent/internal/pkg/utils" "cs-agent/internal/pkg/utils"
) )
const defaultRuntimeKnowledgeContextTokens = 4000
type KnowledgeRetriever struct { type KnowledgeRetriever struct {
AIAgent *models.AIAgent 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 { func NewKnowledgeRetriever(aiAgent *models.AIAgent) *KnowledgeRetriever {
return &KnowledgeRetriever{AIAgent: aiAgent} return &KnowledgeRetriever{AIAgent: aiAgent}
} }
@@ -30,3 +42,24 @@ func (r *KnowledgeRetriever) Retrieve(ctx context.Context, query string) ([]rag.
KnowledgeBaseIDs: ids, 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
}