From 7855f36799050673a2776fae1aec40adfbb707a1 Mon Sep 17 00:00:00 2001 From: wizardchen Date: Tue, 23 Jun 2026 20:33:01 +0800 Subject: [PATCH] fix(model): block deletion when model is referenced by KB or agent Reject model delete with HTTP 400 when the current tenant still has knowledge bases or custom agents bound to the model ID. --- .../application/repository/custom_agent.go | 13 ++ .../application/repository/knowledgebase.go | 14 ++ .../application/repository/model_usage.go | 49 +++++ .../repository/model_usage_test.go | 113 +++++++++++ .../service/knowledgebase_pr3_test.go | 3 + internal/application/service/model.go | 57 +++++- .../application/service/model_delete_test.go | 175 ++++++++++++++++++ .../application/service/vectorstore_test.go | 3 + internal/handler/model.go | 4 + internal/types/interfaces/custom_agent.go | 4 + internal/types/interfaces/knowledgebase.go | 4 + 11 files changed, 437 insertions(+), 2 deletions(-) create mode 100644 internal/application/repository/model_usage.go create mode 100644 internal/application/repository/model_usage_test.go create mode 100644 internal/application/service/model_delete_test.go diff --git a/internal/application/repository/custom_agent.go b/internal/application/repository/custom_agent.go index 02a41c5ba..35bf1112a 100644 --- a/internal/application/repository/custom_agent.go +++ b/internal/application/repository/custom_agent.go @@ -60,3 +60,16 @@ func (r *customAgentRepository) UpdateAgent(ctx context.Context, agent *types.Cu func (r *customAgentRepository) DeleteAgent(ctx context.Context, id string, tenantID uint64) error { return r.db.WithContext(ctx).Where("id = ? AND tenant_id = ?", id, tenantID).Delete(&types.CustomAgent{}).Error } + +// CountByModelID counts active agents whose config references modelID. +func (r *customAgentRepository) CountByModelID( + ctx context.Context, tenantID uint64, modelID string, +) (int64, error) { + var count int64 + query := r.db.WithContext(ctx). + Model(&types.CustomAgent{}). + Where("tenant_id = ?", tenantID) + query = scopeCustomAgentsByModelID(query, modelID) + err := query.Count(&count).Error + return count, err +} diff --git a/internal/application/repository/knowledgebase.go b/internal/application/repository/knowledgebase.go index 798e0f634..b19517298 100644 --- a/internal/application/repository/knowledgebase.go +++ b/internal/application/repository/knowledgebase.go @@ -203,3 +203,17 @@ func (r *knowledgeBaseRepository) CountByVectorStoreID( Count(&count).Error return count, err } + +// CountByModelID counts active knowledge bases that reference modelID in any +// model-binding column (scalar fields or JSON config blobs). +func (r *knowledgeBaseRepository) CountByModelID( + ctx context.Context, tenantID uint64, modelID string, +) (int64, error) { + var count int64 + query := r.db.WithContext(ctx). + Model(&types.KnowledgeBase{}). + Where("tenant_id = ?", tenantID) + query = scopeKnowledgeBasesByModelID(query, modelID) + err := query.Count(&count).Error + return count, err +} diff --git a/internal/application/repository/model_usage.go b/internal/application/repository/model_usage.go new file mode 100644 index 000000000..79d837d13 --- /dev/null +++ b/internal/application/repository/model_usage.go @@ -0,0 +1,49 @@ +package repository + +import ( + "gorm.io/gorm" +) + +// scopeKnowledgeBasesByModelID filters knowledge_bases rows that reference +// modelID in any model-binding field. +func scopeKnowledgeBasesByModelID(db *gorm.DB, modelID string) *gorm.DB { + if db.Dialector.Name() == "postgres" { + return db.Where( + "embedding_model_id = ? OR summary_model_id = ? OR "+ + "image_processing_config->>'model_id' = ? OR "+ + "vlm_config->>'model_id' = ? OR "+ + "asr_config->>'model_id' = ? OR "+ + "wiki_config->>'synthesis_model_id' = ?", + modelID, modelID, modelID, modelID, modelID, modelID, + ) + } + return db.Where( + "embedding_model_id = ? OR summary_model_id = ? OR "+ + "json_extract(image_processing_config, '$.model_id') = ? OR "+ + "json_extract(vlm_config, '$.model_id') = ? OR "+ + "json_extract(asr_config, '$.model_id') = ? OR "+ + "json_extract(wiki_config, '$.synthesis_model_id') = ?", + modelID, modelID, modelID, modelID, modelID, modelID, + ) +} + +// scopeCustomAgentsByModelID filters custom_agents rows whose config JSON +// references modelID in any model-binding field. +func scopeCustomAgentsByModelID(db *gorm.DB, modelID string) *gorm.DB { + if db.Dialector.Name() == "postgres" { + return db.Where( + "config->>'model_id' = ? OR config->>'rerank_model_id' = ? OR "+ + "config->>'vlm_model_id' = ? OR config->>'asr_model_id' = ? OR "+ + "config->>'query_understand_model_id' = ?", + modelID, modelID, modelID, modelID, modelID, + ) + } + return db.Where( + "json_extract(config, '$.model_id') = ? OR "+ + "json_extract(config, '$.rerank_model_id') = ? OR "+ + "json_extract(config, '$.vlm_model_id') = ? OR "+ + "json_extract(config, '$.asr_model_id') = ? OR "+ + "json_extract(config, '$.query_understand_model_id') = ?", + modelID, modelID, modelID, modelID, modelID, + ) +} diff --git a/internal/application/repository/model_usage_test.go b/internal/application/repository/model_usage_test.go new file mode 100644 index 000000000..a8358fef3 --- /dev/null +++ b/internal/application/repository/model_usage_test.go @@ -0,0 +1,113 @@ +package repository + +import ( + "context" + "testing" + + "github.com/Tencent/WeKnora/internal/types" + "github.com/google/uuid" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/gorm" +) + +const customAgentsTestDDL = ` +CREATE TABLE IF NOT EXISTS custom_agents ( + id VARCHAR(36) NOT NULL, + name VARCHAR(255) NOT NULL, + description TEXT, + avatar VARCHAR(64), + is_builtin BOOLEAN NOT NULL DEFAULT 0, + tenant_id INTEGER NOT NULL, + created_by VARCHAR(36), + config TEXT NOT NULL DEFAULT '{}', + created_at DATETIME DEFAULT CURRENT_TIMESTAMP, + updated_at DATETIME DEFAULT CURRENT_TIMESTAMP, + deleted_at DATETIME, + PRIMARY KEY (id, tenant_id) +); +` + +func setupModelUsageTestDB(t *testing.T) *gorm.DB { + t.Helper() + db := setupKBTestDB(t) + require.NoError(t, db.Exec(customAgentsTestDDL).Error) + return db +} + +func TestCountByModelID_KnowledgeBase(t *testing.T) { + ctx := context.Background() + db := setupModelUsageTestDB(t) + repo := NewKnowledgeBaseRepository(db) + modelID := "embed-model-1" + + kb := makeKB(nil) + kb.EmbeddingModelID = modelID + require.NoError(t, db.Create(kb).Error) + + count, err := repo.CountByModelID(ctx, 1, modelID) + require.NoError(t, err) + assert.Equal(t, int64(1), count) + + count, err = repo.CountByModelID(ctx, 1, "other-model") + require.NoError(t, err) + assert.Equal(t, int64(0), count) + + kb2 := makeKB(nil) + kb2.ID = uuid.New().String() + kb2.VLMConfig = types.VLMConfig{Enabled: true, ModelID: modelID} + require.NoError(t, db.Create(kb2).Error) + + count, err = repo.CountByModelID(ctx, 1, modelID) + require.NoError(t, err) + assert.Equal(t, int64(2), count) + + require.NoError(t, db.Delete(kb2).Error) + count, err = repo.CountByModelID(ctx, 1, modelID) + require.NoError(t, err) + assert.Equal(t, int64(1), count) +} + +func TestCountByModelID_CustomAgent(t *testing.T) { + ctx := context.Background() + db := setupModelUsageTestDB(t) + repo := NewCustomAgentRepository(db) + modelID := "chat-model-1" + + agent := &types.CustomAgent{ + ID: uuid.New().String(), + Name: "test-agent", + TenantID: 1, + Config: types.CustomAgentConfig{ + ModelID: modelID, + }, + } + require.NoError(t, repo.CreateAgent(ctx, agent)) + + count, err := repo.CountByModelID(ctx, 1, modelID) + require.NoError(t, err) + assert.Equal(t, int64(1), count) + + agent2 := &types.CustomAgent{ + ID: uuid.New().String(), + Name: "rerank-agent", + TenantID: 1, + Config: types.CustomAgentConfig{ + RerankModelID: modelID, + }, + } + require.NoError(t, repo.CreateAgent(ctx, agent2)) + + count, err = repo.CountByModelID(ctx, 1, modelID) + require.NoError(t, err) + assert.Equal(t, int64(2), count) + + count, err = repo.CountByModelID(ctx, 2, modelID) + require.NoError(t, err) + assert.Equal(t, int64(0), count) + + require.NoError(t, repo.DeleteAgent(ctx, agent2.ID, 1)) + count, err = repo.CountByModelID(ctx, 1, modelID) + require.NoError(t, err) + assert.Equal(t, int64(1), count) +} diff --git a/internal/application/service/knowledgebase_pr3_test.go b/internal/application/service/knowledgebase_pr3_test.go index dac79e992..a0af39f89 100644 --- a/internal/application/service/knowledgebase_pr3_test.go +++ b/internal/application/service/knowledgebase_pr3_test.go @@ -108,6 +108,9 @@ func (r *fakeKBRepo) TogglePinKnowledgeBase(_ context.Context, _ string, _ uint6 func (r *fakeKBRepo) CountByVectorStoreID(_ context.Context, _ *gorm.DB, _ uint64, _ string) (int64, error) { return 0, nil } +func (r *fakeKBRepo) CountByModelID(_ context.Context, _ uint64, _ string) (int64, error) { + return 0, nil +} func (r *fakeKBRepo) SetUserKBPin(_ context.Context, _ uint64, _ string, _ string, _ bool) (*time.Time, error) { return nil, nil } diff --git a/internal/application/service/model.go b/internal/application/service/model.go index e93470e57..78a8f9d00 100644 --- a/internal/application/service/model.go +++ b/internal/application/service/model.go @@ -3,7 +3,9 @@ package service import ( "context" "errors" + "fmt" + apperrors "github.com/Tencent/WeKnora/internal/errors" "github.com/Tencent/WeKnora/internal/logger" "github.com/Tencent/WeKnora/internal/models/asr" "github.com/Tencent/WeKnora/internal/models/chat" @@ -23,6 +25,8 @@ var ErrModelNotFound = errors.New("model not found") // modelService implements the model service interface type modelService struct { repo interfaces.ModelRepository + kbRepo interfaces.KnowledgeBaseRepository + agentRepo interfaces.CustomAgentRepository ollamaService *ollama.OllamaService pooler embedding.EmbedderPooler tenantService interfaces.TenantService @@ -30,12 +34,16 @@ type modelService struct { // NewModelService creates a new model service instance func NewModelService(repo interfaces.ModelRepository, + kbRepo interfaces.KnowledgeBaseRepository, + agentRepo interfaces.CustomAgentRepository, ollamaService *ollama.OllamaService, pooler embedding.EmbedderPooler, tenantService interfaces.TenantService, ) interfaces.ModelService { return &modelService{ repo: repo, + kbRepo: kbRepo, + agentRepo: agentRepo, ollamaService: ollamaService, pooler: pooler, tenantService: tenantService, @@ -341,9 +349,31 @@ func (s *modelService) DeleteModel(ctx context.Context, id string) error { }) return err } - if existingModel != nil && existingModel.IsBuiltin { + if existingModel == nil { + return ErrModelNotFound + } + if existingModel.IsBuiltin { logger.Warnf(ctx, "Attempted to delete builtin model: %s", id) - return errors.New("builtin models cannot be deleted") + return apperrors.NewBadRequestError("builtin models cannot be deleted") + } + + kbCount, err := s.kbRepo.CountByModelID(ctx, tenantID, id) + if err != nil { + logger.ErrorWithFields(ctx, err, map[string]interface{}{ + "model_id": id, + }) + return err + } + agentCount, err := s.agentRepo.CountByModelID(ctx, tenantID, id) + if err != nil { + logger.ErrorWithFields(ctx, err, map[string]interface{}{ + "model_id": id, + }) + return err + } + if kbCount > 0 || agentCount > 0 { + logger.Warnf(ctx, "Model %s is in use: kb=%d agent=%d", id, kbCount, agentCount) + return apperrors.NewBadRequestError(formatModelInUseMessage(kbCount, agentCount)) } // Delete model from repository @@ -582,3 +612,26 @@ func (s *modelService) GetASRModel(ctx context.Context, modelId string) (asr.ASR return sttModel, nil } + +func formatModelInUseMessage(kbCount, agentCount int64) string { + switch { + case kbCount > 0 && agentCount > 0: + return fmt.Sprintf( + "model is used by %d knowledge base(s) and %d agent(s); "+ + "reconfigure or remove those references before deleting", + kbCount, agentCount, + ) + case kbCount > 0: + return fmt.Sprintf( + "model is used by %d knowledge base(s); "+ + "reconfigure or remove those references before deleting", + kbCount, + ) + default: + return fmt.Sprintf( + "model is used by %d agent(s); "+ + "reconfigure or remove those references before deleting", + agentCount, + ) + } +} diff --git a/internal/application/service/model_delete_test.go b/internal/application/service/model_delete_test.go new file mode 100644 index 000000000..fea92be9c --- /dev/null +++ b/internal/application/service/model_delete_test.go @@ -0,0 +1,175 @@ +package service + +import ( + "context" + "testing" + "time" + + apperrors "github.com/Tencent/WeKnora/internal/errors" + "github.com/Tencent/WeKnora/internal/types" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/gorm" +) + +type stubKBRepoForModelDelete struct { + count int64 +} + +func (s *stubKBRepoForModelDelete) CreateKnowledgeBase(context.Context, *types.KnowledgeBase) error { + return nil +} +func (s *stubKBRepoForModelDelete) GetKnowledgeBaseByID(context.Context, string) (*types.KnowledgeBase, error) { + return nil, nil +} +func (s *stubKBRepoForModelDelete) GetKnowledgeBaseByIDAndTenant(context.Context, string, uint64) (*types.KnowledgeBase, error) { + return nil, nil +} +func (s *stubKBRepoForModelDelete) GetKnowledgeBaseByIDs(context.Context, []string) ([]*types.KnowledgeBase, error) { + return nil, nil +} +func (s *stubKBRepoForModelDelete) ListKnowledgeBases(context.Context) ([]*types.KnowledgeBase, error) { + return nil, nil +} +func (s *stubKBRepoForModelDelete) ListKnowledgeBasesByTenantID(context.Context, uint64) ([]*types.KnowledgeBase, error) { + return nil, nil +} +func (s *stubKBRepoForModelDelete) UpdateKnowledgeBase(context.Context, *types.KnowledgeBase) error { + return nil +} +func (s *stubKBRepoForModelDelete) DeleteKnowledgeBase(context.Context, string) error { return nil } +func (s *stubKBRepoForModelDelete) CountByVectorStoreID(context.Context, *gorm.DB, uint64, string) (int64, error) { + return 0, nil +} +func (s *stubKBRepoForModelDelete) CountByModelID(context.Context, uint64, string) (int64, error) { + return s.count, nil +} +func (s *stubKBRepoForModelDelete) SetUserKBPin(context.Context, uint64, string, string, bool) (*time.Time, error) { + return nil, nil +} +func (s *stubKBRepoForModelDelete) ListUserKBPinIDs(context.Context, uint64, string) (map[string]time.Time, error) { + return nil, nil +} + +type stubAgentRepoForModelDelete struct { + count int64 +} + +func (s *stubAgentRepoForModelDelete) CreateAgent(context.Context, *types.CustomAgent) error { + return nil +} +func (s *stubAgentRepoForModelDelete) GetAgentByID(context.Context, string, uint64) (*types.CustomAgent, error) { + return nil, nil +} +func (s *stubAgentRepoForModelDelete) ListAgentsByTenantID(context.Context, uint64) ([]*types.CustomAgent, error) { + return nil, nil +} +func (s *stubAgentRepoForModelDelete) UpdateAgent(context.Context, *types.CustomAgent) error { + return nil +} +func (s *stubAgentRepoForModelDelete) DeleteAgent(context.Context, string, uint64) error { return nil } +func (s *stubAgentRepoForModelDelete) CountByModelID(context.Context, uint64, string) (int64, error) { + return s.count, nil +} + +type stubModelRepoForDelete struct { + model *types.Model + delete func(id string) error +} + +func (s *stubModelRepoForDelete) Create(context.Context, *types.Model) error { return nil } +func (s *stubModelRepoForDelete) GetByID(_ context.Context, _ uint64, id string) (*types.Model, error) { + if s.model != nil && s.model.ID == id { + return s.model, nil + } + return nil, nil +} +func (s *stubModelRepoForDelete) List(context.Context, uint64, types.ModelType, types.ModelSource) ([]*types.Model, error) { + return nil, nil +} +func (s *stubModelRepoForDelete) Update(context.Context, *types.Model) error { return nil } +func (s *stubModelRepoForDelete) Delete(_ context.Context, _ uint64, id string) error { + if s.delete != nil { + return s.delete(id) + } + return nil +} +func (s *stubModelRepoForDelete) ClearDefaultByType(context.Context, uint, types.ModelType, string) error { + return nil +} + +func TestDeleteModel_RejectsWhenReferenced(t *testing.T) { + ctx := context.WithValue(context.Background(), types.TenantIDContextKey, uint64(1)) + modelID := "model-in-use" + + svc := NewModelService( + &stubModelRepoForDelete{model: &types.Model{ID: modelID, TenantID: 1}}, + &stubKBRepoForModelDelete{count: 1}, + &stubAgentRepoForModelDelete{count: 0}, + nil, nil, nil, + ) + + err := svc.DeleteModel(ctx, modelID) + require.Error(t, err) + appErr, ok := apperrors.IsAppError(err) + require.True(t, ok) + assert.Equal(t, apperrors.ErrBadRequest, appErr.Code) + assert.Contains(t, appErr.Message, "knowledge base") +} + +func TestDeleteModel_RejectsWhenUsedByAgent(t *testing.T) { + ctx := context.WithValue(context.Background(), types.TenantIDContextKey, uint64(1)) + modelID := "agent-model" + + svc := NewModelService( + &stubModelRepoForDelete{model: &types.Model{ID: modelID, TenantID: 1}}, + &stubKBRepoForModelDelete{count: 0}, + &stubAgentRepoForModelDelete{count: 2}, + nil, nil, nil, + ) + + err := svc.DeleteModel(ctx, modelID) + require.Error(t, err) + appErr, ok := apperrors.IsAppError(err) + require.True(t, ok) + assert.Contains(t, appErr.Message, "2 agent(s)") +} + +func TestDeleteModel_SucceedsWhenUnreferenced(t *testing.T) { + ctx := context.WithValue(context.Background(), types.TenantIDContextKey, uint64(1)) + modelID := "free-model" + deleted := false + + svc := NewModelService( + &stubModelRepoForDelete{ + model: &types.Model{ID: modelID, TenantID: 1}, + delete: func(id string) error { + assert.Equal(t, modelID, id) + deleted = true + return nil + }, + }, + &stubKBRepoForModelDelete{}, + &stubAgentRepoForModelDelete{}, + nil, nil, nil, + ) + + require.NoError(t, svc.DeleteModel(ctx, modelID)) + assert.True(t, deleted) +} + +func TestFormatModelInUseMessage(t *testing.T) { + t.Parallel() + assert.Equal(t, + "model is used by 1 knowledge base(s); reconfigure or remove those references before deleting", + formatModelInUseMessage(1, 0), + ) + assert.Equal(t, + "model is used by 2 agent(s); reconfigure or remove those references before deleting", + formatModelInUseMessage(0, 2), + ) + assert.Equal(t, + "model is used by 1 knowledge base(s) and 1 agent(s); reconfigure or remove those references before deleting", + formatModelInUseMessage(1, 1), + ) +} diff --git a/internal/application/service/vectorstore_test.go b/internal/application/service/vectorstore_test.go index 40ad18331..3a7fd91aa 100644 --- a/internal/application/service/vectorstore_test.go +++ b/internal/application/service/vectorstore_test.go @@ -921,6 +921,9 @@ func (r *realKBRepo) CountByVectorStoreID(ctx context.Context, db *gorm.DB, tena Count(&count).Error return count, err } +func (r *realKBRepo) CountByModelID(_ context.Context, _ uint64, _ string) (int64, error) { + return 0, nil +} // The remaining methods are not called by the tested code paths; declare them // so realKBRepo satisfies interfaces.KnowledgeBaseRepository. diff --git a/internal/handler/model.go b/internal/handler/model.go index 854c99140..0292aefa6 100644 --- a/internal/handler/model.go +++ b/internal/handler/model.go @@ -346,6 +346,10 @@ func (h *ModelHandler) DeleteModel(c *gin.Context) { c.Error(errors.NewNotFoundError("Model not found")) return } + if appErr, ok := errors.IsAppError(err); ok { + c.Error(appErr) + return + } logger.ErrorWithFields(ctx, err, nil) c.Error(errors.NewInternalServerError(err.Error())) return diff --git a/internal/types/interfaces/custom_agent.go b/internal/types/interfaces/custom_agent.go index cd04182c4..7dd7ed960 100644 --- a/internal/types/interfaces/custom_agent.go +++ b/internal/types/interfaces/custom_agent.go @@ -126,4 +126,8 @@ type CustomAgentRepository interface { // Returns: // - Possible errors such as record not existing, database errors, etc. DeleteAgent(ctx context.Context, id string, tenantID uint64) error + + // CountByModelID counts active agents in the tenant whose config references + // the given model ID (chat, rerank, VLM, ASR, query-understand, etc.). + CountByModelID(ctx context.Context, tenantID uint64, modelID string) (int64, error) } diff --git a/internal/types/interfaces/knowledgebase.go b/internal/types/interfaces/knowledgebase.go index 322d41e41..b5ae76706 100644 --- a/internal/types/interfaces/knowledgebase.go +++ b/internal/types/interfaces/knowledgebase.go @@ -216,6 +216,10 @@ type KnowledgeBaseRepository interface { // scope on KnowledgeBase; implementations MUST NOT add an explicit // `deleted_at IS NULL` predicate (avoids divergence with the auto-scope). CountByVectorStoreID(ctx context.Context, db *gorm.DB, tenantID uint64, storeID string) (int64, error) + + // CountByModelID counts active KBs in the tenant that reference the given + // model ID in any model-binding field (embedding, summary, VLM, ASR, etc.). + CountByModelID(ctx context.Context, tenantID uint64, modelID string) (int64, error) // SetUserKBPin inserts or removes a row in user_kb_pins for the given // (tenant, user, kb) triple. Returns the resulting pinned_at (nil when // pinned=false) and an error. The tenant_id is captured to support