diff --git a/apps/api/internal/store/admission_lock.go b/apps/api/internal/store/admission_lock.go new file mode 100644 index 0000000..970411c --- /dev/null +++ b/apps/api/internal/store/admission_lock.go @@ -0,0 +1,170 @@ +package store + +import ( + "context" + "errors" + "sort" + "sync" + "time" + + "github.com/jackc/pgx/v5" +) + +const ( + admissionLockRetryMin = 10 * time.Millisecond + admissionLockRetryMax = 250 * time.Millisecond +) + +var ( + errAdmissionLockBusy = errors.New("task admission lock is busy") + processAdmissionLocks = newAdmissionLocalLockSet() +) + +type admissionLocalLock struct { + token chan struct{} + refs int +} + +type admissionLocalLockSet struct { + mu sync.Mutex + locks map[string]*admissionLocalLock +} + +func newAdmissionLocalLockSet() *admissionLocalLockSet { + return &admissionLocalLockSet{locks: make(map[string]*admissionLocalLock)} +} + +// acquire serializes contenders inside one API or Worker process before they +// consume a PostgreSQL connection. Sorted acquisition keeps operations that +// span task, platform-model, and user-group keys deadlock-free. +func (s *admissionLocalLockSet) acquire(ctx context.Context, keys []string) (func(), error) { + keys = normalizedAdmissionLockKeys(keys) + acquired := make([]struct { + key string + lock *admissionLocalLock + }, 0, len(keys)) + for _, key := range keys { + lock := s.reference(key) + select { + case lock.token <- struct{}{}: + acquired = append(acquired, struct { + key string + lock *admissionLocalLock + }{key: key, lock: lock}) + case <-ctx.Done(): + s.unreference(key, lock) + for index := len(acquired) - 1; index >= 0; index-- { + <-acquired[index].lock.token + s.unreference(acquired[index].key, acquired[index].lock) + } + return nil, ctx.Err() + } + } + return func() { + for index := len(acquired) - 1; index >= 0; index-- { + <-acquired[index].lock.token + s.unreference(acquired[index].key, acquired[index].lock) + } + }, nil +} + +func (s *admissionLocalLockSet) reference(key string) *admissionLocalLock { + s.mu.Lock() + defer s.mu.Unlock() + lock := s.locks[key] + if lock == nil { + lock = &admissionLocalLock{token: make(chan struct{}, 1)} + s.locks[key] = lock + } + lock.refs++ + return lock +} + +func (s *admissionLocalLockSet) unreference(key string, lock *admissionLocalLock) { + s.mu.Lock() + defer s.mu.Unlock() + lock.refs-- + if lock.refs == 0 && s.locks[key] == lock { + delete(s.locks, key) + } +} + +func normalizedAdmissionLockKeys(keys []string) []string { + seen := make(map[string]struct{}, len(keys)) + normalized := make([]string, 0, len(keys)) + for _, key := range keys { + if key == "" { + continue + } + if _, exists := seen[key]; exists { + continue + } + seen[key] = struct{}{} + normalized = append(normalized, key) + } + sort.Strings(normalized) + return normalized +} + +func admissionOperationLockKeys(input TaskAdmissionInput) []string { + keys := []string{"task-admission:" + input.TaskID} + for _, scope := range normalizedAdmissionScopes(input.Scopes) { + keys = append(keys, admissionLockKey(scope)) + } + return normalizedAdmissionLockKeys(keys) +} + +func tryAdmissionTransactionLock(ctx context.Context, tx pgx.Tx, key string) error { + var locked bool + if err := tx.QueryRow( + ctx, + `SELECT pg_try_advisory_xact_lock(hashtextextended($1, 0))`, + key, + ).Scan(&locked); err != nil { + return err + } + if !locked { + return errAdmissionLockBusy + } + return nil +} + +func retryAdmissionOperation[T any]( + ctx context.Context, + keys []string, + operation func() (T, error), +) (T, error) { + var zero T + for attempt := 0; ; attempt++ { + release, err := processAdmissionLocks.acquire(ctx, keys) + if err != nil { + return zero, err + } + result, operationErr := operation() + release() + if !errors.Is(operationErr, errAdmissionLockBusy) { + return result, operationErr + } + if err := waitAdmissionLockRetry(ctx, attempt); err != nil { + return zero, err + } + } +} + +func waitAdmissionLockRetry(ctx context.Context, attempt int) error { + delay := admissionLockRetryMin + for index := 0; index < attempt && delay < admissionLockRetryMax; index++ { + delay *= 2 + } + if delay > admissionLockRetryMax { + delay = admissionLockRetryMax + } + timer := time.NewTimer(delay) + defer timer.Stop() + select { + case <-timer.C: + return nil + case <-ctx.Done(): + return ctx.Err() + } +} diff --git a/apps/api/internal/store/admission_lock_test.go b/apps/api/internal/store/admission_lock_test.go new file mode 100644 index 0000000..99e3b83 --- /dev/null +++ b/apps/api/internal/store/admission_lock_test.go @@ -0,0 +1,221 @@ +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) + } +} diff --git a/apps/api/internal/store/admission_queue.go b/apps/api/internal/store/admission_queue.go index f8c6e2c..1f30866 100644 --- a/apps/api/internal/store/admission_queue.go +++ b/apps/api/internal/store/admission_queue.go @@ -130,13 +130,23 @@ func (s *Store) QueueTaskAdmissionWithHook( if input.Mode != "async" { return TaskAdmission{}, errors.New("queued admission registration requires asynchronous mode") } + return retryAdmissionOperation(ctx, admissionOperationLockKeys(input), func() (TaskAdmission, error) { + return s.queueTaskAdmissionOnce(ctx, input, onQueued) + }) +} + +func (s *Store) queueTaskAdmissionOnce( + ctx context.Context, + input TaskAdmissionInput, + onQueued func(pgx.Tx) error, +) (TaskAdmission, error) { tx, err := s.pool.Begin(ctx) if err != nil { return TaskAdmission{}, err } defer rollbackTransaction(tx) - if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended($1, 0))`, "task-admission:"+input.TaskID); err != nil { + if err := tryAdmissionTransactionLock(ctx, tx, "task-admission:"+input.TaskID); err != nil { return TaskAdmission{}, err } var taskActive bool @@ -163,7 +173,7 @@ WHERE id = $1::uuid`, input.TaskID).Scan(&taskActive); err != nil { ) } for _, scope := range normalizedAdmissionScopes(lockScopes) { - if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended($1, 0))`, admissionLockKey(scope)); err != nil { + if err := tryAdmissionTransactionLock(ctx, tx, admissionLockKey(scope)); err != nil { return TaskAdmission{}, err } } @@ -228,13 +238,23 @@ func (s *Store) tryTaskAdmission( if err := validateTaskAdmissionInput(input); err != nil { return TaskAdmissionResult{}, err } + return retryAdmissionOperation(ctx, admissionOperationLockKeys(input), func() (TaskAdmissionResult, error) { + return s.tryTaskAdmissionOnce(ctx, input, onAdmitted) + }) +} + +func (s *Store) tryTaskAdmissionOnce( + ctx context.Context, + input TaskAdmissionInput, + onAdmitted func(pgx.Tx) error, +) (TaskAdmissionResult, error) { tx, err := s.pool.Begin(ctx) if err != nil { return TaskAdmissionResult{}, err } defer rollbackTransaction(tx) - if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended($1, 0))`, "task-admission:"+input.TaskID); err != nil { + if err := tryAdmissionTransactionLock(ctx, tx, "task-admission:"+input.TaskID); err != nil { return TaskAdmissionResult{}, err } var taskActive bool @@ -261,7 +281,7 @@ WHERE id = $1::uuid`, input.TaskID).Scan(&taskActive); err != nil { ) } for _, scope := range normalizedAdmissionScopes(lockScopes) { - if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended($1, 0))`, admissionLockKey(scope)); err != nil { + if err := tryAdmissionTransactionLock(ctx, tx, admissionLockKey(scope)); err != nil { return TaskAdmissionResult{}, err } } @@ -452,12 +472,18 @@ func (s *Store) RebindWaitingTaskAdmission(ctx context.Context, input TaskAdmiss if err := validateTaskAdmissionInput(input); err != nil { return TaskAdmission{}, err } + return retryAdmissionOperation(ctx, admissionOperationLockKeys(input), func() (TaskAdmission, error) { + return s.rebindWaitingTaskAdmissionOnce(ctx, input) + }) +} + +func (s *Store) rebindWaitingTaskAdmissionOnce(ctx context.Context, input TaskAdmissionInput) (TaskAdmission, error) { tx, err := s.pool.Begin(ctx) if err != nil { return TaskAdmission{}, err } defer rollbackTransaction(tx) - if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended($1, 0))`, "task-admission:"+input.TaskID); err != nil { + if err := tryAdmissionTransactionLock(ctx, tx, "task-admission:"+input.TaskID); err != nil { return TaskAdmission{}, err } current, found, err := loadTaskAdmissionTx(ctx, tx, input.TaskID) @@ -477,7 +503,7 @@ func (s *Store) RebindWaitingTaskAdmission(ctx context.Context, input TaskAdmiss AdmissionScope{ScopeType: "user_group", ScopeKey: current.UserGroupID}, ) for _, scope := range normalizedAdmissionScopes(lockScopes) { - if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended($1, 0))`, admissionLockKey(scope)); err != nil { + if err := tryAdmissionTransactionLock(ctx, tx, admissionLockKey(scope)); err != nil { return TaskAdmission{}, err } } @@ -629,12 +655,19 @@ WHERE admission.task_id = owned.task_id } func (s *Store) DeleteTaskAdmission(ctx context.Context, taskID string) error { + _, err := retryAdmissionOperation(ctx, []string{"task-admission:" + taskID}, func() (struct{}, error) { + return struct{}{}, s.deleteTaskAdmissionOnce(ctx, taskID) + }) + return err +} + +func (s *Store) deleteTaskAdmissionOnce(ctx context.Context, taskID string) error { tx, err := s.pool.Begin(ctx) if err != nil { return err } defer rollbackTransaction(tx) - if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended($1, 0))`, "task-admission:"+taskID); err != nil { + if err := tryAdmissionTransactionLock(ctx, tx, "task-admission:"+taskID); err != nil { return err } if _, err := tx.Exec(ctx, `