feat: refactor retrieve logic to enhance trace handling and knowledge base preparation
This commit is contained in:
@@ -31,34 +31,18 @@ type RetrieveTrace struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (s *retrieve) RetrieveWithTrace(ctx context.Context, req RetrieveRequest) ([]RetrieveResult, *RetrieveTrace, error) {
|
func (s *retrieve) RetrieveWithTrace(ctx context.Context, req RetrieveRequest) ([]RetrieveResult, *RetrieveTrace, error) {
|
||||||
trace := &RetrieveTrace{}
|
trace := newRetrieveTrace()
|
||||||
if req.Query == "" {
|
retrievableKnowledgeBases, _, ok := s.prepareRetrievableKnowledgeBases(req, trace)
|
||||||
return nil, trace, nil
|
if !ok {
|
||||||
}
|
|
||||||
knowledgeBaseIDs := normalizeKnowledgeBaseIDs(req.KnowledgeBaseIDs)
|
|
||||||
if len(knowledgeBaseIDs) == 0 {
|
|
||||||
return nil, trace, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
retrievableKnowledgeBases := s.loadRetrievableKnowledgeBases(knowledgeBaseIDs)
|
|
||||||
if len(retrievableKnowledgeBases) == 0 {
|
|
||||||
slog.Info("Skip retrieve for non-enabled knowledge bases",
|
|
||||||
"knowledge_base_ids", fmt.Sprint(knowledgeBaseIDs))
|
|
||||||
return nil, trace, nil
|
return nil, trace, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
searchResults, searchTrace, err := s.searchKnowledgeBaseVectors(ctx, req, retrievableKnowledgeBases)
|
searchResults, searchTrace, err := s.searchKnowledgeBaseVectors(ctx, req, retrievableKnowledgeBases)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if searchTrace != nil {
|
applySearchTrace(trace, searchTrace)
|
||||||
trace.EmbeddingMs = searchTrace.EmbeddingMs
|
|
||||||
trace.VectorSearchMs = searchTrace.VectorSearchMs
|
|
||||||
}
|
|
||||||
return nil, trace, err
|
return nil, trace, err
|
||||||
}
|
}
|
||||||
if searchTrace != nil {
|
applySearchTrace(trace, searchTrace)
|
||||||
trace.EmbeddingMs = searchTrace.EmbeddingMs
|
|
||||||
trace.VectorSearchMs = searchTrace.VectorSearchMs
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(searchResults) == 0 {
|
if len(searchResults) == 0 {
|
||||||
return nil, trace, nil
|
return nil, trace, nil
|
||||||
|
|||||||
@@ -0,0 +1,38 @@
|
|||||||
|
package rag
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"log/slog"
|
||||||
|
|
||||||
|
"cs-agent/internal/models"
|
||||||
|
)
|
||||||
|
|
||||||
|
func newRetrieveTrace() *RetrieveTrace {
|
||||||
|
return &RetrieveTrace{}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *retrieve) prepareRetrievableKnowledgeBases(req RetrieveRequest, trace *RetrieveTrace) ([]models.KnowledgeBase, []int64, bool) {
|
||||||
|
if req.Query == "" {
|
||||||
|
return nil, nil, false
|
||||||
|
}
|
||||||
|
knowledgeBaseIDs := normalizeKnowledgeBaseIDs(req.KnowledgeBaseIDs)
|
||||||
|
if len(knowledgeBaseIDs) == 0 {
|
||||||
|
return nil, nil, false
|
||||||
|
}
|
||||||
|
|
||||||
|
retrievableKnowledgeBases := s.loadRetrievableKnowledgeBases(knowledgeBaseIDs)
|
||||||
|
if len(retrievableKnowledgeBases) == 0 {
|
||||||
|
slog.Info("Skip retrieve for non-enabled knowledge bases",
|
||||||
|
"knowledge_base_ids", fmt.Sprint(knowledgeBaseIDs))
|
||||||
|
return nil, knowledgeBaseIDs, false
|
||||||
|
}
|
||||||
|
return retrievableKnowledgeBases, knowledgeBaseIDs, true
|
||||||
|
}
|
||||||
|
|
||||||
|
func applySearchTrace(target *RetrieveTrace, source *RetrieveTrace) {
|
||||||
|
if target == nil || source == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
target.EmbeddingMs = source.EmbeddingMs
|
||||||
|
target.VectorSearchMs = source.VectorSearchMs
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user