From dd965e760c13a65de60556504358ed80611cc7f5 Mon Sep 17 00:00:00 2001 From: mlogclub Date: Mon, 13 Apr 2026 20:03:04 +0800 Subject: [PATCH] feat: refactor retrieve logic to enhance trace handling and knowledge base preparation --- internal/ai/rag/retrieve.go | 26 +++++----------------- internal/ai/rag/retrieve_flow.go | 38 ++++++++++++++++++++++++++++++++ 2 files changed, 43 insertions(+), 21 deletions(-) create mode 100644 internal/ai/rag/retrieve_flow.go diff --git a/internal/ai/rag/retrieve.go b/internal/ai/rag/retrieve.go index 9dc0191..c6aacbd 100644 --- a/internal/ai/rag/retrieve.go +++ b/internal/ai/rag/retrieve.go @@ -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 diff --git a/internal/ai/rag/retrieve_flow.go b/internal/ai/rag/retrieve_flow.go new file mode 100644 index 0000000..24792a0 --- /dev/null +++ b/internal/ai/rag/retrieve_flow.go @@ -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 +}