summaryrefslogtreecommitdiff
path: root/internal/controlplane/store_integration_test.go
diff options
context:
space:
mode:
Diffstat (limited to 'internal/controlplane/store_integration_test.go')
-rw-r--r--internal/controlplane/store_integration_test.go84
1 files changed, 73 insertions, 11 deletions
diff --git a/internal/controlplane/store_integration_test.go b/internal/controlplane/store_integration_test.go
index 9bf093c..b0038c1 100644
--- a/internal/controlplane/store_integration_test.go
+++ b/internal/controlplane/store_integration_test.go
@@ -6,6 +6,7 @@ import (
"fmt"
"net/url"
"os"
+ "strings"
"testing"
"time"
@@ -54,25 +55,30 @@ func TestAPIKeyRestrictionsRoundTripIntoRuntimeSnapshotPostgres(t *testing.T) {
if err != nil {
t.Fatal(err)
}
- provider, _, err := store.CreateProvider(ctx, CreateProviderInput{Name: "key-control-provider", Protocol: "openai", WireAPI: "responses", BaseURL: "https://example.invalid/v1", APIKey: "provider-secret"})
+ 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)
}
- model, _, err := store.CreateModel(ctx, CreateModelInput{PublicID: "model/key-control", DisplayName: "Key Control", PriceCurrency: "usd", InputPriceMicrosPerMillion: 100_000, OutputPriceMicrosPerMillion: 200_000, Routes: []RouteInput{{ProviderID: provider.ID, UpstreamModel: "upstream-key-control", Weight: 1}}})
+ 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, ExpiresAt: &expiresAt})
+ 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 == "" || created.MonthlySpendMicros != 25_000_000 || len(created.AllowedModels) != 1 || len(created.Tags) != 2 {
+ 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
@@ -93,12 +99,15 @@ func TestAPIKeyRestrictionsRoundTripIntoRuntimeSnapshotPostgres(t *testing.T) {
if err != nil {
t.Fatal(err)
}
- if len(keys) != 1 || keys[0].AllowedModels[0] != model.PublicID || keys[0].ExpiresAt == nil || !keys[0].ExpiresAt.Equal(expiresAt) {
+ 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)
@@ -106,8 +115,11 @@ func TestAPIKeyRestrictionsRoundTripIntoRuntimeSnapshotPostgres(t *testing.T) {
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.ExpiresAt == nil {
+ 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 {
@@ -132,6 +144,38 @@ func TestAPIKeyRestrictionsRoundTripIntoRuntimeSnapshotPostgres(t *testing.T) {
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) {
@@ -192,10 +236,21 @@ func TestMigrationUpgradesPreviousVersionAndIsIdempotentPostgres(t *testing.T) {
t.Fatal(err)
}
defer scopedDB.Close()
- var tenantID, providerID, modelID string
+ 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)
@@ -210,7 +265,8 @@ func TestMigrationUpgradesPreviousVersionAndIsIdempotentPostgres(t *testing.T) {
t.Fatal(err)
}
store := &Store{db: scopedDB}
- prefs, err := store.SetDeveloperPreferences(ctx, SetDeveloperPreferencesInput{TenantID: tenantID, DefaultModel: "model/preference-test"})
+ 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)
}
@@ -219,15 +275,21 @@ func TestMigrationUpgradesPreviousVersionAndIsIdempotentPostgres(t *testing.T) {
}
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}, 5_000_000); err != nil {
+ LowBalanceEnabled: &enabled, LowBalanceThresholdMicros: &threshold, SpendAnomalyEnabled: &anomalyEnabled,
+ SpendAnomalyMultiplier: &anomalyMultiplier, SpendAnomalyMinMicros: &anomalyMinimum}, preferenceDefaults); err != nil {
t.Fatal(err)
}
- prefs, err = store.GetTenantPreferences(ctx, tenantID, 5_000_000)
+ prefs, err = store.GetTenantPreferences(ctx, tenantID, preferenceDefaults)
if err != nil {
t.Fatal(err)
}
- if prefs.LowBalanceEnabled || prefs.LowBalanceThresholdMicros != threshold || prefs.DefaultModel != "model/preference-test" {
+ 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)
}
}