//go:build integration package seed import ( "context" "fmt" "os" "sync" "testing" "time" "github.com/stretchr/testify/require" "golang.org/x/crypto/bcrypt" "gorm.io/driver/postgres" "gorm.io/gorm" "gorm.io/gorm/logger" "senlinai-agent/backend/internal/models" ) type demoSeedRunContextKey struct{} type demoSeedRunMarker string const ( firstDemoSeedRun demoSeedRunMarker = "first" secondDemoSeedRun demoSeedRunMarker = "second" ) func TestPostgresDemoSeedSerializesConcurrentRunsForSameOwner(t *testing.T) { dsn := os.Getenv("TEST_DATABASE_URL") require.NotEmpty(t, dsn, "TEST_DATABASE_URL is required for integration tests and must point to an isolated database") database, err := gorm.Open(postgres.Open(dsn), &gorm.Config{ TranslateError: true, Logger: logger.Default.LogMode(logger.Silent), }) require.NoError(t, err) require.NoError(t, models.AutoMigrate(database)) suffix := fmt.Sprint(time.Now().UnixNano()) email := "postgres-seed-" + suffix + "@example.com" passwordHash, err := bcrypt.GenerateFromPassword([]byte("original-password"), bcrypt.MinCost) require.NoError(t, err) user := models.SaUser{Email: email, DisplayName: "保留用户名称", PasswordHash: string(passwordHash), Role: "user"} require.NoError(t, database.Create(&user).Error) projects := []models.SaProject{ {OwnerID: user.ID, Name: "项目 A1", Identifier: "A1"}, {OwnerID: user.ID, Name: "项目 A2", Identifier: "A2"}, {OwnerID: user.ID, Name: "数据中台", Identifier: "DATA"}, {OwnerID: user.ID, Name: "运营自动化", Identifier: "AUTO"}, } for index := range projects { require.NoError(t, database.Create(&projects[index]).Error) } uiTag := models.SaTag{ProjectID: projects[0].ID, Name: "UI"} require.NoError(t, database.Create(&uiTag).Error) require.NoError(t, database.Create(&models.SaTask{ ProjectID: projects[0].ID, CreatedBy: user.ID, TagID: &uiTag.ID, Title: "智能报表导出功能", Status: "open", }).Error) t.Cleanup(func() { cleanupPostgresDemoSeed(database, user.ID) }) firstAtChildQuery := make(chan struct{}, 1) secondOwnerQueryAttempted := make(chan struct{}, 1) secondEnteredChildQuery := make(chan struct{}, 1) releaseFirst := make(chan struct{}) var firstChildOnce sync.Once var releaseOnce sync.Once releaseFirstRun := func() { releaseOnce.Do(func() { close(releaseFirst) }) } callbackName := "test:coordinate_concurrent_seed_" + suffix require.NoError(t, database.Callback().Query().Before("gorm:query").Register(callbackName, func(tx *gorm.DB) { if tx.Statement.Schema == nil { return } marker, ok := tx.Statement.Context.Value(demoSeedRunContextKey{}).(demoSeedRunMarker) if !ok { return } if tx.Statement.Schema.Table == (models.SaUser{}).TableName() { if _, locking := tx.Statement.Clauses["FOR"]; locking && marker == secondDemoSeedRun { signalDemoSeedStage(secondOwnerQueryAttempted) } return } switch marker { case firstDemoSeedRun: firstChildOnce.Do(func() { signalDemoSeedStage(firstAtChildQuery) <-releaseFirst }) case secondDemoSeedRun: signalDemoSeedStage(secondEnteredChildQuery) } })) ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) var wait sync.WaitGroup firstResult := make(chan error, 1) secondResult := make(chan error, 1) t.Cleanup(func() { cancel() releaseFirstRun() wait.Wait() database.Callback().Query().Remove(callbackName) }) wait.Add(1) go func() { defer wait.Done() firstContext := context.WithValue(ctx, demoSeedRunContextKey{}, firstDemoSeedRun) _, err := Demo(database.WithContext(firstContext), DemoOptions{Email: email, DisplayName: "不应覆盖", Password: "new-password"}) firstResult <- err }() waitForDemoSeedStage(t, ctx, firstAtChildQuery, "first Demo to hold the owner lock and reach its first child query") wait.Add(1) go func() { defer wait.Done() secondContext := context.WithValue(ctx, demoSeedRunContextKey{}, secondDemoSeedRun) _, err := Demo(database.WithContext(secondContext), DemoOptions{Email: email, DisplayName: "不应覆盖", Password: "new-password"}) secondResult <- err }() waitForDemoSeedStage(t, ctx, secondOwnerQueryAttempted, "second Demo to attempt the owner FOR UPDATE query") select { case <-secondEnteredChildQuery: t.Fatal("second Demo entered a child query before the first released the owner lock") case err := <-secondResult: t.Fatalf("second Demo returned before the first released the owner lock: %v", err) case <-time.After(200 * time.Millisecond): case <-ctx.Done(): t.Fatal("timed out while proving the second Demo is blocked on the owner lock") } releaseFirstRun() require.NoError(t, waitForDemoSeedResult(t, ctx, firstResult, "first Demo result")) require.NoError(t, waitForDemoSeedResult(t, ctx, secondResult, "second Demo result")) var storedUser models.SaUser require.NoError(t, database.Where("email = ?", email).First(&storedUser).Error) require.Equal(t, "保留用户名称", storedUser.DisplayName) require.NoError(t, bcrypt.CompareHashAndPassword([]byte(storedUser.PasswordHash), []byte("original-password"))) require.Error(t, bcrypt.CompareHashAndPassword([]byte(storedUser.PasswordHash), []byte("new-password"))) projectIDs := []uint{projects[0].ID, projects[1].ID, projects[2].ID, projects[3].ID} assertPostgresDemoCount(t, database, &models.SaProject{}, "owner_id = ?", []any{user.ID}, 4) assertPostgresDemoCount(t, database, &models.SaTag{}, "project_id = ?", []any{projects[0].ID}, 4) assertPostgresDemoCount(t, database, &models.SaTask{}, "project_id IN ?", []any{projectIDs}, 6) assertPostgresDemoCount(t, database, &models.SaInboxItem{}, "project_id IN ?", []any{projectIDs}, 7) assertPostgresDemoCount(t, database, &models.SaAISession{}, "project_id = ?", []any{projects[0].ID}, 4) assertPostgresDemoCount(t, database, &models.SaDocumentTree{}, "project_id = ?", []any{projects[0].ID}, 2) assertPostgresDemoCount(t, database, &models.SaCronPlan{}, "project_id = ?", []any{projects[0].ID}, 3) assertPostgresDemoCount(t, database, &models.SaProjectChannel{}, "project_id = ?", []any{projects[0].ID}, 2) } func signalDemoSeedStage(stage chan<- struct{}) { select { case stage <- struct{}{}: default: } } func waitForDemoSeedStage(t *testing.T, ctx context.Context, stage <-chan struct{}, description string) { t.Helper() select { case <-stage: case <-ctx.Done(): t.Fatalf("timed out waiting for %s", description) } } func waitForDemoSeedResult(t *testing.T, ctx context.Context, result <-chan error, description string) error { t.Helper() select { case err := <-result: return err case <-ctx.Done(): t.Fatalf("timed out waiting for %s", description) return ctx.Err() } } func assertPostgresDemoCount(t *testing.T, database *gorm.DB, model any, where string, args []any, expected int64) { t.Helper() var count int64 require.NoError(t, database.Model(model).Where(where, args...).Count(&count).Error) require.Equal(t, expected, count) } func cleanupPostgresDemoSeed(database *gorm.DB, userID uint) { var projects []models.SaProject if database.Where("owner_id = ?", userID).Find(&projects).Error != nil { return } projectIDs := make([]uint, 0, len(projects)) for _, project := range projects { projectIDs = append(projectIDs, project.ID) } if len(projectIDs) != 0 { var tasks []models.SaTask database.Where("project_id IN ?", projectIDs).Find(&tasks) taskIDs := make([]uint, 0, len(tasks)) for _, task := range tasks { taskIDs = append(taskIDs, task.ID) } if len(taskIDs) != 0 { database.Where("task_id IN ?", taskIDs).Delete(&models.SaTaskShare{}) } var inboxItems []models.SaInboxItem database.Where("project_id IN ?", projectIDs).Find(&inboxItems) inboxIDs := make([]uint, 0, len(inboxItems)) for _, item := range inboxItems { inboxIDs = append(inboxIDs, item.ID) } if len(inboxIDs) != 0 { database.Where("inbox_item_id IN ?", inboxIDs).Delete(&models.SaInboxSuggestion{}) } database.Where("project_id IN ?", projectIDs).Delete(&models.SaProjectEvent{}) database.Where("project_id IN ?", projectIDs).Delete(&models.SaTask{}) database.Where("project_id IN ?", projectIDs).Delete(&models.SaDocumentTree{}) database.Where("project_id IN ?", projectIDs).Delete(&models.SaAISession{}) database.Where("project_id IN ?", projectIDs).Delete(&models.SaTag{}) database.Where("project_id IN ?", projectIDs).Delete(&models.SaProjectChannel{}) database.Where("project_id IN ?", projectIDs).Delete(&models.SaCronPlan{}) database.Where("project_id IN ?", projectIDs).Delete(&models.SaInboxItem{}) database.Where("id IN ?", projectIDs).Delete(&models.SaProject{}) } database.Delete(&models.SaUser{}, userID) }