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...)
|
messages = append(messages, history.Messages...)
|
||||||
|
|
||||||
retriever := retrievers.NewKnowledgeRetriever(req.AIAgent)
|
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)
|
summary.RetrieverCount = len(retrieveResult.Hits)
|
||||||
collector.Data.Retriever.Count = len(retrieveResult.Hits)
|
collector.Data.Retriever.Count = len(retrieveResult.Hits)
|
||||||
for _, item := range retrieveResult.Hits {
|
collector.Data.Retriever.Items = append(collector.Data.Retriever.Items, retrieveResult.TraceItems...)
|
||||||
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),
|
|
||||||
})
|
|
||||||
}
|
|
||||||
if strings.TrimSpace(retrieveResult.ContextText) != "" {
|
if strings.TrimSpace(retrieveResult.ContextText) != "" {
|
||||||
messages = append(messages, schema.SystemMessage(retrieveResult.ContextText))
|
messages = append(messages, schema.SystemMessage(retrieveResult.ContextText))
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ import (
|
|||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
"cs-agent/internal/ai/rag"
|
"cs-agent/internal/ai/rag"
|
||||||
|
"cs-agent/internal/ai/runtime/internal/impl/callbacks"
|
||||||
"cs-agent/internal/models"
|
"cs-agent/internal/models"
|
||||||
"cs-agent/internal/pkg/utils"
|
"cs-agent/internal/pkg/utils"
|
||||||
)
|
)
|
||||||
@@ -15,6 +16,13 @@ type KnowledgeRetriever struct {
|
|||||||
AIAgent *models.AIAgent
|
AIAgent *models.AIAgent
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type KnowledgeRetrieveOptions struct {
|
||||||
|
ContextMaxTokens int
|
||||||
|
TopK int
|
||||||
|
ScoreThreshold float64
|
||||||
|
QueryPreview string
|
||||||
|
}
|
||||||
|
|
||||||
type KnowledgeRetrieveResult struct {
|
type KnowledgeRetrieveResult struct {
|
||||||
KnowledgeBaseIDs []int64
|
KnowledgeBaseIDs []int64
|
||||||
Query string
|
Query string
|
||||||
@@ -22,6 +30,7 @@ type KnowledgeRetrieveResult struct {
|
|||||||
ContextResults []rag.RetrieveResult
|
ContextResults []rag.RetrieveResult
|
||||||
ContextText string
|
ContextText string
|
||||||
Trace *rag.RetrieveTrace
|
Trace *rag.RetrieveTrace
|
||||||
|
TraceItems []callbacks.RetrieverTraceItem
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewKnowledgeRetriever(aiAgent *models.AIAgent) *KnowledgeRetriever {
|
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) {
|
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()
|
ids := r.KnowledgeBaseIDs()
|
||||||
return rag.Retrieve.RetrieveWithTrace(ctx, rag.RetrieveRequest{
|
return rag.Retrieve.RetrieveWithTrace(ctx, rag.RetrieveRequest{
|
||||||
Query: query,
|
Query: query,
|
||||||
KnowledgeBaseIDs: ids,
|
KnowledgeBaseIDs: ids,
|
||||||
|
TopK: opts.TopK,
|
||||||
|
ScoreThreshold: opts.ScoreThreshold,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *KnowledgeRetriever) RetrieveContext(ctx context.Context, query string) (*KnowledgeRetrieveResult, error) {
|
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)
|
query = strings.TrimSpace(query)
|
||||||
knowledgeBaseIDs := r.KnowledgeBaseIDs()
|
knowledgeBaseIDs := r.KnowledgeBaseIDs()
|
||||||
|
contextMaxTokens := opts.ContextMaxTokens
|
||||||
|
if contextMaxTokens <= 0 {
|
||||||
|
contextMaxTokens = defaultRuntimeKnowledgeContextTokens
|
||||||
|
}
|
||||||
|
queryPreview := strings.TrimSpace(opts.QueryPreview)
|
||||||
|
if queryPreview == "" {
|
||||||
|
queryPreview = query
|
||||||
|
}
|
||||||
ret := &KnowledgeRetrieveResult{
|
ret := &KnowledgeRetrieveResult{
|
||||||
KnowledgeBaseIDs: append([]int64(nil), knowledgeBaseIDs...),
|
KnowledgeBaseIDs: append([]int64(nil), knowledgeBaseIDs...),
|
||||||
Query: query,
|
Query: query,
|
||||||
@@ -53,13 +80,36 @@ func (r *KnowledgeRetriever) RetrieveContext(ctx context.Context, query string)
|
|||||||
if query == "" || len(knowledgeBaseIDs) == 0 {
|
if query == "" || len(knowledgeBaseIDs) == 0 {
|
||||||
return ret, nil
|
return ret, nil
|
||||||
}
|
}
|
||||||
results, trace, err := r.Retrieve(ctx, query)
|
results, trace, err := r.RetrieveByOptions(ctx, opts, query)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
ret.Hits = append([]rag.RetrieveResult(nil), results...)
|
ret.Hits = append([]rag.RetrieveResult(nil), results...)
|
||||||
ret.Trace = trace
|
ret.Trace = trace
|
||||||
ret.ContextResults = rag.Retrieve.SelectContextResults(results, defaultRuntimeKnowledgeContextTokens)
|
ret.ContextResults = rag.Retrieve.SelectContextResults(results, contextMaxTokens)
|
||||||
ret.ContextText = strings.TrimSpace(rag.Retrieve.BuildContext(ctx, results, defaultRuntimeKnowledgeContextTokens))
|
ret.ContextText = strings.TrimSpace(rag.Retrieve.BuildContext(ctx, results, contextMaxTokens))
|
||||||
|
ret.TraceItems = buildRetrieverTraceItems(queryPreview, results, trace)
|
||||||
return ret, nil
|
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