feat: refactor retrieve logic to enhance trace handling and knowledge base preparation

This commit is contained in:
mlogclub
2026-04-13 20:03:04 +08:00
parent a5cf31656e
commit dd965e760c
2 changed files with 43 additions and 21 deletions
+5 -21
View File
@@ -31,34 +31,18 @@ type RetrieveTrace struct {
}
func (s *retrieve) RetrieveWithTrace(ctx context.Context, req RetrieveRequest) ([]RetrieveResult, *RetrieveTrace, error) {
trace := &RetrieveTrace{}
if req.Query == "" {
return nil, trace, nil
}
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))
trace := newRetrieveTrace()
retrievableKnowledgeBases, _, ok := s.prepareRetrievableKnowledgeBases(req, trace)
if !ok {
return nil, trace, nil
}
searchResults, searchTrace, err := s.searchKnowledgeBaseVectors(ctx, req, retrievableKnowledgeBases)
if err != nil {
if searchTrace != nil {
trace.EmbeddingMs = searchTrace.EmbeddingMs
trace.VectorSearchMs = searchTrace.VectorSearchMs
}
applySearchTrace(trace, searchTrace)
return nil, trace, err
}
if searchTrace != nil {
trace.EmbeddingMs = searchTrace.EmbeddingMs
trace.VectorSearchMs = searchTrace.VectorSearchMs
}
applySearchTrace(trace, searchTrace)
if len(searchResults) == 0 {
return nil, trace, nil
+38
View File
@@ -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
}