summaryrefslogtreecommitdiff
path: root/internal/controlplane/usage_integration_test.go
diff options
context:
space:
mode:
Diffstat (limited to 'internal/controlplane/usage_integration_test.go')
-rw-r--r--internal/controlplane/usage_integration_test.go152
1 files changed, 152 insertions, 0 deletions
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)
+ }
+}