test: strengthen concurrent seed coverage

This commit is contained in:
2026-07-21 20:58:17 +08:00
parent d9440b5334
commit 7f71c971eb
2 changed files with 152 additions and 66 deletions

View File

@@ -5,7 +5,6 @@ import (
"fmt" "fmt"
"os" "os"
"sync" "sync"
"sync/atomic"
"testing" "testing"
"time" "time"
@@ -17,6 +16,15 @@ import (
"senlinai-agent/backend/internal/models" "senlinai-agent/backend/internal/models"
) )
type demoSeedRunContextKey struct{}
type demoSeedRunMarker string
const (
firstDemoSeedRun demoSeedRunMarker = "first"
secondDemoSeedRun demoSeedRunMarker = "second"
)
func TestPostgresDemoSeedSerializesConcurrentRunsForSameOwner(t *testing.T) { func TestPostgresDemoSeedSerializesConcurrentRunsForSameOwner(t *testing.T) {
dsn := os.Getenv("DATABASE_URL") dsn := os.Getenv("DATABASE_URL")
if dsn == "" { if dsn == "" {
@@ -50,71 +58,83 @@ func TestPostgresDemoSeedSerializesConcurrentRunsForSameOwner(t *testing.T) {
ProjectID: projects[0].ID, CreatedBy: user.ID, TagID: &uiTag.ID, ProjectID: projects[0].ID, CreatedBy: user.ID, TagID: &uiTag.ID,
Title: "智能报表导出功能", Status: "open", Title: "智能报表导出功能", Status: "open",
}).Error) }).Error)
t.Cleanup(func() { cleanupPostgresDemoSeed(database, user.ID) })
var ownerLockObserved atomic.Bool firstAtChildQuery := make(chan struct{}, 1)
var tagMisses atomic.Int32 secondOwnerQueryAttempted := make(chan struct{}, 1)
releaseTagMisses := make(chan struct{}) secondEnteredChildQuery := make(chan struct{}, 1)
releaseFirst := make(chan struct{})
var firstChildOnce sync.Once
var releaseOnce sync.Once var releaseOnce sync.Once
releaseFirstRun := func() { releaseOnce.Do(func() { close(releaseFirst) }) }
callbackName := "test:coordinate_concurrent_seed_" + suffix 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 { if tx.Statement.Schema == nil {
return return
} }
marker, ok := tx.Statement.Context.Value(demoSeedRunContextKey{}).(demoSeedRunMarker)
if !ok {
return
}
if tx.Statement.Schema.Table == (models.SenlinAgentUser{}).TableName() { if tx.Statement.Schema.Table == (models.SenlinAgentUser{}).TableName() {
if _, ok := tx.Statement.Clauses["FOR"]; ok { if _, locking := tx.Statement.Clauses["FOR"]; locking && marker == secondDemoSeedRun {
ownerLockObserved.Store(true) signalDemoSeedStage(secondOwnerQueryAttempted)
} }
return return
} }
if ownerLockObserved.Load() || tx.Statement.Schema.Table != (models.SenlinAgentTag{}).TableName() || tx.RowsAffected != 0 { switch marker {
return case firstDemoSeedRun:
} firstChildOnce.Do(func() {
if tagMisses.Add(1) >= 2 { signalDemoSeedStage(firstAtChildQuery)
releaseOnce.Do(func() { close(releaseTagMisses) }) <-releaseFirst
} })
select { case secondDemoSeedRun:
case <-releaseTagMisses: signalDemoSeedStage(secondEnteredChildQuery)
case <-time.After(5 * time.Second):
tx.AddError(fmt.Errorf("timed out coordinating concurrent tag misses"))
} }
})) }))
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() { t.Cleanup(func() {
cancel()
releaseFirstRun()
wait.Wait()
database.Callback().Query().Remove(callbackName) 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) wait.Add(1)
go func() { go func() {
defer wait.Done() defer wait.Done()
<-start firstContext := context.WithValue(ctx, demoSeedRunContextKey{}, firstDemoSeedRun)
_, err := Demo(database.WithContext(ctx), DemoOptions{Email: email, DisplayName: "不应覆盖", Password: "new-password"}) _, err := Demo(database.WithContext(firstContext), DemoOptions{Email: email, DisplayName: "不应覆盖", Password: "new-password"})
results <- err firstResult <- err
}() }()
} waitForDemoSeedStage(t, ctx, firstAtChildQuery, "first Demo to hold the owner lock and reach its first child query")
close(start)
done := make(chan struct{}) wait.Add(1)
go func() { go func() {
wait.Wait() defer wait.Done()
close(done) secondContext := context.WithValue(ctx, demoSeedRunContextKey{}, secondDemoSeedRun)
}() _, err := Demo(database.WithContext(secondContext), DemoOptions{Email: email, DisplayName: "不应覆盖", Password: "new-password"})
select { secondResult <- err
case <-done: }()
case <-ctx.Done(): waitForDemoSeedStage(t, ctx, secondOwnerQueryAttempted, "second Demo to attempt the owner FOR UPDATE query")
<-done
t.Fatal("timed out waiting for concurrent Demo calls") select {
} case <-secondEnteredChildQuery:
close(results) t.Fatal("second Demo entered a child query before the first released the owner lock")
for err := range results { case err := <-secondResult:
require.NoError(t, err) 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 var storedUser models.SenlinAgentUser
require.NoError(t, database.Where("email = ?", email).First(&storedUser).Error) require.NoError(t, database.Where("email = ?", email).First(&storedUser).Error)
require.Equal(t, "保留用户名称", storedUser.DisplayName) 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) 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) { func assertPostgresDemoCount(t *testing.T, database *gorm.DB, model any, where string, args []any, expected int64) {
t.Helper() t.Helper()
var count int64 var count int64

View File

@@ -2,6 +2,7 @@ package seed
import ( import (
"fmt" "fmt"
"strings"
"sync/atomic" "sync/atomic"
"testing" "testing"
"time" "time"
@@ -134,27 +135,65 @@ func TestDemoSeedAcceptsAConcurrentWinnerWithoutOverwritingUser(t *testing.T) {
} }
func TestDemoSeedRetriesSerializationFailureOutsideTheTransaction(t *testing.T) { func TestDemoSeedRetriesSerializationFailureOutsideTheTransaction(t *testing.T) {
for _, sqlState := range []string{"40001", "40P01"} {
t.Run(sqlState, func(t *testing.T) {
database := newTestDB(t) database := newTestDB(t)
var failures atomic.Int32 email := "retry-" + strings.ToLower(sqlState) + "@senlin.ai"
callbackName := "test:fail_first_seed_user_create" var injected atomic.Bool
require.NoError(t, database.Callback().Create().Before("gorm:create").Register(callbackName, func(tx *gorm.DB) { var sawProjectWrite atomic.Bool
if tx.Statement.Schema == nil || tx.Statement.Schema.Table != (models.SenlinAgentUser{}).TableName() { 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 return
} }
if failures.Add(1) == 1 { if !injected.CompareAndSwap(false, true) {
tx.AddError(testSQLStateError{state: "40001"}) 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) }) 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.NoError(t, err)
require.Equal(t, int32(2), failures.Load()) require.True(t, injected.Load())
require.Equal(t, "retry@senlin.ai", result.User.Email) require.True(t, sawProjectWrite.Load(), "failed transaction must contain the first inserted project")
var userCount int64 require.NotEqual(t, failedOwnerIdentity, result.User.Identity, "retry must recreate the rolled-back user")
require.NoError(t, database.Model(&models.SenlinAgentUser{}).Where("email = ?", result.User.Email).Count(&userCount).Error) require.NotEqual(t, failedProjectIdentity, result.Projects[0].Identity, "retry must recreate the rolled-back project")
require.Equal(t, int64(1), userCount) 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 { type testSQLStateError struct {