package store import ( "context" "os" "strings" "sync" "sync/atomic" "testing" "time" "github.com/google/uuid" "github.com/jackc/pgx/v5" ) func TestAdmissionLocalLockSetSerializesSharedKeys(t *testing.T) { lockSet := newAdmissionLocalLockSet() ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() releaseFirst, err := lockSet.acquire(ctx, []string{"scope:b", "scope:a", "scope:a"}) if err != nil { t.Fatalf("acquire first lock set: %v", err) } const contenders = 64 var active atomic.Int32 var maxActive atomic.Int32 var waitGroup sync.WaitGroup started := make(chan struct{}, contenders) for index := 0; index < contenders; index++ { waitGroup.Add(1) go func() { defer waitGroup.Done() started <- struct{}{} release, acquireErr := lockSet.acquire(ctx, []string{"scope:a", "scope:b"}) if acquireErr != nil { t.Errorf("acquire contended lock set: %v", acquireErr) return } current := active.Add(1) for { observed := maxActive.Load() if current <= observed || maxActive.CompareAndSwap(observed, current) { break } } time.Sleep(time.Millisecond) active.Add(-1) release() }() } for index := 0; index < contenders; index++ { <-started } time.Sleep(20 * time.Millisecond) if active.Load() != 0 { t.Fatalf("contenders entered while first lock was held: active=%d", active.Load()) } releaseFirst() waitGroup.Wait() if maxActive.Load() != 1 { t.Fatalf("maximum concurrent holders = %d, want 1", maxActive.Load()) } lockSet.mu.Lock() defer lockSet.mu.Unlock() if len(lockSet.locks) != 0 { t.Fatalf("local lock entries leaked: %d", len(lockSet.locks)) } } func TestRetryAdmissionOperationReleasesLocalLockBeforeBackoff(t *testing.T) { ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() lockKey := "scope:retry-test:" + uuid.NewString() var firstAttempts atomic.Int32 firstStarted := make(chan struct{}) firstDone := make(chan error, 1) go func() { _, err := retryAdmissionOperation(ctx, []string{lockKey}, func() (struct{}, error) { if firstAttempts.Add(1) == 1 { close(firstStarted) return struct{}{}, errAdmissionLockBusy } return struct{}{}, nil }) firstDone <- err }() <-firstStarted secondRan := make(chan struct{}) if _, err := retryAdmissionOperation(ctx, []string{lockKey}, func() (struct{}, error) { close(secondRan) return struct{}{}, nil }); err != nil { t.Fatalf("second operation: %v", err) } select { case <-secondRan: default: t.Fatal("second operation did not run during first operation backoff") } if err := <-firstDone; err != nil { t.Fatalf("first operation retry: %v", err) } if firstAttempts.Load() != 2 { t.Fatalf("first attempts = %d, want 2", firstAttempts.Load()) } } func TestAdmissionDatabaseLockContentionDoesNotOccupyPool(t *testing.T) { databaseURL := strings.TrimSpace(os.Getenv("AI_GATEWAY_TEST_DATABASE_URL")) if databaseURL == "" { t.Skip("set AI_GATEWAY_TEST_DATABASE_URL to run task admission PostgreSQL integration tests") } ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) defer cancel() db, err := ConnectWithMaxConns(ctx, databaseURL, 8) if err != nil { t.Fatalf("connect store: %v", err) } defer db.Close() var databaseName string if err := db.pool.QueryRow(ctx, `SELECT current_database()`).Scan(&databaseName); err != nil { t.Fatalf("read test database name: %v", err) } if !strings.Contains(strings.ToLower(databaseName), "test") { t.Fatalf("refusing to use non-test database %q", databaseName) } lockKey := "task-admission-contention-test:" + uuid.NewString() holder, err := db.pool.Begin(ctx) if err != nil { t.Fatalf("begin lock holder: %v", err) } defer rollbackTransaction(holder) if _, err := holder.Exec( ctx, `SELECT pg_advisory_xact_lock(hashtextextended($1, 0))`, lockKey, ); err != nil { t.Fatalf("hold advisory lock: %v", err) } const contenders = 64 var peakAcquired atomic.Int32 sampleDone := make(chan struct{}) go func() { ticker := time.NewTicker(time.Millisecond) defer ticker.Stop() for { select { case <-ticker.C: current := int32(db.pool.Stat().AcquiredConns()) for { observed := peakAcquired.Load() if current <= observed || peakAcquired.CompareAndSwap(observed, current) { break } } case <-sampleDone: return } } }() start := make(chan struct{}) results := make(chan error, contenders) var waitGroup sync.WaitGroup for index := 0; index < contenders; index++ { waitGroup.Add(1) go func() { defer waitGroup.Done() <-start _, retryErr := retryAdmissionOperation(ctx, []string{lockKey}, func() (struct{}, error) { tx, beginErr := db.pool.Begin(ctx) if beginErr != nil { return struct{}{}, beginErr } defer rollbackTransaction(tx) if lockErr := tryAdmissionTransactionLock(ctx, tx, lockKey); lockErr != nil { return struct{}{}, lockErr } return struct{}{}, tx.Commit(ctx) }) results <- retryErr }() } close(start) time.Sleep(100 * time.Millisecond) if current := db.pool.Stat().AcquiredConns(); current > 2 { t.Fatalf("connections acquired during lock contention = %d, want at most 2", current) } if err := holder.Commit(ctx); err != nil { t.Fatalf("release advisory lock: %v", err) } waitGroup.Wait() close(sampleDone) close(results) for retryErr := range results { if retryErr != nil { t.Fatalf("contended operation: %v", retryErr) } } if peakAcquired.Load() > 2 { t.Fatalf("peak acquired connections = %d, want at most 2", peakAcquired.Load()) } var advisoryWaiters int if err := db.pool.QueryRow(ctx, ` SELECT COUNT(*) FROM pg_stat_activity WHERE datname = current_database() AND wait_event = 'advisory'`).Scan(&advisoryWaiters); err != nil && err != pgx.ErrNoRows { t.Fatalf("count advisory waiters: %v", err) } if advisoryWaiters != 0 { t.Fatalf("advisory waiters after contention = %d, want 0", advisoryWaiters) } }