From 7f71c971ebc49ec64d897c44ab4384226cbdb921 Mon Sep 17 00:00:00 2001 From: yanweidong Date: Tue, 21 Jul 2026 20:58:17 +0800 Subject: [PATCH] test: strengthen concurrent seed coverage --- backend/internal/seed/demo_postgres_test.go | 141 +++++++++++++------- backend/internal/seed/demo_test.go | 77 ++++++++--- 2 files changed, 152 insertions(+), 66 deletions(-) diff --git a/backend/internal/seed/demo_postgres_test.go b/backend/internal/seed/demo_postgres_test.go index fcebc6b..acbb7b3 100644 --- a/backend/internal/seed/demo_postgres_test.go +++ b/backend/internal/seed/demo_postgres_test.go @@ -5,7 +5,6 @@ import ( "fmt" "os" "sync" - "sync/atomic" "testing" "time" @@ -17,6 +16,15 @@ import ( "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("DATABASE_URL") if dsn == "" { @@ -50,71 +58,83 @@ func TestPostgresDemoSeedSerializesConcurrentRunsForSameOwner(t *testing.T) { ProjectID: projects[0].ID, CreatedBy: user.ID, TagID: &uiTag.ID, Title: "智能报表导出功能", Status: "open", }).Error) + t.Cleanup(func() { cleanupPostgresDemoSeed(database, user.ID) }) - var ownerLockObserved atomic.Bool - var tagMisses atomic.Int32 - releaseTagMisses := make(chan struct{}) + 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().After("gorm:query").Register(callbackName, func(tx *gorm.DB) { + 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.SenlinAgentUser{}).TableName() { - if _, ok := tx.Statement.Clauses["FOR"]; ok { - ownerLockObserved.Store(true) + if _, locking := tx.Statement.Clauses["FOR"]; locking && marker == secondDemoSeedRun { + signalDemoSeedStage(secondOwnerQueryAttempted) } 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")) + 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) - 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{}) + wait.Add(1) go func() { - wait.Wait() - close(done) + defer wait.Done() + firstContext := context.WithValue(ctx, demoSeedRunContextKey{}, firstDemoSeedRun) + _, err := Demo(database.WithContext(firstContext), DemoOptions{Email: email, DisplayName: "不应覆盖", Password: "new-password"}) + firstResult <- err }() - 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) - } + 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")) - 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) @@ -133,6 +153,33 @@ func TestPostgresDemoSeedSerializesConcurrentRunsForSameOwner(t *testing.T) { assertPostgresDemoCount(t, database, &models.SenlinAgentProjectChannel{}, "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 diff --git a/backend/internal/seed/demo_test.go b/backend/internal/seed/demo_test.go index 4c3c134..af5b0b0 100644 --- a/backend/internal/seed/demo_test.go +++ b/backend/internal/seed/demo_test.go @@ -2,6 +2,7 @@ package seed import ( "fmt" + "strings" "sync/atomic" "testing" "time" @@ -134,27 +135,65 @@ func TestDemoSeedAcceptsAConcurrentWinnerWithoutOverwritingUser(t *testing.T) { } 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) }) + for _, sqlState := range []string{"40001", "40P01"} { + t.Run(sqlState, func(t *testing.T) { + database := newTestDB(t) + email := "retry-" + strings.ToLower(sqlState) + "@senlin.ai" + var injected atomic.Bool + var sawProjectWrite atomic.Bool + var failedOwnerIdentity string + var failedProjectIdentity string + callbackName := "test:fail_first_seed_project_" + sqlState + require.NoError(t, database.Callback().Create().After("gorm:create").Register(callbackName, func(tx *gorm.DB) { + if tx.Error != nil || tx.Statement.Schema == nil || tx.Statement.Schema.Table != (models.SenlinAgentProject{}).TableName() { + return + } + if !injected.CompareAndSwap(false, true) { + return + } + project, ok := tx.Statement.Dest.(*models.SenlinAgentProject) + if !ok { + tx.AddError(fmt.Errorf("unexpected project create destination %T", tx.Statement.Dest)) + return + } + if tx.RowsAffected == 1 && project.ID != 0 && project.Identity != "" && project.OwnerIdentity != "" { + sawProjectWrite.Store(true) + failedOwnerIdentity = project.OwnerIdentity + failedProjectIdentity = project.Identity + } + tx.AddError(testSQLStateError{state: sqlState}) + })) + t.Cleanup(func() { database.Callback().Create().Remove(callbackName) }) - result, err := Demo(database, DemoOptions{Email: "retry@senlin.ai", DisplayName: "重试用户", Password: "password123"}) + result, err := Demo(database, DemoOptions{Email: email, 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) + require.NoError(t, err) + require.True(t, injected.Load()) + require.True(t, sawProjectWrite.Load(), "failed transaction must contain the first inserted project") + require.NotEqual(t, failedOwnerIdentity, result.User.Identity, "retry must recreate the rolled-back user") + require.NotEqual(t, failedProjectIdentity, result.Projects[0].Identity, "retry must recreate the rolled-back project") + require.Equal(t, email, result.User.Email) + require.Len(t, result.Projects, 4) + projectIDs := []uint{result.Projects[0].ID, result.Projects[1].ID, result.Projects[2].ID, result.Projects[3].ID} + requireDemoModelCount(t, database, &models.SenlinAgentUser{}, "email = ?", []any{email}, 1) + requireDemoModelCount(t, database, &models.SenlinAgentProject{}, "owner_id = ?", []any{result.User.ID}, 4) + requireDemoModelCount(t, database, &models.SenlinAgentTag{}, "project_id = ?", []any{result.Projects[0].ID}, 4) + requireDemoModelCount(t, database, &models.SenlinAgentTask{}, "project_id IN ?", []any{projectIDs}, 6) + requireDemoModelCount(t, database, &models.SenlinAgentInboxItem{}, "project_id IN ?", []any{projectIDs}, 7) + requireDemoModelCount(t, database, &models.SenlinAgentAISession{}, "project_id = ?", []any{result.Projects[0].ID}, 4) + requireDemoModelCount(t, database, &models.SenlinAgentNote{}, "project_id = ?", []any{result.Projects[0].ID}, 2) + requireDemoModelCount(t, database, &models.SenlinAgentSource{}, "project_id = ?", []any{result.Projects[0].ID}, 2) + requireDemoModelCount(t, database, &models.SenlinAgentCronPlan{}, "project_id = ?", []any{result.Projects[0].ID}, 3) + requireDemoModelCount(t, database, &models.SenlinAgentProjectChannel{}, "project_id = ?", []any{result.Projects[0].ID}, 2) + }) + } +} + +func requireDemoModelCount(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) } type testSQLStateError struct {