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 var storedLimit float64 if err := first.Pool().QueryRow(ctx, ` SELECT COUNT(*), COALESCE(MAX(limit_value), 0)::float8 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, &storedLimit); err != nil { t.Fatalf("count active leases: %v", err) } if successes.Load() != 64 || active != 64 || storedLimit != 64 { t.Fatalf("successful reservations=%d active leases=%d stored limit=%.0f, want exactly 64", successes.Load(), active, storedLimit) } if peak.Load() > 64 { t.Fatalf("active lease peak=%d, want <=64", peak.Load()) } } func TestConcurrencyLeaseTimestampStartsAtReservationStatement(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() database, err := Connect(ctx, databaseURL) if err != nil { t.Fatalf("connect store: %v", err) } defer database.Close() scopeKey := "statement-clock-" + time.Now().UTC().Format("20060102150405.000000000") taskIDs := createLeaseTestTasks(t, ctx, database, 1, scopeKey) defer deleteLeaseTestTasks(t, database, taskIDs) tx, err := database.pool.Begin(ctx) if err != nil { t.Fatalf("begin reservation transaction: %v", err) } defer rollbackTransaction(tx) var transactionStartedAt time.Time if err := tx.QueryRow(ctx, `SELECT now()`).Scan(&transactionStartedAt); err != nil { t.Fatalf("read transaction start: %v", err) } time.Sleep(1100 * time.Millisecond) lease, err := reserveConcurrencyLease(ctx, tx, taskIDs[0], "", RateLimitReservation{ ScopeType: "platform_model", ScopeKey: scopeKey, Metric: "concurrent", Limit: 1, Amount: 1, LeaseTTLSeconds: 30, }) if err != nil { t.Fatalf("reserve concurrency lease: %v", err) } if err := tx.Commit(ctx); err != nil { t.Fatalf("commit reservation transaction: %v", err) } var acquiredAt, expiresAt time.Time if err := database.pool.QueryRow(ctx, ` SELECT acquired_at, expires_at FROM gateway_concurrency_leases WHERE id = $1::uuid`, lease.ID).Scan(&acquiredAt, &expiresAt); err != nil { t.Fatalf("read lease timestamps: %v", err) } if elapsed := acquiredAt.Sub(transactionStartedAt); elapsed < time.Second { t.Fatalf("lease acquired_at advanced by %s, want at least 1s after transaction start", elapsed) } if ttl := expiresAt.Sub(acquiredAt); ttl != 30*time.Second { t.Fatalf("lease ttl=%s, want 30s", ttl) } } func TestCounterWindowReservationIsAtomicAcrossPools(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 rate limit 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() tests := []struct { metric string limit float64 amount float64 wantSuccesses int64 }{ {metric: "rpm", limit: 37, amount: 1, wantSuccesses: 37}, {metric: "tpm_total", limit: 1_000, amount: 25, wantSuccesses: 40}, } for _, test := range tests { t.Run(test.metric, func(t *testing.T) { scopeKey := "atomic-" + test.metric + "-" + time.Now().UTC().Format("20060102150405.000000000") taskIDs := createLeaseTestTasks(t, ctx, first, 128, scopeKey) defer deleteLeaseTestTasks(t, first, taskIDs) var successes atomic.Int64 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: test.metric, Limit: test.limit, Amount: test.amount, WindowSeconds: 3600, }}) if err == nil { successes.Add(1) return } if !errors.Is(err, ErrRateLimited) { errs <- err } }(index, taskID) } wg.Wait() close(errs) for err := range errs { t.Fatalf("unexpected reservation error: %v", err) } var current float64 if err := first.Pool().QueryRow(ctx, ` SELECT COALESCE(MAX(used_value + reserved_value), 0)::float8 FROM gateway_rate_limit_counters WHERE scope_type = 'platform_model' AND scope_key = $1 AND metric = $2`, scopeKey, test.metric).Scan(¤t); err != nil { t.Fatalf("read %s counter: %v", test.metric, err) } if successes.Load() != test.wantSuccesses { t.Fatalf("successful %s reservations=%d, want exactly %d", test.metric, successes.Load(), test.wantSuccesses) } wantCurrent := float64(test.wantSuccesses) * test.amount if current != wantCurrent || current > test.limit { t.Fatalf("%s current=%.0f, want %.0f and <= %.0f", test.metric, current, wantCurrent, test.limit) } }) } } 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) } }