refactor: streamline document and FAQ index removal methods and enhance knowledge base deletion logic
This commit is contained in:
@@ -1,8 +1,11 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"agent-desk/internal/ai/rag"
|
||||
"agent-desk/internal/models"
|
||||
"agent-desk/internal/pkg/dto"
|
||||
"agent-desk/internal/pkg/dto/request"
|
||||
@@ -125,16 +128,32 @@ func (s *knowledgeBaseService) DeleteKnowledgeBase(id int64) error {
|
||||
if current == nil {
|
||||
return errorsx.InvalidParam("知识库不存在")
|
||||
}
|
||||
docCount := repositories.KnowledgeDocumentRepository.CountByKnowledgeBaseID(sqls.DB(), id)
|
||||
if docCount > 0 {
|
||||
return errorsx.InvalidParam("知识库下存在文档,无法删除")
|
||||
|
||||
referencingAgents := repositories.AIAgentRepository.FindByKnowledgeBaseID(sqls.DB(), id)
|
||||
if len(referencingAgents) > 0 {
|
||||
if len(referencingAgents) == 1 {
|
||||
return errorsx.Forbidden(fmt.Sprintf("知识库已被 AI Agent「%s」引用,请先解除绑定", referencingAgents[0].Name))
|
||||
}
|
||||
return errorsx.Forbidden(fmt.Sprintf("知识库已被 %d 个 AI Agent 引用,请先解除绑定", len(referencingAgents)))
|
||||
}
|
||||
faqCount := repositories.KnowledgeFAQRepository.CountByKnowledgeBaseID(sqls.DB(), id)
|
||||
if faqCount > 0 {
|
||||
return errorsx.InvalidParam("知识库下存在FAQ,无法删除")
|
||||
|
||||
chunks := repositories.KnowledgeChunkRepository.Find(sqls.DB(), sqls.NewCnd().Eq("knowledge_base_id", id))
|
||||
if err := sqls.WithTransaction(func(ctx *sqls.TxContext) error {
|
||||
if err := repositories.KnowledgeChunkRepository.DeleteByKnowledgeBaseID(ctx.Tx, id); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := repositories.KnowledgeDocumentRepository.DeleteByKnowledgeBaseID(ctx.Tx, id); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := repositories.KnowledgeFAQRepository.DeleteByKnowledgeBaseID(ctx.Tx, id); err != nil {
|
||||
return err
|
||||
}
|
||||
return ctx.Tx.Delete(&models.KnowledgeBase{}, "id = ?", id).Error
|
||||
}); err != nil {
|
||||
return err
|
||||
}
|
||||
repositories.KnowledgeBaseRepository.Delete(sqls.DB(), id)
|
||||
return nil
|
||||
|
||||
return rag.Index.RemoveKnowledgeBaseIndexByChunkModels(context.Background(), id, chunks)
|
||||
}
|
||||
|
||||
func (s *knowledgeBaseService) UpdateSort(ids []int64) error {
|
||||
|
||||
@@ -1,9 +1,18 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"agent-desk/internal/models"
|
||||
"agent-desk/internal/pkg/dto/request"
|
||||
"agent-desk/internal/pkg/enums"
|
||||
"agent-desk/internal/repositories"
|
||||
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/mlogclub/simple/sqls"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func TestBuildKnowledgeBaseModelUsesLowerDefaultScoreThreshold(t *testing.T) {
|
||||
@@ -15,3 +24,111 @@ func TestBuildKnowledgeBaseModelUsesLowerDefaultScoreThreshold(t *testing.T) {
|
||||
t.Fatalf("expected default score threshold 0.2, got %v", item.DefaultScoreThreshold)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDeleteKnowledgeBaseRejectsAIAgentReference(t *testing.T) {
|
||||
setupKnowledgeBaseServiceTestDB(t)
|
||||
kb := createKnowledgeBaseServiceTestBase(t, "Referenced KB")
|
||||
otherKB := createKnowledgeBaseServiceTestBase(t, "Other KB")
|
||||
if err := repositories.AIAgentRepository.Create(sqls.DB(), &models.AIAgent{
|
||||
Name: "Support Agent",
|
||||
Status: enums.StatusOk,
|
||||
KnowledgeIDs: "12",
|
||||
}); err != nil {
|
||||
t.Fatalf("create unrelated ai agent: %v", err)
|
||||
}
|
||||
if err := repositories.AIAgentRepository.Create(sqls.DB(), &models.AIAgent{
|
||||
Name: "Knowledge Agent",
|
||||
Status: enums.StatusOk,
|
||||
KnowledgeIDs: fmt.Sprintf("12,%d,%d", kb.ID, otherKB.ID),
|
||||
}); err != nil {
|
||||
t.Fatalf("create ai agent: %v", err)
|
||||
}
|
||||
|
||||
err := KnowledgeBaseService.DeleteKnowledgeBase(kb.ID)
|
||||
if err == nil {
|
||||
t.Fatal("DeleteKnowledgeBase() error is nil, want referenced knowledge base error")
|
||||
}
|
||||
if got := err.Error(); !strings.Contains(got, "Knowledge Agent") {
|
||||
t.Fatalf("DeleteKnowledgeBase() error = %q, want agent name", got)
|
||||
}
|
||||
if repositories.KnowledgeBaseRepository.Get(sqls.DB(), kb.ID) == nil {
|
||||
t.Fatal("knowledge base was deleted despite ai agent reference")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDeleteKnowledgeBaseCascadesContentWhenNotReferenced(t *testing.T) {
|
||||
setupKnowledgeBaseServiceTestDB(t)
|
||||
kb := createKnowledgeBaseServiceTestBase(t, "Delete KB")
|
||||
document := &models.KnowledgeDocument{
|
||||
KnowledgeBaseID: kb.ID,
|
||||
Title: "Doc",
|
||||
ContentType: enums.KnowledgeDocumentContentTypeMarkdown,
|
||||
Content: "content",
|
||||
Status: enums.StatusOk,
|
||||
IndexStatus: enums.KnowledgeDocumentIndexStatusIndexed,
|
||||
}
|
||||
if err := repositories.KnowledgeDocumentRepository.Create(sqls.DB(), document); err != nil {
|
||||
t.Fatalf("create document: %v", err)
|
||||
}
|
||||
faq := &models.KnowledgeFAQ{
|
||||
KnowledgeBaseID: kb.ID,
|
||||
Question: "Question",
|
||||
Answer: "Answer",
|
||||
Status: enums.StatusOk,
|
||||
IndexStatus: enums.KnowledgeDocumentIndexStatusIndexed,
|
||||
}
|
||||
if err := repositories.KnowledgeFAQRepository.Create(sqls.DB(), faq); err != nil {
|
||||
t.Fatalf("create faq: %v", err)
|
||||
}
|
||||
if err := repositories.KnowledgeChunkRepository.BatchCreate(sqls.DB(), []models.KnowledgeChunk{
|
||||
{KnowledgeBaseID: kb.ID, DocumentID: document.ID, ChunkNo: 1, Status: enums.StatusOk},
|
||||
{KnowledgeBaseID: kb.ID, FaqID: faq.ID, ChunkNo: 1, Status: enums.StatusOk},
|
||||
}); err != nil {
|
||||
t.Fatalf("create chunks: %v", err)
|
||||
}
|
||||
|
||||
if err := KnowledgeBaseService.DeleteKnowledgeBase(kb.ID); err != nil {
|
||||
t.Fatalf("DeleteKnowledgeBase() error = %v", err)
|
||||
}
|
||||
|
||||
assertKnowledgeBaseServiceTestCount(t, &models.KnowledgeBase{}, "id = ?", kb.ID, 0)
|
||||
assertKnowledgeBaseServiceTestCount(t, &models.KnowledgeDocument{}, "knowledge_base_id = ?", kb.ID, 0)
|
||||
assertKnowledgeBaseServiceTestCount(t, &models.KnowledgeFAQ{}, "knowledge_base_id = ?", kb.ID, 0)
|
||||
assertKnowledgeBaseServiceTestCount(t, &models.KnowledgeChunk{}, "knowledge_base_id = ?", kb.ID, 0)
|
||||
}
|
||||
|
||||
func setupKnowledgeBaseServiceTestDB(t *testing.T) {
|
||||
t.Helper()
|
||||
db, err := gorm.Open(sqlite.Open("file:"+t.Name()+"?mode=memory&cache=shared"), &gorm.Config{})
|
||||
if err != nil {
|
||||
t.Fatalf("open sqlite db: %v", err)
|
||||
}
|
||||
if err := db.AutoMigrate(&models.KnowledgeBase{}, &models.KnowledgeDocument{}, &models.KnowledgeFAQ{}, &models.KnowledgeChunk{}, &models.AIAgent{}); err != nil {
|
||||
t.Fatalf("auto migrate: %v", err)
|
||||
}
|
||||
sqls.SetDB(db)
|
||||
}
|
||||
|
||||
func createKnowledgeBaseServiceTestBase(t *testing.T, name string) *models.KnowledgeBase {
|
||||
t.Helper()
|
||||
item := &models.KnowledgeBase{
|
||||
Name: name,
|
||||
KnowledgeType: string(enums.KnowledgeBaseTypeDocument),
|
||||
Status: enums.StatusOk,
|
||||
}
|
||||
if err := repositories.KnowledgeBaseRepository.Create(sqls.DB(), item); err != nil {
|
||||
t.Fatalf("create knowledge base: %v", err)
|
||||
}
|
||||
return item
|
||||
}
|
||||
|
||||
func assertKnowledgeBaseServiceTestCount(t *testing.T, model any, query string, arg any, want int64) {
|
||||
t.Helper()
|
||||
var count int64
|
||||
if err := sqls.DB().Model(model).Where(query, arg).Count(&count).Error; err != nil {
|
||||
t.Fatalf("count %T: %v", model, err)
|
||||
}
|
||||
if count != want {
|
||||
t.Fatalf("count %T = %d, want %d", model, count, want)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -155,7 +155,7 @@ func (s *knowledgeDocumentService) UpdateKnowledgeDocument(req request.UpdateKno
|
||||
}
|
||||
|
||||
if oldKnowledgeBaseID != item.KnowledgeBaseID {
|
||||
if err := rag.Index.RemoveDocumentIndexFromKnowledgeBase(context.Background(), oldKnowledgeBaseID, req.ID); err != nil {
|
||||
if err := rag.Index.RemoveDocumentIndex(context.Background(), req.ID); err != nil {
|
||||
slog.Error("failed to remove old document index after knowledge base change", "document_id", req.ID, "knowledge_base_id", oldKnowledgeBaseID, "error", err)
|
||||
}
|
||||
}
|
||||
@@ -167,10 +167,6 @@ func (s *knowledgeDocumentService) UpdateKnowledgeDocument(req request.UpdateKno
|
||||
}
|
||||
|
||||
func (s *knowledgeDocumentService) DeleteKnowledgeDocument(id int64) error {
|
||||
current := s.Get(id)
|
||||
if current == nil {
|
||||
return errorsx.InvalidParam("文档不存在")
|
||||
}
|
||||
chunks := repositories.KnowledgeChunkRepository.FindByDocumentID(sqls.DB(), id)
|
||||
if err := sqls.WithTransaction(func(ctx *sqls.TxContext) error {
|
||||
_ = repositories.KnowledgeDocumentRepository.Updates(ctx.Tx, id, map[string]any{
|
||||
@@ -182,7 +178,7 @@ func (s *knowledgeDocumentService) DeleteKnowledgeDocument(id int64) error {
|
||||
}); err != nil {
|
||||
return err
|
||||
}
|
||||
return rag.Index.RemoveDocumentIndexByChunkModels(context.Background(), current.KnowledgeBaseID, id, chunks)
|
||||
return rag.Index.RemoveDocumentIndexByChunkModels(context.Background(), id, chunks)
|
||||
}
|
||||
|
||||
func (s *knowledgeDocumentService) buildKnowledgeDocumentModel(req request.CreateKnowledgeDocumentRequest) (*models.KnowledgeDocument, error) {
|
||||
|
||||
@@ -113,7 +113,7 @@ func (s *knowledgeFAQService) DeleteKnowledgeFAQ(id int64) error {
|
||||
}); err != nil {
|
||||
return err
|
||||
}
|
||||
return rag.Index.RemoveFAQIndexByChunkModels(context.Background(), current.KnowledgeBaseID, id, chunks)
|
||||
return rag.Index.RemoveFAQIndexByChunkModels(context.Background(), id, chunks)
|
||||
}
|
||||
|
||||
func (s *knowledgeFAQService) buildKnowledgeFAQModel(req request.CreateKnowledgeFAQRequest) (*models.KnowledgeFAQ, error) {
|
||||
|
||||
Reference in New Issue
Block a user