diff --git a/backend/internal/seed/demo.go b/backend/internal/seed/demo.go index 06a9155..19aa58b 100644 --- a/backend/internal/seed/demo.go +++ b/backend/internal/seed/demo.go @@ -1,14 +1,18 @@ package seed import ( + "errors" "strings" "time" "golang.org/x/crypto/bcrypt" "gorm.io/gorm" + "gorm.io/gorm/clause" "senlinai-agent/backend/internal/models" ) +const demoTransactionAttempts = 3 + type DemoOptions struct { Email string DisplayName string @@ -35,43 +39,68 @@ func Demo(database *gorm.DB, options DemoOptions) (DemoResult, error) { } var result DemoResult - err := database.Transaction(func(tx *gorm.DB) error { - user, err := upsertUser(tx, email, displayName, password) - if err != nil { - return err - } - result.User = user + var err error + for attempt := 0; attempt < demoTransactionAttempts; attempt++ { + result = DemoResult{} + err = database.Transaction(func(tx *gorm.DB) error { + user, err := upsertAndLockUser(tx, email, displayName, password) + if err != nil { + return err + } + result.User = user - projects, err := seedProjects(tx, user.ID) - if err != nil { - return err + projects, err := seedProjects(tx, user.ID) + if err != nil { + return err + } + result.Projects = projects + return nil + }) + if err == nil || !isRetryableDemoTransactionError(err) || attempt == demoTransactionAttempts-1 { + return result, err } - result.Projects = projects - return nil - }) + // 重试发生在事务完全回滚之后,避免持锁等待自身的新事务。 + time.Sleep(time.Duration(attempt+1) * 10 * time.Millisecond) + } return result, err } -func upsertUser(tx *gorm.DB, email string, displayName string, password string) (models.SenlinAgentUser, error) { - var user models.SenlinAgentUser - err := tx.Where("email = ?", email).First(&user).Error - if err == nil { - return user, nil - } - if err != gorm.ErrRecordNotFound { - return user, err - } +func upsertAndLockUser(tx *gorm.DB, email string, displayName string, password string) (models.SenlinAgentUser, error) { hash, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost) if err != nil { - return user, err + return models.SenlinAgentUser{}, err } - user = models.SenlinAgentUser{ + candidate := models.SenlinAgentUser{ Email: email, DisplayName: displayName, PasswordHash: string(hash), Role: "user", } - return user, tx.Create(&user).Error + if err := tx.Clauses(clause.OnConflict{ + Columns: []clause.Column{{Name: "email"}}, + DoNothing: true, + }).Create(&candidate).Error; err != nil { + return models.SenlinAgentUser{}, err + } + + // PostgreSQL 会先等待并发的同 email INSERT 提交,再在 READ COMMITTED 的下一条语句中看见该用户。 + // 锁定 owner 后,项目、标签和子对象的 find-or-create 在不同进程间也会按同一用户串行执行。 + var user models.SenlinAgentUser + err = tx.Clauses(clause.Locking{Strength: "UPDATE"}).Where("email = ?", email).First(&user).Error + return user, err +} + +func isRetryableDemoTransactionError(err error) bool { + var sqlState interface{ SQLState() string } + if !errors.As(err, &sqlState) { + return false + } + switch sqlState.SQLState() { + case "40001", "40P01": + return true + default: + return false + } } func seedProjects(tx *gorm.DB, ownerID uint) ([]models.SenlinAgentProject, error) { diff --git a/backend/internal/seed/demo_postgres_test.go b/backend/internal/seed/demo_postgres_test.go new file mode 100644 index 0000000..fcebc6b --- /dev/null +++ b/backend/internal/seed/demo_postgres_test.go @@ -0,0 +1,183 @@ +package seed + +import ( + "context" + "fmt" + "os" + "sync" + "sync/atomic" + "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" +) + +func TestPostgresDemoSeedSerializesConcurrentRunsForSameOwner(t *testing.T) { + dsn := os.Getenv("DATABASE_URL") + if dsn == "" { + t.Skip("DATABASE_URL is not configured; skipping PostgreSQL seed concurrency test") + } + 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.SenlinAgentUser{Email: email, DisplayName: "保留用户名称", PasswordHash: string(passwordHash), Role: "user"} + require.NoError(t, database.Create(&user).Error) + projects := []models.SenlinAgentProject{ + {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.SenlinAgentTag{ProjectID: projects[0].ID, Name: "UI"} + require.NoError(t, database.Create(&uiTag).Error) + require.NoError(t, database.Create(&models.SenlinAgentTask{ + ProjectID: projects[0].ID, CreatedBy: user.ID, TagID: &uiTag.ID, + Title: "智能报表导出功能", Status: "open", + }).Error) + + var ownerLockObserved atomic.Bool + var tagMisses atomic.Int32 + releaseTagMisses := make(chan struct{}) + var releaseOnce sync.Once + callbackName := "test:coordinate_concurrent_seed_" + suffix + require.NoError(t, database.Callback().Query().After("gorm:query").Register(callbackName, func(tx *gorm.DB) { + if tx.Statement.Schema == nil { + return + } + if tx.Statement.Schema.Table == (models.SenlinAgentUser{}).TableName() { + if _, ok := tx.Statement.Clauses["FOR"]; ok { + ownerLockObserved.Store(true) + } + return + } + if ownerLockObserved.Load() || tx.Statement.Schema.Table != (models.SenlinAgentTag{}).TableName() || tx.RowsAffected != 0 { + return + } + if tagMisses.Add(1) >= 2 { + releaseOnce.Do(func() { close(releaseTagMisses) }) + } + select { + case <-releaseTagMisses: + case <-time.After(5 * time.Second): + tx.AddError(fmt.Errorf("timed out coordinating concurrent tag misses")) + } + })) + t.Cleanup(func() { + database.Callback().Query().Remove(callbackName) + cleanupPostgresDemoSeed(database, user.ID) + }) + + start := make(chan struct{}) + results := make(chan error, 2) + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + var wait sync.WaitGroup + for range 2 { + wait.Add(1) + go func() { + defer wait.Done() + <-start + _, err := Demo(database.WithContext(ctx), DemoOptions{Email: email, DisplayName: "不应覆盖", Password: "new-password"}) + results <- err + }() + } + close(start) + done := make(chan struct{}) + go func() { + wait.Wait() + close(done) + }() + select { + case <-done: + case <-ctx.Done(): + <-done + t.Fatal("timed out waiting for concurrent Demo calls") + } + close(results) + for err := range results { + require.NoError(t, err) + } + + require.True(t, ownerLockObserved.Load(), "seed must lock the owner user row before creating child objects") + var storedUser models.SenlinAgentUser + 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.SenlinAgentProject{}, "owner_id = ?", []any{user.ID}, 4) + assertPostgresDemoCount(t, database, &models.SenlinAgentTag{}, "project_id = ?", []any{projects[0].ID}, 4) + assertPostgresDemoCount(t, database, &models.SenlinAgentTask{}, "project_id IN ?", []any{projectIDs}, 6) + assertPostgresDemoCount(t, database, &models.SenlinAgentInboxItem{}, "project_id IN ?", []any{projectIDs}, 7) + assertPostgresDemoCount(t, database, &models.SenlinAgentAISession{}, "project_id = ?", []any{projects[0].ID}, 4) + assertPostgresDemoCount(t, database, &models.SenlinAgentNote{}, "project_id = ?", []any{projects[0].ID}, 2) + assertPostgresDemoCount(t, database, &models.SenlinAgentSource{}, "project_id = ?", []any{projects[0].ID}, 2) + assertPostgresDemoCount(t, database, &models.SenlinAgentCronPlan{}, "project_id = ?", []any{projects[0].ID}, 3) + assertPostgresDemoCount(t, database, &models.SenlinAgentProjectChannel{}, "project_id = ?", []any{projects[0].ID}, 2) +} + +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.SenlinAgentProject + 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.SenlinAgentTask + 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.SenlinAgentTaskShare{}) + } + var inboxItems []models.SenlinAgentInboxItem + 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.SenlinAgentInboxSuggestion{}) + } + database.Where("project_id IN ?", projectIDs).Delete(&models.SenlinAgentProjectEvent{}) + database.Where("project_id IN ?", projectIDs).Delete(&models.SenlinAgentTask{}) + database.Where("project_id IN ?", projectIDs).Delete(&models.SenlinAgentNote{}) + database.Where("project_id IN ?", projectIDs).Delete(&models.SenlinAgentSource{}) + database.Where("project_id IN ?", projectIDs).Delete(&models.SenlinAgentAISession{}) + database.Where("project_id IN ?", projectIDs).Delete(&models.SenlinAgentTag{}) + database.Where("project_id IN ?", projectIDs).Delete(&models.SenlinAgentProjectChannel{}) + database.Where("project_id IN ?", projectIDs).Delete(&models.SenlinAgentCronPlan{}) + database.Where("project_id IN ?", projectIDs).Delete(&models.SenlinAgentInboxItem{}) + database.Where("id IN ?", projectIDs).Delete(&models.SenlinAgentProject{}) + } + database.Delete(&models.SenlinAgentUser{}, userID) +} diff --git a/backend/internal/seed/demo_test.go b/backend/internal/seed/demo_test.go index aaa0f8c..4c3c134 100644 --- a/backend/internal/seed/demo_test.go +++ b/backend/internal/seed/demo_test.go @@ -2,10 +2,13 @@ package seed import ( "fmt" + "sync/atomic" "testing" + "time" "github.com/glebarez/sqlite" "github.com/stretchr/testify/require" + "golang.org/x/crypto/bcrypt" "gorm.io/gorm" "gorm.io/gorm/logger" "senlinai-agent/backend/internal/logic/auth" @@ -87,6 +90,7 @@ func TestDemoSeedUsesIdentifierAsStableProjectKey(t *testing.T) { require.NoError(t, err) require.Equal(t, original.ID, second.Projects[0].ID) require.Equal(t, "A1", second.Projects[0].Identifier) + require.Equal(t, "本地重命名项目", second.Projects[0].Name) var projectCount int64 require.NoError(t, database.Model(&models.SenlinAgentProject{}). Where("owner_id = ? AND identifier = ?", first.User.ID, "A1"). @@ -94,6 +98,77 @@ func TestDemoSeedUsesIdentifierAsStableProjectKey(t *testing.T) { require.Equal(t, int64(1), projectCount) } +func TestDemoSeedAcceptsAConcurrentWinnerWithoutOverwritingUser(t *testing.T) { + database := newTestDB(t) + const email = "seed-race@senlin.ai" + winnerHash, err := bcrypt.GenerateFromPassword([]byte("winner-password"), bcrypt.MinCost) + require.NoError(t, err) + + var injected atomic.Bool + callbackName := "test:inject_concurrent_seed_user" + require.NoError(t, database.Callback().Create().Before("gorm:create").Register(callbackName, func(tx *gorm.DB) { + if tx.Statement.Schema == nil || tx.Statement.Schema.Table != (models.SenlinAgentUser{}).TableName() { + return + } + if !injected.CompareAndSwap(false, true) { + return + } + now := time.Now().UTC() + tx.AddError(tx.Exec( + "INSERT INTO senlin_agent_users (identity, email, display_name, password_hash, role, created_at, updated_at) VALUES (?, ?, ?, ?, ?, ?, ?)", + "00000000-0000-7000-8000-000000000001", email, "并发胜出用户", string(winnerHash), "user", now, now, + ).Error) + })) + t.Cleanup(func() { database.Callback().Create().Remove(callbackName) }) + + result, err := Demo(database, DemoOptions{Email: email, DisplayName: "不应覆盖", Password: "loser-password"}) + + require.NoError(t, err) + require.True(t, injected.Load()) + require.Equal(t, "并发胜出用户", result.User.DisplayName) + require.NoError(t, bcrypt.CompareHashAndPassword([]byte(result.User.PasswordHash), []byte("winner-password"))) + require.Error(t, bcrypt.CompareHashAndPassword([]byte(result.User.PasswordHash), []byte("loser-password"))) + var userCount int64 + require.NoError(t, database.Model(&models.SenlinAgentUser{}).Where("email = ?", email).Count(&userCount).Error) + require.Equal(t, int64(1), userCount) +} + +func TestDemoSeedRetriesSerializationFailureOutsideTheTransaction(t *testing.T) { + database := newTestDB(t) + var failures atomic.Int32 + callbackName := "test:fail_first_seed_user_create" + require.NoError(t, database.Callback().Create().Before("gorm:create").Register(callbackName, func(tx *gorm.DB) { + if tx.Statement.Schema == nil || tx.Statement.Schema.Table != (models.SenlinAgentUser{}).TableName() { + return + } + if failures.Add(1) == 1 { + tx.AddError(testSQLStateError{state: "40001"}) + } + })) + t.Cleanup(func() { database.Callback().Create().Remove(callbackName) }) + + result, err := Demo(database, DemoOptions{Email: "retry@senlin.ai", DisplayName: "重试用户", Password: "password123"}) + + require.NoError(t, err) + require.Equal(t, int32(2), failures.Load()) + require.Equal(t, "retry@senlin.ai", result.User.Email) + var userCount int64 + require.NoError(t, database.Model(&models.SenlinAgentUser{}).Where("email = ?", result.User.Email).Count(&userCount).Error) + require.Equal(t, int64(1), userCount) +} + +type testSQLStateError struct { + state string +} + +func (err testSQLStateError) Error() string { + return "test PostgreSQL transaction failure " + err.state +} + +func (err testSQLStateError) SQLState() string { + return err.state +} + func TestDemoSeedKeepsSameNamedTagsScopedToTheirProjects(t *testing.T) { database := newTestDB(t)