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, TTFTMS: 47, 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.Data) != 1 || records.Data[0].RequestID != events[0].RequestID || records.Data[0].ProjectName != "Production" || records.Data[0].KeyName != "Production key" || records.Data[0].MeteringStatus != "reported" || records.Data[0].TTFTMS != 47 { 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.Data) != 1 || records.Data[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].P50DurationMS != 235 || points[0].P95DurationMS < 120 || points[0].P50TTFTMS != 47 || points[0].P95TTFTMS != 47 { 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, TTFTMS: 30, 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) } } requestDetail, err := store.ListUsage(ctx, UsageQuery{TenantID: tenantID, RequestID: events[0].RequestID, Limit: 1}) if err != nil { t.Fatal(err) } if len(requestDetail.Data) != 1 || requestDetail.Data[0].RequestID != events[0].RequestID { t.Fatalf("tenant could not retrieve its ledger-linked request: %+v", requestDetail) } crossTenantDetail, err := store.ListUsage(ctx, UsageQuery{TenantID: tenantID, RequestID: otherTenantEvent.RequestID, Limit: 1}) if err != nil { t.Fatal(err) } if len(crossTenantDetail.Data) != 0 { t.Fatalf("tenant retrieved another tenant's request by request id: %+v", crossTenantDetail) } firstPage, err := store.ListUsage(ctx, UsageQuery{TenantID: tenantID, From: started.Add(-time.Hour), To: started.Add(10 * time.Minute), Limit: 1}) if err != nil { t.Fatal(err) } if len(firstPage.Data) != 1 || firstPage.NextCursor == "" { t.Fatalf("expected a cursor for the first usage page: %+v", firstPage) } secondPage, err := store.ListUsage(ctx, UsageQuery{TenantID: tenantID, From: started.Add(-time.Hour), To: started.Add(10 * time.Minute), Limit: 1, Cursor: firstPage.NextCursor}) if err != nil { t.Fatal(err) } if len(secondPage.Data) != 1 || secondPage.Data[0].RequestID == firstPage.Data[0].RequestID { t.Fatalf("cursor did not advance usage page: first=%+v second=%+v", firstPage, secondPage) } 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].P50DurationMS != 200 || analytics.Models[0].P95DurationMS < 200 || analytics.Models[0].P50TTFTMS != 47 || analytics.Models[0].P95TTFTMS != 47 || 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.Keys) != 1 || analytics.Keys[0].KeyID != keyID || analytics.Keys[0].RequestCount != 3 || analytics.Keys[0].ChargedMicros != 125 { t.Fatalf("unexpected key analytics: %+v", analytics.Keys) } 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) } }