package controlplane import ( "context" "encoding/base64" "fmt" "net/url" "os" "strings" "testing" "time" "github.com/jackc/pgx/v5/pgxpool" ) func TestAPIKeyRestrictionsRoundTripIntoRuntimeSnapshotPostgres(t *testing.T) { databaseURL := os.Getenv("AIGW_TEST_DATABASE_URL") if databaseURL == "" { t.Skip("AIGW_TEST_DATABASE_URL is not set") } ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) defer cancel() rootDB, err := pgxpool.New(ctx, databaseURL) if err != nil { t.Fatal(err) } defer rootDB.Close() schema := fmt.Sprintf("key_controls_%d", time.Now().UnixNano()) if _, err := rootDB.Exec(ctx, "CREATE SCHEMA "+schema); err != nil { t.Fatal(err) } t.Cleanup(func() { _, _ = rootDB.Exec(context.Background(), "DROP SCHEMA "+schema+" CASCADE") }) parsed, err := url.Parse(databaseURL) if err != nil { t.Fatal(err) } query := parsed.Query() query.Set("search_path", schema) parsed.RawQuery = query.Encode() isolatedURL := parsed.String() if err := MigrateDatabase(ctx, isolatedURL); err != nil { t.Fatal(err) } credentialKey := base64.StdEncoding.EncodeToString([]byte("01234567890123456789012345678901")) store, err := NewStore(ctx, Options{DatabaseURL: isolatedURL, CredentialKey: credentialKey}) if err != nil { t.Fatal(err) } t.Cleanup(func() { _ = store.Close() }) tenant, _, err := store.CreateTenant(ctx, CreateTenantInput{Slug: "key-control-test", Name: "Key Control Test"}) if err != nil { t.Fatal(err) } project, _, err := store.CreateProject(ctx, CreateProjectInput{TenantID: tenant.ID, Slug: "production", Name: "Production"}) if err != nil { t.Fatal(err) } provider, _, err := store.CreateProvider(ctx, CreateProviderInput{Name: "key-control-provider", Protocol: "openai", WireAPI: "embeddings", BaseURL: "https://example.invalid/v1", APIKey: "provider-secret"}) if err != nil { t.Fatal(err) } if provider.Slug != "key-control-provider" { t.Fatalf("derived provider slug = %q, want key-control-provider", provider.Slug) } if _, _, err := store.CreateProvider(ctx, CreateProviderInput{Name: "invalid-anthropic-embeddings", Protocol: "anthropic", WireAPI: "embeddings", BaseURL: "https://example.invalid/v1", APIKey: "provider-secret"}); err == nil { t.Fatal("Anthropic provider accepted the OpenAI Embeddings wire API") } model, _, err := store.CreateModel(ctx, CreateModelInput{PublicID: "model/key-control", DisplayName: "Key Control", Capabilities: []string{"embeddings"}, PriceCurrency: "usd", InputPriceMicrosPerMillion: 100_000, OutputPriceMicrosPerMillion: 200_000, Routes: []RouteInput{{ProviderID: provider.ID, UpstreamModel: "upstream-key-control", Weight: 1}}}) if err != nil { t.Fatal(err) } expiresAt := time.Now().Add(24 * time.Hour).UTC().Truncate(time.Microsecond) created, _, err := store.CreateAPIKey(ctx, CreateAPIKeyInput{TenantID: tenant.ID, ProjectID: project.ID, Name: "production backend", Scopes: []string{"inference"}, Tags: []string{"production", "backend"}, AllowedModels: []string{model.PublicID}, MonthlySpendMicros: 25_000_000, DailySpendMicros: 5_000_000, RequestsPerMinute: 12, TokensPerMinute: 34_000, ExpiresAt: &expiresAt}) if err != nil { t.Fatal(err) } if created.Key == "" || len(created.KeySuffix) != 6 || !strings.HasSuffix(created.Key, created.KeySuffix) || created.MonthlySpendMicros != 25_000_000 || created.DailySpendMicros != 5_000_000 || created.RequestsPerMinute != 12 || created.TokensPerMinute != 34_000 || len(created.AllowedModels) != 1 || len(created.Tags) != 2 { t.Fatalf("unexpected created key: %+v", created.APIKey) } if _, err := store.db.Exec(ctx, `INSERT INTO usage_events (request_id,tenant_id,project_id,key_id,public_model,protocol,status_code,success,started_at,cost_micros) VALUES ('req_key_current_month',$1,$2,$3,$4,'responses',200,TRUE,now(),42000), ('req_key_previous_month',$1,$2,$3,$4,'responses',200,TRUE,now()-interval '2 months',99000)`, tenant.ID, project.ID, created.ID, model.PublicID); err != nil { t.Fatal(err) } if _, err := store.db.Exec(ctx, `INSERT INTO billing_reservations (request_id,tenant_id,project_id,key_id,public_model,currency,reserved_micros,status, input_price_micros_per_million,output_price_micros_per_million,cache_read_price_micros_per_million,cache_write_price_micros_per_million) VALUES ('req_key_pending',$1,$2,$3,$4,'usd',9000,'pending',100000,200000,0,0)`, tenant.ID, project.ID, created.ID, model.PublicID); err != nil { t.Fatal(err) } keys, err := store.ListAPIKeysFor(ctx, tenant.ID) if err != nil { t.Fatal(err) } if len(keys) != 1 || keys[0].KeySuffix != created.KeySuffix || keys[0].AllowedModels[0] != model.PublicID || keys[0].ExpiresAt == nil || !keys[0].ExpiresAt.Equal(expiresAt) { t.Fatalf("key restrictions did not round trip: %+v", keys) } if keys[0].CurrentMonthSpendMicros != 42_000 || keys[0].CurrentMonthReservedMicros != 9_000 || keys[0].CurrentMonthRequests != 1 { t.Fatalf("key month activity is incorrect: %+v", keys[0]) } if keys[0].CurrentDaySpendMicros != 42_000 || keys[0].CurrentDayReservedMicros != 9_000 || keys[0].CurrentDayRequests != 1 { t.Fatalf("key daily activity is incorrect: %+v", keys[0]) } snapshot, err := store.LoadSnapshot(ctx) if err != nil { t.Fatal(err) } if len(snapshot.APIKeys) != 1 { t.Fatalf("snapshot API keys = %d, want 1", len(snapshot.APIKeys)) } if len(snapshot.Models) != 1 || len(snapshot.Models[0].Routes) != 1 || snapshot.Models[0].Routes[0].Provider.EffectiveWireAPI() != "embeddings" || len(snapshot.Models[0].Capabilities) != 1 || snapshot.Models[0].Capabilities[0] != "embeddings" { t.Fatalf("Embeddings model did not round trip into runtime snapshot: %+v", snapshot.Models) } principal := snapshot.APIKeys[0].Principal if principal.MonthlySpendMicros != 25_000_000 || principal.DailySpendMicros != 5_000_000 || principal.RequestsPerMinute != 12 || principal.TokensPerMinute != 34_000 || principal.ExpiresAt == nil { t.Fatalf("snapshot lost API key controls: %+v", principal) } if _, ok := principal.AllowedModels[model.PublicID]; !ok { t.Fatalf("snapshot lost allowed model: %+v", principal.AllowedModels) } otherTenant, _, err := store.CreateTenant(ctx, CreateTenantInput{Slug: "other-key-control-test", Name: "Other Key Control Test"}) if err != nil { t.Fatal(err) } otherProject, _, err := store.CreateProject(ctx, CreateProjectInput{TenantID: otherTenant.ID, Slug: "default", Name: "Other Default"}) if err != nil { t.Fatal(err) } if _, _, err := store.CreateAPIKey(ctx, CreateAPIKeyInput{TenantID: tenant.ID, ProjectID: otherProject.ID, Name: "cross tenant"}); err == nil { t.Fatal("cross-tenant project binding should be rejected by the composite foreign key") } if _, _, err := store.CreateAPIKey(ctx, CreateAPIKeyInput{TenantID: tenant.ID, ProjectID: project.ID, Name: "invalid model", AllowedModels: []string{"model/does-not-exist"}}); err == nil { t.Fatal("unknown allowed model should reject the whole API key transaction") } keys, err = store.ListAPIKeysFor(ctx, tenant.ID) if err != nil || len(keys) != 1 { t.Fatalf("failed key transaction leaked a row: keys=%d err=%v", len(keys), err) } if _, err := store.DisableAPIKey(ctx, created.ID); err != nil { t.Fatal(err) } snapshot, err = store.LoadSnapshot(ctx) if err != nil { t.Fatal(err) } if len(snapshot.APIKeys) != 0 { t.Fatalf("disabled key remained in runtime snapshot: %+v", snapshot.APIKeys) } if _, err := store.EnableAPIKey(ctx, created.ID); err != nil { t.Fatal(err) } rotated, _, err := store.RotateAPIKey(ctx, created.ID) if err != nil { t.Fatal(err) } if rotated.ID == created.ID || rotated.Key == "" || rotated.Key == created.Key || len(rotated.KeySuffix) != 6 || !strings.HasSuffix(rotated.Key, rotated.KeySuffix) || rotated.KeySuffix == created.KeySuffix || rotated.DailySpendMicros != created.DailySpendMicros || rotated.RequestsPerMinute != created.RequestsPerMinute || len(rotated.AllowedModels) != 1 || rotated.AllowedModels[0] != model.PublicID { t.Fatalf("rotated key did not preserve controls: old=%+v new=%+v", created.APIKey, rotated.APIKey) } snapshot, err = store.LoadSnapshot(ctx) if err != nil { t.Fatal(err) } if len(snapshot.APIKeys) != 1 || snapshot.APIKeys[0].Principal.KeyID != rotated.ID { t.Fatalf("rotation snapshot = %+v, want only new key %s", snapshot.APIKeys, rotated.ID) } keys, err = store.ListAPIKeysFor(ctx, tenant.ID) if err != nil || len(keys) != 2 || keys[0].Status != "active" || keys[1].Status != "revoked" { t.Fatalf("rotation lifecycle rows are incorrect: keys=%+v err=%v", keys, err) } } func TestMigrationUpgradesPreviousVersionAndIsIdempotentPostgres(t *testing.T) { databaseURL := os.Getenv("AIGW_TEST_DATABASE_URL") if databaseURL == "" { t.Skip("AIGW_TEST_DATABASE_URL is not set") } ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) defer cancel() db, err := pgxpool.New(ctx, databaseURL) if err != nil { t.Fatal(err) } defer db.Close() schema := fmt.Sprintf("migration_drill_%d", time.Now().UnixNano()) if _, err := db.Exec(ctx, "CREATE SCHEMA "+schema); err != nil { t.Fatal(err) } t.Cleanup(func() { _, _ = db.Exec(context.Background(), "DROP SCHEMA "+schema+" CASCADE") }) if _, err := db.Exec(ctx, `CREATE TABLE `+schema+`.schema_migrations ( version BIGINT PRIMARY KEY,name TEXT NOT NULL,checksum TEXT NOT NULL,applied_at TIMESTAMPTZ NOT NULL DEFAULT now())`); err != nil { t.Fatal(err) } if _, err := db.Exec(ctx, `INSERT INTO `+schema+`.schema_migrations(version,name,checksum) VALUES ($1,'previous-release','immutable-previous-checksum')`, migrationVersion-1); err != nil { t.Fatal(err) } parsed, err := url.Parse(databaseURL) if err != nil { t.Fatal(err) } query := parsed.Query() query.Set("search_path", schema) parsed.RawQuery = query.Encode() isolatedURL := parsed.String() if err := MigrateDatabase(ctx, isolatedURL); err != nil { t.Fatal(err) } if err := MigrateDatabase(ctx, isolatedURL); err != nil { t.Fatalf("second migration must be idempotent: %v", err) } status, err := MigrationStatusDatabase(ctx, isolatedURL) if err != nil { t.Fatal(err) } if status.Version != migrationVersion { t.Fatalf("migration version = %d, want %d", status.Version, migrationVersion) } var count int if err := db.QueryRow(ctx, `SELECT count(*) FROM `+schema+`.schema_migrations`).Scan(&count); err != nil { t.Fatal(err) } if count != 2 { t.Fatalf("migration history contains %d rows, want previous and current", count) } scopedDB, err := pgxpool.New(ctx, isolatedURL) if err != nil { t.Fatal(err) } defer scopedDB.Close() var tenantID, projectID, providerID, modelID string if err := scopedDB.QueryRow(ctx, `INSERT INTO tenants (slug,name) VALUES ('preference-test','Preference Test') RETURNING id::text`).Scan(&tenantID); err != nil { t.Fatal(err) } if err := scopedDB.QueryRow(ctx, `INSERT INTO projects (tenant_id,slug,name) VALUES ($1,'default','Default') RETURNING id::text`, tenantID).Scan(&projectID); err != nil { t.Fatal(err) } if _, err := scopedDB.Exec(ctx, `INSERT INTO api_keys (tenant_id,project_id,name,key_prefix,key_hash) VALUES ($1,$2,'Legacy key','sk-aigw-legacy...',decode(repeat('01',32),'hex'))`, tenantID, projectID); err != nil { t.Fatalf("legacy key without suffix did not retain its empty default: %v", err) } if _, err := scopedDB.Exec(ctx, `INSERT INTO api_keys (tenant_id,project_id,name,key_prefix,key_suffix,key_hash) VALUES ($1,$2,'Invalid display suffix','sk-aigw-invalid...','too-long',decode(repeat('02',32),'hex'))`, tenantID, projectID); err == nil { t.Fatal("invalid API key display suffix was accepted") } if err := scopedDB.QueryRow(ctx, `INSERT INTO providers (slug,name,protocol,wire_api,base_url,api_key_ciphertext) VALUES ('preference-provider','Preference provider','openai','responses','https://example.invalid',decode('00','hex')) RETURNING id::text`).Scan(&providerID); err != nil { t.Fatal(err) } if err := scopedDB.QueryRow(ctx, `INSERT INTO models (public_id,display_name) VALUES ('model/preference-test','Preference Model') RETURNING id::text`).Scan(&modelID); err != nil { t.Fatal(err) } if _, err := scopedDB.Exec(ctx, `INSERT INTO model_price_versions (model_id,version,currency) VALUES ($1,1,'usd')`, modelID); err != nil { t.Fatal(err) } if _, err := scopedDB.Exec(ctx, `INSERT INTO model_routes (model_id,provider_id,upstream_model) VALUES ($1,$2,'upstream-test')`, modelID, providerID); err != nil { t.Fatal(err) } store := &Store{db: scopedDB} preferenceDefaults := BillingPreferenceDefaults{LowBalanceThresholdMicros: 5_000_000, SpendAnomalyMultiplier: 3, SpendAnomalyMinMicros: 10_000_000} prefs, err := store.SetDeveloperPreferences(ctx, SetDeveloperPreferencesInput{TenantID: tenantID, DefaultModel: "model/preference-test"}, preferenceDefaults) if err != nil { t.Fatal(err) } if prefs.DefaultModel != "model/preference-test" || prefs.FallbackModel != "" { t.Fatalf("unexpected developer preferences: %+v", prefs) } enabled := false threshold := int64(9_750_000) anomalyEnabled := false anomalyMultiplier := int64(7) anomalyMinimum := int64(8_250_000) if _, err := store.SetBillingPreferences(ctx, SetBillingPreferencesInput{TenantID: tenantID, LowBalanceEnabled: &enabled, LowBalanceThresholdMicros: &threshold, SpendAnomalyEnabled: &anomalyEnabled, SpendAnomalyMultiplier: &anomalyMultiplier, SpendAnomalyMinMicros: &anomalyMinimum}, preferenceDefaults); err != nil { t.Fatal(err) } prefs, err = store.GetTenantPreferences(ctx, tenantID, preferenceDefaults) if err != nil { t.Fatal(err) } if prefs.LowBalanceEnabled || prefs.LowBalanceThresholdMicros != threshold || prefs.SpendAnomalyEnabled || prefs.SpendAnomalyMultiplier != anomalyMultiplier || prefs.SpendAnomalyMinMicros != anomalyMinimum || prefs.DefaultModel != "model/preference-test" { t.Fatalf("preferences did not round trip: %+v", prefs) } }