feat: refactor knowledge retrieval to introduce RetrieveContext method and streamline context handling
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user