From 4cd31ce454eb29ebadb2f318a6f3d79195ddf7a6 Mon Sep 17 00:00:00 2001 From: mlogclub Date: Mon, 13 Apr 2026 20:03:28 +0800 Subject: [PATCH] feat: add document and FAQ loading methods to streamline indexing logic --- internal/ai/rag/index.go | 27 ++++++++---------- internal/ai/rag/index_flow.go | 52 +++++++++++++++++++++++++++++++++++ 2 files changed, 64 insertions(+), 15 deletions(-) create mode 100644 internal/ai/rag/index_flow.go diff --git a/internal/ai/rag/index.go b/internal/ai/rag/index.go index 46f6685..179a312 100644 --- a/internal/ai/rag/index.go +++ b/internal/ai/rag/index.go @@ -48,9 +48,9 @@ var Index = &index{ } func (s *index) IndexDocumentByID(ctx context.Context, documentID int64) error { - document := repositories.KnowledgeDocumentRepository.Get(sqls.DB(), documentID) - if document == nil { - return fmt.Errorf("document not found: %d", documentID) + document, err := s.loadDocumentByID(documentID) + if err != nil { + return err } return s.IndexDocument(ctx, document) } @@ -69,9 +69,9 @@ func (s *index) IndexDocument(ctx context.Context, document *models.KnowledgeDoc } // TODO 这里每次都查询下知识库不太友好 - knowledgeBase := repositories.KnowledgeBaseRepository.Get(sqls.DB(), document.KnowledgeBaseID) - if knowledgeBase == nil { - return fail(fmt.Errorf("knowledge base not found: %d", document.KnowledgeBaseID)) + knowledgeBase, err := s.loadDocumentKnowledgeBase(document) + if err != nil { + return fail(err) } existingChunks := repositories.KnowledgeChunkRepository.FindByDocumentID(sqls.DB(), document.ID) @@ -130,9 +130,9 @@ func (s *index) IndexDocument(ctx context.Context, document *models.KnowledgeDoc } func (s *index) IndexFAQByID(ctx context.Context, faqID int64) error { - faq := repositories.KnowledgeFAQRepository.Get(sqls.DB(), faqID) - if faq == nil { - return fmt.Errorf("faq not found: %d", faqID) + faq, err := s.loadFAQByID(faqID) + if err != nil { + return err } if err := s.markFAQIndexPending(faq.ID); err != nil { slog.Error("Failed to mark knowledge faq index as pending", "faq_id", faq.ID, "error", err) @@ -143,12 +143,9 @@ func (s *index) IndexFAQByID(ctx context.Context, faqID int64) error { } return err } - knowledgeBase := repositories.KnowledgeBaseRepository.Get(sqls.DB(), faq.KnowledgeBaseID) - if knowledgeBase == nil { - return fail(fmt.Errorf("knowledge base not found: %d", faq.KnowledgeBaseID)) - } - if knowledgeBase.KnowledgeType != string(enums.KnowledgeBaseTypeFAQ) { - return fail(fmt.Errorf("knowledge base %d is not faq type", knowledgeBase.ID)) + knowledgeBase, err := s.loadFAQKnowledgeBase(faq) + if err != nil { + return fail(err) } existingChunks := repositories.KnowledgeChunkRepository.FindByFaqID(sqls.DB(), faq.ID) content := buildFAQChunkContent(faq) diff --git a/internal/ai/rag/index_flow.go b/internal/ai/rag/index_flow.go new file mode 100644 index 0000000..26b3b20 --- /dev/null +++ b/internal/ai/rag/index_flow.go @@ -0,0 +1,52 @@ +package rag + +import ( + "fmt" + + "cs-agent/internal/models" + "cs-agent/internal/pkg/enums" + "cs-agent/internal/repositories" + + "github.com/mlogclub/simple/sqls" +) + +func (s *index) loadDocumentByID(documentID int64) (*models.KnowledgeDocument, error) { + document := repositories.KnowledgeDocumentRepository.Get(sqls.DB(), documentID) + if document == nil { + return nil, fmt.Errorf("document not found: %d", documentID) + } + return document, nil +} + +func (s *index) loadFAQByID(faqID int64) (*models.KnowledgeFAQ, error) { + faq := repositories.KnowledgeFAQRepository.Get(sqls.DB(), faqID) + if faq == nil { + return nil, fmt.Errorf("faq not found: %d", faqID) + } + return faq, nil +} + +func (s *index) loadDocumentKnowledgeBase(document *models.KnowledgeDocument) (*models.KnowledgeBase, error) { + if document == nil { + return nil, fmt.Errorf("document is nil") + } + knowledgeBase := repositories.KnowledgeBaseRepository.Get(sqls.DB(), document.KnowledgeBaseID) + if knowledgeBase == nil { + return nil, fmt.Errorf("knowledge base not found: %d", document.KnowledgeBaseID) + } + return knowledgeBase, nil +} + +func (s *index) loadFAQKnowledgeBase(faq *models.KnowledgeFAQ) (*models.KnowledgeBase, error) { + if faq == nil { + return nil, fmt.Errorf("faq is nil") + } + knowledgeBase := repositories.KnowledgeBaseRepository.Get(sqls.DB(), faq.KnowledgeBaseID) + if knowledgeBase == nil { + return nil, fmt.Errorf("knowledge base not found: %d", faq.KnowledgeBaseID) + } + if knowledgeBase.KnowledgeType != string(enums.KnowledgeBaseTypeFAQ) { + return nil, fmt.Errorf("knowledge base %d is not faq type", knowledgeBase.ID) + } + return knowledgeBase, nil +}