package store import ( "context" "errors" "os" "strings" "sync" "sync/atomic" "testing" "time" ) func TestConcurrencyLeaseReservationIsAtomicAcrossPools(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 concurrency lease PostgreSQL integration tests") } ctx, cancel := context.WithTimeout(context.Background(), 45*time.Second) defer cancel() first, err := Connect(ctx, databaseURL) if err != nil { t.Fatalf("connect first store: %v", err) } defer first.Close() second, err := Connect(ctx, databaseURL) if err != nil { t.Fatalf("connect second store: %v", err) } defer second.Close() scopeKey := "atomic-" + time.Now().UTC().Format("20060102150405.000000000") taskIDs := createLeaseTestTasks(t, ctx, first, 256, scopeKey) defer deleteLeaseTestTasks(t, first, taskIDs) var successes atomic.Int64 var peak atomic.Int64 monitorCtx, stopMonitor := context.WithCancel(ctx) var monitorWG sync.WaitGroup monitorWG.Add(1) go func() { defer monitorWG.Done() ticker := time.NewTicker(2 * time.Millisecond) defer ticker.Stop() for { select { case <-monitorCtx.Done(): return case <-ticker.C: var active int64 if err := first.Pool().QueryRow(monitorCtx, ` SELECT COUNT(*) FROM gateway_concurrency_leases WHERE scope_type = 'platform_model' AND scope_key = $1 AND released_at IS NULL AND expires_at > now()`, scopeKey).Scan(&active); err == nil { for active > peak.Load() && !peak.CompareAndSwap(peak.Load(), active) { } } } } }() var wg sync.WaitGroup errs := make(chan error, len(taskIDs)) for index, taskID := range taskIDs { wg.Add(1) go func(index int, taskID string) { defer wg.Done() target := first if index%2 == 1 { target = second } _, err := target.ReserveRateLimits(ctx, taskID, "", []RateLimitReservation{{ ScopeType: "platform_model", ScopeKey: scopeKey, Metric: "concurrent", Limit: 64, Amount: 1, LeaseTTLSeconds: 30, }}) if err == nil { successes.Add(1) return } if !errors.Is(err, ErrRateLimited) { errs <- err } }(index, taskID) } wg.Wait() stopMonitor() monitorWG.Wait() close(errs) for err := range errs { t.Fatalf("unexpected reservation error: %v", err) } var active int64 if err := first.Pool().QueryRow(ctx, ` SELECT COUNT(*) FROM gateway_concurrency_leases WHERE scope_type = 'platform_model' AND scope_key = $1 AND released_at IS NULL AND expires_at > now()`, scopeKey).Scan(&active); err != nil { t.Fatalf("count active leases: %v", err) } if successes.Load() != 64 || active != 64 { t.Fatalf("successful reservations=%d active leases=%d, want exactly 64", successes.Load(), active) } if peak.Load() > 64 { t.Fatalf("active lease peak=%d, want <=64", peak.Load()) } } func TestConcurrencyLeaseRenewalExtendsAndReleases(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 concurrency lease PostgreSQL integration tests") } ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) defer cancel() db, err := Connect(ctx, databaseURL) if err != nil { t.Fatalf("connect store: %v", err) } defer db.Close() scopeKey := "renew-" + time.Now().UTC().Format("20060102150405.000000000") taskIDs := createLeaseTestTasks(t, ctx, db, 1, scopeKey) defer deleteLeaseTestTasks(t, db, taskIDs) result, err := db.ReserveRateLimits(ctx, taskIDs[0], "", []RateLimitReservation{{ ScopeType: "platform_model", ScopeKey: scopeKey, Metric: "concurrent", Limit: 1, Amount: 1, LeaseTTLSeconds: 2, }}) if err != nil { t.Fatalf("reserve short lease: %v", err) } time.Sleep(time.Second) if err := db.RenewConcurrencyLeases(ctx, result.Leases); err != nil { t.Fatalf("renew short lease: %v", err) } time.Sleep(1500 * time.Millisecond) var active bool if err := db.Pool().QueryRow(ctx, ` SELECT EXISTS ( SELECT 1 FROM gateway_concurrency_leases WHERE id = $1::uuid AND released_at IS NULL AND expires_at > now() )`, result.Leases[0].ID).Scan(&active); err != nil { t.Fatalf("read renewed lease: %v", err) } if !active { t.Fatal("renewed lease expired at its original TTL") } if err := db.ReleaseConcurrencyLeases(ctx, result.Leases); err != nil { t.Fatalf("release renewed lease: %v", err) } } func createLeaseTestTasks(t *testing.T, ctx context.Context, db *Store, count int, marker string) []string { t.Helper() rows, err := db.Pool().Query(ctx, ` INSERT INTO gateway_tasks (kind, run_mode, user_id, model, model_type, request, status, queue_key) SELECT 'lease-test', 'simulation', $2, 'lease-test', 'text_generate', '{}'::jsonb, 'queued', $2 FROM generate_series(1, $1) RETURNING id::text`, count, marker) if err != nil { t.Fatalf("create lease test tasks: %v", err) } defer rows.Close() ids := make([]string, 0, count) for rows.Next() { var id string if err := rows.Scan(&id); err != nil { t.Fatalf("scan lease test task: %v", err) } ids = append(ids, id) } if err := rows.Err(); err != nil { t.Fatalf("create lease test tasks: %v", err) } return ids } func deleteLeaseTestTasks(t *testing.T, db *Store, taskIDs []string) { t.Helper() cleanupCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second) defer cancel() if _, err := db.Pool().Exec(cleanupCtx, `DELETE FROM gateway_tasks WHERE id = ANY($1::uuid[])`, taskIDs); err != nil { t.Errorf("delete lease test tasks: %v", err) } }