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