feat: enhance knowledge retrieval by introducing RetrieveContextByOptions method and streamline context handling
This commit is contained in:
@@ -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))
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user