From 41e322c53d7b4b796eb377d0df9c29ecd10ba431 Mon Sep 17 00:00:00 2001 From: Chia Date: Thu, 6 Aug 2026 09:29:41 +1200 Subject: feat: complete commercial control plane, billing, auth, and model catalog - add PostgreSQL control-plane persistence with Redis-degraded hot reload - implement prepaid balance, usage ledger, Stripe top-up and reconciliation - add registration, email verification, password reset, invitations and RBAC - support TOTP, Passkey MFA, device sessions, quotas and rate limits - add tenant billing profiles, audit logs and operational readiness checks - build authenticated admin console, Quickstart, Playground and usage analytics - add public model catalog with pricing, filtering and cost estimation - support OpenAI Responses providers and provider health failover - validate real upstream usage reporting and balance settlement --- internal/controlplane/usage_integration_test.go | 152 ++++++++++++++++++++++++ 1 file changed, 152 insertions(+) create mode 100644 internal/controlplane/usage_integration_test.go (limited to 'internal/controlplane/usage_integration_test.go') diff --git a/internal/controlplane/usage_integration_test.go b/internal/controlplane/usage_integration_test.go new file mode 100644 index 0000000..2545c1d --- /dev/null +++ b/internal/controlplane/usage_integration_test.go @@ -0,0 +1,152 @@ +package controlplane + +import ( + "context" + "crypto/sha256" + "fmt" + "os" + "testing" + "time" + + "aigw/internal/domain" + + "github.com/jackc/pgx/v5/pgxpool" +) + +func TestUsageFiltersAndDailyAggregationPostgres(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() + if err := MigrateDatabase(ctx, databaseURL); err != nil { + t.Fatal(err) + } + db, err := pgxpool.New(ctx, databaseURL) + if err != nil { + t.Fatal(err) + } + t.Cleanup(db.Close) + store := &Store{db: db} + suffix := time.Now().UnixNano() + var tenantID, projectID, keyID, providerID, otherTenantID, otherProjectID, otherKeyID string + if err := db.QueryRow(ctx, `INSERT INTO tenants (slug,name) VALUES ($1,'Usage integration') RETURNING id::text`, fmt.Sprintf("usage-%d", suffix)).Scan(&tenantID); err != nil { + t.Fatal(err) + } + if err := db.QueryRow(ctx, `INSERT INTO projects (tenant_id,slug,name) VALUES ($1,'production','Production') RETURNING id::text`, tenantID).Scan(&projectID); err != nil { + t.Fatal(err) + } + keyHash := sha256.Sum256([]byte(fmt.Sprint(suffix))) + if err := db.QueryRow(ctx, `INSERT INTO api_keys (tenant_id,project_id,name,key_prefix,key_hash) VALUES ($1,$2,'Production key','sk-usage',$3) RETURNING id::text`, tenantID, projectID, keyHash[:]).Scan(&keyID); err != nil { + t.Fatal(err) + } + if err := db.QueryRow(ctx, `INSERT INTO providers (slug,name,protocol,wire_api,base_url,api_key_ciphertext) + VALUES ($1,$2,'openai','responses','https://usage.test',$3) RETURNING id::text`, fmt.Sprintf("usage-provider-%d", suffix), fmt.Sprintf("Usage provider %d", suffix), []byte("encrypted-test-value")).Scan(&providerID); err != nil { + t.Fatal(err) + } + if err := db.QueryRow(ctx, `INSERT INTO tenants (slug,name) VALUES ($1,'Other usage tenant') RETURNING id::text`, fmt.Sprintf("usage-other-%d", suffix)).Scan(&otherTenantID); err != nil { + t.Fatal(err) + } + if err := db.QueryRow(ctx, `INSERT INTO projects (tenant_id,slug,name) VALUES ($1,'production','Other production') RETURNING id::text`, otherTenantID).Scan(&otherProjectID); err != nil { + t.Fatal(err) + } + otherKeyHash := sha256.Sum256([]byte(fmt.Sprintf("other-%d", suffix))) + if err := db.QueryRow(ctx, `INSERT INTO api_keys (tenant_id,project_id,name,key_prefix,key_hash) VALUES ($1,$2,'Other key','sk-other',$3) RETURNING id::text`, otherTenantID, otherProjectID, otherKeyHash[:]).Scan(&otherKeyID); err != nil { + t.Fatal(err) + } + t.Cleanup(func() { + cleanupCtx := context.Background() + for _, statement := range []struct{ query, arg string }{ + {`DELETE FROM usage_monthly_rollups WHERE tenant_id=$1`, tenantID}, + {`DELETE FROM usage_monthly_rollups WHERE tenant_id=$1`, otherTenantID}, + {`DELETE FROM usage_events WHERE tenant_id=$1`, tenantID}, + {`DELETE FROM usage_events WHERE tenant_id=$1`, otherTenantID}, + {`DELETE FROM providers WHERE id=$1`, providerID}, + {`DELETE FROM api_keys WHERE id=$1`, keyID}, + {`DELETE FROM api_keys WHERE id=$1`, otherKeyID}, + {`DELETE FROM projects WHERE id=$1`, projectID}, + {`DELETE FROM projects WHERE id=$1`, otherProjectID}, + {`DELETE FROM tenants WHERE id=$1`, tenantID}, + {`DELETE FROM tenants WHERE id=$1`, otherTenantID}, + } { + if _, cleanupErr := db.Exec(cleanupCtx, statement.query, statement.arg); cleanupErr != nil { + t.Errorf("cleanup usage integration data: %v", cleanupErr) + } + } + }) + + started := time.Now().UTC().Add(-time.Hour).Truncate(time.Second) + events := []domain.UsageEvent{ + {RequestID: fmt.Sprintf("req_usage_ok_%d", suffix), TenantID: tenantID, ProjectID: projectID, KeyID: keyID, PublicModel: "openai/test", ProviderID: providerID, Protocol: domain.ProtocolOpenAI, StatusCode: 200, Success: true, UsageReported: true, StartedAt: started, DurationMS: 120, Usage: domain.Usage{InputTokens: 10, OutputTokens: 4, TotalTokens: 14, CacheReadInputTokens: 5}}, + {RequestID: fmt.Sprintf("req_usage_error_%d", suffix), TenantID: tenantID, ProjectID: projectID, KeyID: keyID, PublicModel: "openai/test", ProviderID: providerID, Protocol: domain.ProtocolOpenAIResponses, Stream: true, StatusCode: 502, Success: false, ErrorType: "provider_error", StartedAt: started.Add(time.Minute), DurationMS: 350}, + } + for _, event := range events { + if err := store.RecordUsage(ctx, event); err != nil { + t.Fatal(err) + } + } + if _, err := db.Exec(ctx, `UPDATE usage_events SET charged_micros=125,cost_micros=125 WHERE request_id=$1`, events[0].RequestID); err != nil { + t.Fatal(err) + } + + records, err := store.ListUsage(ctx, UsageQuery{TenantID: tenantID, KeyID: keyID, Status: "success", From: started.Add(-time.Minute), To: started.Add(time.Hour), Limit: 20}) + if err != nil { + t.Fatal(err) + } + if len(records) != 1 || records[0].RequestID != events[0].RequestID || records[0].ProjectName != "Production" || records[0].KeyName != "Production key" || records[0].MeteringStatus != "reported" { + t.Fatalf("unexpected filtered usage: %+v", records) + } + streaming := true + providerSlug := fmt.Sprintf("usage-provider-%d", suffix) + records, err = store.ListUsage(ctx, UsageQuery{TenantID: tenantID, Provider: providerSlug, Protocol: string(domain.ProtocolOpenAIResponses), ErrorType: "provider_error", Stream: &streaming, From: started.Add(-time.Minute), To: started.Add(time.Hour), Limit: 20}) + if err != nil { + t.Fatal(err) + } + if len(records) != 1 || records[0].RequestID != events[1].RequestID { + t.Fatalf("unexpected provider/protocol/transport usage filter: %+v", records) + } + filteredPoints, err := store.UsageDaily(ctx, UsageQuery{TenantID: tenantID, Provider: providerSlug, Protocol: string(domain.ProtocolOpenAIResponses), ErrorType: "provider_error", Stream: &streaming, From: started.Add(-time.Minute), To: started.Add(time.Hour)}) + if err != nil { + t.Fatal(err) + } + if len(filteredPoints) != 1 || filteredPoints[0].RequestCount != 1 || filteredPoints[0].SuccessfulRequests != 0 { + t.Fatalf("unexpected filtered daily usage: %+v", filteredPoints) + } + points, err := store.UsageDaily(ctx, UsageQuery{TenantID: tenantID, From: started.Add(-time.Hour), To: started.Add(2 * time.Hour)}) + if err != nil { + t.Fatal(err) + } + if len(points) != 1 || points[0].RequestCount != 2 || points[0].SuccessfulRequests != 1 || points[0].TotalTokens != 14 || points[0].ChargedMicros != 125 || points[0].P95DurationMS < 120 { + t.Fatalf("unexpected daily usage: %+v", points) + } + + previous := domain.UsageEvent{RequestID: fmt.Sprintf("req_usage_previous_%d", suffix), TenantID: tenantID, ProjectID: projectID, KeyID: keyID, PublicModel: "openai/test", ProviderID: providerID, Protocol: domain.ProtocolOpenAI, StatusCode: 200, Success: true, UsageReported: true, StartedAt: started.Add(-5 * time.Minute), DurationMS: 80, Usage: domain.Usage{InputTokens: 3, OutputTokens: 2, TotalTokens: 5}} + missing := domain.UsageEvent{RequestID: fmt.Sprintf("req_usage_missing_%d", suffix), TenantID: tenantID, ProjectID: projectID, KeyID: keyID, PublicModel: "openai/test", ProviderID: providerID, Protocol: domain.ProtocolOpenAI, StatusCode: 200, Success: true, StartedAt: started.Add(2 * time.Minute), DurationMS: 200} + otherTenantEvent := domain.UsageEvent{RequestID: fmt.Sprintf("req_usage_other_%d", suffix), TenantID: otherTenantID, ProjectID: otherProjectID, KeyID: otherKeyID, PublicModel: "openai/test", ProviderID: providerID, Protocol: domain.ProtocolOpenAI, StatusCode: 200, Success: true, UsageReported: true, StartedAt: started.Add(3 * time.Minute), DurationMS: 900, Usage: domain.Usage{InputTokens: 1000, OutputTokens: 1000, TotalTokens: 2000}} + for _, event := range []domain.UsageEvent{previous, missing, otherTenantEvent} { + if err := store.RecordUsage(ctx, event); err != nil { + t.Fatal(err) + } + } + if _, err := db.Exec(ctx, `UPDATE usage_events SET charged_micros=25,cost_micros=25 WHERE request_id=$1`, previous.RequestID); err != nil { + t.Fatal(err) + } + analytics, err := store.UsageAnalytics(ctx, UsageQuery{TenantID: tenantID, From: started.Add(-time.Minute), To: started.Add(10 * time.Minute)}) + if err != nil { + t.Fatal(err) + } + if len(analytics.Models) != 1 || analytics.Models[0].RequestCount != 3 || analytics.Models[0].SuccessfulRequests != 2 || analytics.Models[0].ProviderCount != 1 || analytics.Models[0].CacheReadInputTokens != 5 || analytics.Models[0].MissingUsageRequests != 1 || analytics.Models[0].ChargedMicros != 125 || analytics.Models[0].PreviousChargedMicros != 25 || analytics.Models[0].ChargeChangePercent == nil || *analytics.Models[0].ChargeChangePercent != 400 { + t.Fatalf("unexpected model analytics: %+v", analytics.Models) + } + if len(analytics.Providers) != 1 || analytics.Providers[0].ProviderID != providerID || analytics.Providers[0].WireAPI != "responses" || analytics.Providers[0].RequestCount != 3 || analytics.Providers[0].MissingUsageRequests != 1 || analytics.Providers[0].P95DurationMS < 200 { + t.Fatalf("unexpected provider analytics: %+v", analytics.Providers) + } + filteredAnalytics, err := store.UsageAnalytics(ctx, UsageQuery{TenantID: tenantID, Provider: providerSlug, Protocol: string(domain.ProtocolOpenAIResponses), ErrorType: "provider_error", Stream: &streaming, From: started.Add(-time.Minute), To: started.Add(10 * time.Minute)}) + if err != nil { + t.Fatal(err) + } + if len(filteredAnalytics.Models) != 1 || filteredAnalytics.Models[0].RequestCount != 1 || filteredAnalytics.Models[0].ErrorCount != 1 || len(filteredAnalytics.Providers) != 1 || filteredAnalytics.Providers[0].RequestCount != 1 { + t.Fatalf("unexpected filtered analytics: %+v", filteredAnalytics) + } +} -- cgit v1.2.3