From 5135793bd27fe2c0572c1291cbcd419e28906b5a Mon Sep 17 00:00:00 2001 From: Rogee Date: Fri, 21 Aug 2026 13:25:39 +0800 Subject: [PATCH] test(autoassignment): cover concurrent CAS loser (#71) Co-authored-by: Rogee --- .../internal/autoassignment/coverage9_test.go | 67 ++++++++++++++----- 1 file changed, 50 insertions(+), 17 deletions(-) diff --git a/backend/internal/autoassignment/coverage9_test.go b/backend/internal/autoassignment/coverage9_test.go index a805fbcf..deb2bf7c 100644 --- a/backend/internal/autoassignment/coverage9_test.go +++ b/backend/internal/autoassignment/coverage9_test.go @@ -5,7 +5,9 @@ import ( "fmt" "strings" "sync" + "sync/atomic" "testing" + "time" "github.com/alicebob/miniredis/v2" "github.com/gochat/gochat/internal/channel" @@ -19,8 +21,11 @@ import ( func setupFullAADB_Cov9(t *testing.T) (*gorm.DB, *redis.Client) { t.Helper() - db, err := gorm.Open(sqlite.Open(fmt.Sprintf("file:%s?mode=memory&cache=shared", strings.ReplaceAll(t.Name(), "/", "_"))), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)}) + db, err := gorm.Open(sqlite.Open(fmt.Sprintf("file:%s?mode=memory&cache=shared&_busy_timeout=5000", strings.ReplaceAll(t.Name(), "/", "_"))), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)}) require.NoError(t, err) + sqlDB, err := db.DB() + require.NoError(t, err) + t.Cleanup(func() { _ = sqlDB.Close() }) require.NoError(t, db.AutoMigrate( &model.Account{}, &model.User{}, &model.AccountUser{}, &model.Inbox{}, &model.InboxMember{}, &model.Contact{}, &model.ContactInbox{}, @@ -164,40 +169,68 @@ func TestAssignmentServiceAutoAssignmentDoesNotOverwriteNewConversationState(t * func TestAssignmentServiceOnlyOneConcurrentWorkerWins(t *testing.T) { db, rdb := setupFullAADB_Cov9(t) account, agent, inbox, conversation := seedAssignableConversation_Cov9(t, db) + db = db.Session(&gorm.Session{SkipDefaultTransaction: true}) sqlDB, err := db.DB() require.NoError(t, err) - sqlDB.SetMaxOpenConns(1) + sqlDB.SetMaxOpenConns(2) + sqlDB.SetMaxIdleConns(2) + + var arrivals atomic.Int32 + readBarrier := make(chan struct{}) + type barrierContextKey struct{} + barrierContext := context.WithValue(context.Background(), barrierContextKey{}, true) + require.NoError(t, db.Callback().Query().After("gorm:query").Register("test:concurrent_auto_assignment_barrier", func(tx *gorm.DB) { + if tx.Statement.Table != "conversations" || tx.Statement.Context == nil || tx.Statement.Context.Value(barrierContextKey{}) == nil { + return + } + if arrivals.Add(1) == 2 { + close(readBarrier) + } + select { + case <-readBarrier: + case <-time.After(5 * time.Second): + tx.AddError(fmt.Errorf("concurrent assignment barrier timed out after %d arrivals", arrivals.Load())) + } + })) + t.Cleanup(func() { _ = db.Callback().Query().Remove("test:concurrent_auto_assignment_barrier") }) start := make(chan struct{}) - results := make(chan uint, 2) - errs := make(chan error, 2) + type result struct { + assignedIDs []uint + err error + } + results := make(chan result, 2) var wg sync.WaitGroup for range 2 { wg.Add(1) go func() { defer wg.Done() <-start - agentID, err := NewAssignmentService(db, rdb).AssignConversation(context.Background(), conversation.ID, inbox.ID, account.ID) - results <- agentID - errs <- err + assignedIDs, err := NewAssignmentService(db, rdb).AssignUnassignedConversations(barrierContext, inbox.ID, account.ID) + results <- result{assignedIDs: assignedIDs, err: err} }() } close(start) wg.Wait() close(results) - close(errs) - winners := 0 - for err := range errs { - require.NoError(t, err) - } - for agentID := range results { - if agentID != 0 { - require.Equal(t, agent.ID, agentID) - winners++ + assignedCount := 0 + for result := range results { + require.NoError(t, result.err) + assignedCount += len(result.assignedIDs) + for _, conversationID := range result.assignedIDs { + require.Equal(t, conversation.ID, conversationID) } } - require.Equal(t, 1, winners) + require.Equal(t, int32(2), arrivals.Load(), "both workers must read the same unassigned conversation") + require.Equal(t, 1, assignedCount) + + var updated model.Conversation + require.NoError(t, db.First(&updated, conversation.ID).Error) + require.Equal(t, agent.ID, *updated.AssigneeID) + rateCount, err := NewRateLimiter(rdb).GetCount(context.Background(), inbox.ID, agent.ID, 300) + require.NoError(t, err) + require.Equal(t, 1, rateCount, "the CAS loser must not increment the rate-limit key") } func TestAssignmentServiceExcludesNonAgentInboxMembers(t *testing.T) {