package billing import ( "context" "crypto/sha256" "encoding/json" "fmt" "net/http" "net/http/httptest" "os" "strings" "testing" "time" "aigw/internal/controlplane" "aigw/internal/domain" "github.com/stripe/stripe-go/v86" "github.com/stripe/stripe-go/v86/webhook" ) func TestUsageCostUsesFixedPointAndRoundsOnce(t *testing.T) { usage := domain.Usage{InputTokens: 3, OutputTokens: 2, CacheReadInputTokens: 5} cost, err := usageCost(usage, 150_000, 600_000, 30_000, 0) if err != nil { t.Fatal(err) } // (3*150000 + 2*600000 + 5*30000) / 1e6 = 1.8 micros. if cost != 2 { t.Fatalf("cost = %d, want 2", cost) } } func TestReservationUsesExplicitOutputLimit(t *testing.T) { model := domain.Model{InputPriceMicrosPerMillion: 1_000_000, OutputPriceMicrosPerMillion: 1_000_000} body := []byte(`{"max_tokens":8192}`) cost, err := reservationCost(model, body, 4096) if err != nil { t.Fatal(err) } want := int64(len(body)) + 8192 if cost != want { t.Fatalf("reservation = %d, want %d", cost, want) } } func TestMinorToMicrosSupportsCurrencyExponents(t *testing.T) { tests := []struct { currency string minor int64 want int64 }{{"usd", 123, 1_230_000}, {"jpy", 123, 123_000_000}, {"bhd", 123, 123_000}} for _, test := range tests { got, err := minorToMicros(test.currency, test.minor) if err != nil { t.Fatalf("%s: %v", test.currency, err) } if got != test.want { t.Fatalf("%s: got %d, want %d", test.currency, got, test.want) } } } func TestCollectibleChargePreservesOtherReservations(t *testing.T) { charged, err := collectibleCharge(100, 100, 80, 40) if err != nil { t.Fatal(err) } if charged != 60 { t.Fatalf("charged = %d, want 60", charged) } if remainingBalance, remainingHeld := int64(100)-charged, int64(80)-40; remainingBalance < remainingHeld { t.Fatalf("remaining balance %d does not cover held %d", remainingBalance, remainingHeld) } } func TestIntegrationIdentifierSuffixUsesLetters(t *testing.T) { value := randomLetters(8) if len(value) != 8 { t.Fatalf("length = %d", len(value)) } for _, char := range value { if char < 'a' || char > 'z' { t.Fatalf("non-letter suffix %q", value) } } } func TestCheckoutReturnURLPreservesCallbackAndSessionPlaceholder(t *testing.T) { success := checkoutReturnURL("https://console.example.test/admin/?topup=success", "order-123", true) if !strings.Contains(success, "topup=success") || !strings.Contains(success, "order_id=order-123") || !strings.Contains(success, "session_id={CHECKOUT_SESSION_ID}") { t.Fatalf("unexpected success URL %q", success) } cancel := checkoutReturnURL("https://console.example.test/admin/?topup=cancel&session_id=stale", "order-123", false) if !strings.Contains(cancel, "topup=cancel") || !strings.Contains(cancel, "order_id=order-123") || strings.Contains(cancel, "session_id=") { t.Fatalf("unexpected cancel URL %q", cancel) } } func TestSettlementSpoolWritesOneDurableRecord(t *testing.T) { path := t.TempDir() + "/settlements.jsonl" service := &Service{settlementSpoolPath: path} event := domain.UsageEvent{RequestID: "req_spool_test", TenantID: "tenant", StartedAt: time.Now().UTC()} payload, err := json.Marshal(event) if err != nil { t.Fatal(err) } if err := service.appendSettlementSpool(payload); err != nil { t.Fatal(err) } contents, err := os.ReadFile(path) if err != nil { t.Fatal(err) } if strings.Count(string(contents), "req_spool_test") != 1 || !strings.HasSuffix(string(contents), "\n") { t.Fatalf("unexpected spool contents %q", contents) } } func TestWebhookRejectsInvalidSignatureBeforeProcessing(t *testing.T) { service := &Service{stripeWebhookSecret: "whsec_test"} request := httptest.NewRequest(http.MethodPost, "/billing/stripe/webhook", strings.NewReader(`{"id":"evt_fake"}`)) request.Header.Set("Stripe-Signature", "invalid") response := httptest.NewRecorder() service.WebhookHandler().ServeHTTP(response, request) if response.Code != http.StatusBadRequest { t.Fatalf("status = %d, want 400", response.Code) } } func TestStripeWebhookCreditsPaidOrderExactlyOncePostgres(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 := controlplane.MigrateDatabase(ctx, databaseURL); err != nil { t.Fatal(err) } service, err := New(ctx, Options{ DatabaseURL: databaseURL, Currency: "usd", MinTopUpMinor: 500, MaxTopUpMinor: 1_000_000, StripeEnabled: true, StripeAPIKey: "rk_test_placeholder", StripeWebhookSecret: "whsec_integration_test", }) if err != nil { t.Fatal(err) } t.Cleanup(service.Close) var tenantID string slug := fmt.Sprintf("stripe-%d", time.Now().UnixNano()) if err := service.db.QueryRow(ctx, `INSERT INTO tenants (slug, name) VALUES ($1,'Stripe integration') RETURNING id::text`, slug).Scan(&tenantID); err != nil { t.Fatal(err) } eventID := fmt.Sprintf("evt_aigw_%d", time.Now().UnixNano()) followupEventID := eventID + "_async" t.Cleanup(func() { cleanupCtx := context.Background() for _, statement := range []struct { query string arg string }{ {`DELETE FROM stripe_webhook_events WHERE event_id=$1`, eventID}, {`DELETE FROM stripe_webhook_events WHERE event_id=$1`, followupEventID}, {`DELETE FROM billing_ledger WHERE tenant_id=$1`, tenantID}, {`DELETE FROM tenant_wallets WHERE tenant_id=$1`, tenantID}, {`DELETE FROM topup_orders WHERE tenant_id=$1`, tenantID}, {`DELETE FROM tenants WHERE id=$1`, tenantID}, } { if _, cleanupErr := service.db.Exec(cleanupCtx, statement.query, statement.arg); cleanupErr != nil { t.Errorf("cleanup Stripe integration data: %v", cleanupErr) } } }) const amountMinor int64 = 2500 orderID, amountMicros, err := service.createTopUpOrder(ctx, CheckoutInput{TenantID: tenantID, AmountMinor: amountMinor}) if err != nil { t.Fatal(err) } sessionID := fmt.Sprintf("cs_test_aigw_%d", time.Now().UnixNano()) payload, err := json.Marshal(map[string]any{ "id": eventID, "object": "event", "api_version": stripe.APIVersion, "type": string(stripe.EventTypeCheckoutSessionCompleted), "data": map[string]any{"object": map[string]any{ "id": sessionID, "object": "checkout.session", "client_reference_id": orderID, "amount_total": amountMinor, "currency": "usd", "payment_status": "paid", }}, }) if err != nil { t.Fatal(err) } signed := webhook.GenerateTestSignedPayload(&webhook.UnsignedPayload{Payload: payload, Secret: service.stripeWebhookSecret}) for delivery := 0; delivery < 2; delivery++ { request := httptest.NewRequest(http.MethodPost, "/billing/stripe/webhook", strings.NewReader(string(payload))) request.Header.Set("Stripe-Signature", signed.Header) response := httptest.NewRecorder() service.WebhookHandler().ServeHTTP(response, request) if response.Code != http.StatusOK { t.Fatalf("delivery %d status = %d, body = %s", delivery+1, response.Code, response.Body.String()) } } if _, err := service.db.Exec(ctx, `UPDATE topup_orders SET status='refunded',refunded_micros=amount_micros WHERE id=$1`, orderID); err != nil { t.Fatal(err) } followupPayload, err := json.Marshal(map[string]any{ "id": followupEventID, "object": "event", "api_version": stripe.APIVersion, "type": string(stripe.EventTypeCheckoutSessionAsyncPaymentSucceeded), "data": map[string]any{"object": map[string]any{ "id": sessionID, "object": "checkout.session", "client_reference_id": orderID, "amount_total": amountMinor, "currency": "usd", "payment_status": "paid", }}, }) if err != nil { t.Fatal(err) } followupSigned := webhook.GenerateTestSignedPayload(&webhook.UnsignedPayload{Payload: followupPayload, Secret: service.stripeWebhookSecret}) followupRequest := httptest.NewRequest(http.MethodPost, "/billing/stripe/webhook", strings.NewReader(string(followupPayload))) followupRequest.Header.Set("Stripe-Signature", followupSigned.Header) followupResponse := httptest.NewRecorder() service.WebhookHandler().ServeHTTP(followupResponse, followupRequest) if followupResponse.Code != http.StatusOK { t.Fatalf("follow-up status = %d, body = %s", followupResponse.Code, followupResponse.Body.String()) } var balance int64 if err := service.db.QueryRow(ctx, `SELECT balance_micros FROM tenant_wallets WHERE tenant_id=$1`, tenantID).Scan(&balance); err != nil { t.Fatal(err) } if balance != amountMicros { t.Fatalf("balance = %d, want %d", balance, amountMicros) } var ledgerCount, webhookCount int var orderStatus string if err := service.db.QueryRow(ctx, `SELECT count(*) FROM billing_ledger WHERE source_type='stripe_checkout' AND source_id=$1`, sessionID).Scan(&ledgerCount); err != nil { t.Fatal(err) } if err := service.db.QueryRow(ctx, `SELECT count(*) FROM stripe_webhook_events WHERE event_id IN ($1,$2)`, eventID, followupEventID).Scan(&webhookCount); err != nil { t.Fatal(err) } if err := service.db.QueryRow(ctx, `SELECT status FROM topup_orders WHERE id=$1`, orderID).Scan(&orderStatus); err != nil { t.Fatal(err) } if ledgerCount != 1 || webhookCount != 2 || orderStatus != "refunded" { t.Fatalf("ledger=%d webhook=%d order=%s, want 1/2/refunded", ledgerCount, webhookCount, orderStatus) } } func TestSettlementWorkerPersistsUsageAndReleasesReservationPostgres(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 := controlplane.MigrateDatabase(ctx, databaseURL); err != nil { t.Fatal(err) } service, err := New(ctx, Options{DatabaseURL: databaseURL, Currency: "usd", SettlementSpoolPath: t.TempDir() + "/settlements.jsonl"}) if err != nil { t.Fatal(err) } t.Cleanup(service.Close) slug := fmt.Sprintf("settle-%d", time.Now().UnixNano()) var tenantID, projectID, keyID string if err := service.db.QueryRow(ctx, `INSERT INTO tenants (slug,name) VALUES ($1,'Settlement integration') RETURNING id::text`, slug).Scan(&tenantID); err != nil { t.Fatal(err) } if err := service.db.QueryRow(ctx, `INSERT INTO projects (tenant_id,slug,name) VALUES ($1,'default','Settlement') RETURNING id::text`, tenantID).Scan(&projectID); err != nil { t.Fatal(err) } keyHash := sha256.Sum256([]byte(slug)) if err := service.db.QueryRow(ctx, `INSERT INTO api_keys (tenant_id,project_id,name,key_prefix,key_hash) VALUES ($1,$2,'integration','sk-test',$3) RETURNING id::text`, tenantID, projectID, keyHash[:]).Scan(&keyID); err != nil { t.Fatal(err) } requestID := fmt.Sprintf("req_settle_%d", time.Now().UnixNano()) missingRequestID := requestID + "_missing_usage" t.Cleanup(func() { for _, statement := range []struct { query string arg string }{ {`DELETE FROM billing_ledger WHERE tenant_id=$1`, tenantID}, {`DELETE FROM usage_monthly_rollups WHERE tenant_id=$1`, tenantID}, {`DELETE FROM usage_events WHERE tenant_id=$1`, tenantID}, {`DELETE FROM billing_settlement_jobs WHERE request_id=$1`, requestID}, {`DELETE FROM billing_settlement_jobs WHERE request_id=$1`, missingRequestID}, {`DELETE FROM billing_reservations WHERE request_id=$1`, requestID}, {`DELETE FROM billing_reservations WHERE request_id=$1`, missingRequestID}, {`DELETE FROM tenant_wallets WHERE tenant_id=$1`, tenantID}, {`DELETE FROM api_keys WHERE id=$1`, keyID}, {`DELETE FROM projects WHERE id=$1`, projectID}, {`DELETE FROM tenants WHERE id=$1`, tenantID}, } { if _, cleanupErr := service.db.Exec(context.Background(), statement.query, statement.arg); cleanupErr != nil { t.Errorf("cleanup settlement integration data: %v", cleanupErr) } } }) if _, err := service.db.Exec(ctx, `INSERT INTO tenant_wallets (tenant_id,currency,balance_micros,reserved_micros) VALUES ($1,'usd',1000,20)`, tenantID); err != nil { t.Fatal(err) } if _, err := service.db.Exec(ctx, `INSERT INTO billing_reservations (request_id,tenant_id,project_id,key_id,public_model,currency,reserved_micros,input_price_micros_per_million,output_price_micros_per_million,cache_read_price_micros_per_million,cache_write_price_micros_per_million) VALUES ($1,$2,$3,$4,'demo/model','usd',20,1000000,1000000,0,0)`, requestID, tenantID, projectID, keyID); err != nil { t.Fatal(err) } if _, err := service.db.Exec(ctx, `INSERT INTO billing_settlement_jobs (request_id) VALUES ($1)`, requestID); err != nil { t.Fatal(err) } if err := service.EnqueueSettlement(ctx, domain.UsageEvent{RequestID: requestID, TenantID: tenantID, ProjectID: projectID, KeyID: keyID, PublicModel: "demo/model", Protocol: domain.ProtocolOpenAI, StatusCode: 200, Success: true, UsageReported: true, StartedAt: time.Now().UTC(), Usage: domain.Usage{InputTokens: 10}}); err != nil { t.Fatal(err) } processed, err := service.processSettlementJob(ctx) if err != nil || !processed { t.Fatalf("process settlement: processed=%v err=%v", processed, err) } var reservationStatus, jobStatus string var balance, reserved, charged, usageCount int64 if err := service.db.QueryRow(ctx, `SELECT status,charged_micros FROM billing_reservations WHERE request_id=$1`, requestID).Scan(&reservationStatus, &charged); err != nil { t.Fatal(err) } if err := service.db.QueryRow(ctx, `SELECT status FROM billing_settlement_jobs WHERE request_id=$1`, requestID).Scan(&jobStatus); err != nil { t.Fatal(err) } if err := service.db.QueryRow(ctx, `SELECT balance_micros,reserved_micros FROM tenant_wallets WHERE tenant_id=$1`, tenantID).Scan(&balance, &reserved); err != nil { t.Fatal(err) } if err := service.db.QueryRow(ctx, `SELECT count(*) FROM usage_events WHERE request_id=$1`, requestID).Scan(&usageCount); err != nil { t.Fatal(err) } if reservationStatus != "settled" || jobStatus != "done" || balance != 990 || reserved != 0 || charged != 10 || usageCount != 1 { t.Fatalf("reservation=%s job=%s balance=%d reserved=%d charged=%d usage=%d", reservationStatus, jobStatus, balance, reserved, charged, usageCount) } if _, err := service.db.Exec(ctx, `UPDATE tenant_wallets SET reserved_micros=reserved_micros+30 WHERE tenant_id=$1`, tenantID); err != nil { t.Fatal(err) } if _, err := service.db.Exec(ctx, `INSERT INTO billing_reservations (request_id,tenant_id,project_id,key_id,public_model,currency,reserved_micros,input_price_micros_per_million,output_price_micros_per_million,cache_read_price_micros_per_million,cache_write_price_micros_per_million) VALUES ($1,$2,$3,$4,'demo/model','usd',30,1000000,1000000,0,0)`, missingRequestID, tenantID, projectID, keyID); err != nil { t.Fatal(err) } if _, err := service.db.Exec(ctx, `INSERT INTO billing_settlement_jobs (request_id) VALUES ($1)`, missingRequestID); err != nil { t.Fatal(err) } if err := service.EnqueueSettlement(ctx, domain.UsageEvent{RequestID: missingRequestID, TenantID: tenantID, ProjectID: projectID, KeyID: keyID, PublicModel: "demo/model", Protocol: domain.ProtocolOpenAI, StatusCode: 200, Success: true, UsageReported: false, StartedAt: time.Now().UTC()}); err != nil { t.Fatal(err) } processed, err = service.processSettlementJob(ctx) if err != nil || !processed { t.Fatalf("process missing-usage settlement: processed=%v err=%v", processed, err) } var missingReservationStatus, missingJobStatus, meteringStatus string if err := service.db.QueryRow(ctx, `SELECT status FROM billing_reservations WHERE request_id=$1`, missingRequestID).Scan(&missingReservationStatus); err != nil { t.Fatal(err) } if err := service.db.QueryRow(ctx, `SELECT status FROM billing_settlement_jobs WHERE request_id=$1`, missingRequestID).Scan(&missingJobStatus); err != nil { t.Fatal(err) } if err := service.db.QueryRow(ctx, `SELECT balance_micros,reserved_micros FROM tenant_wallets WHERE tenant_id=$1`, tenantID).Scan(&balance, &reserved); err != nil { t.Fatal(err) } if err := service.db.QueryRow(ctx, `SELECT metering_status FROM usage_events WHERE request_id=$1`, missingRequestID).Scan(&meteringStatus); err != nil { t.Fatal(err) } if missingReservationStatus != "metering_failed" || missingJobStatus != "done" || balance != 990 || reserved != 30 || meteringStatus != "missing" { t.Fatalf("missing usage reservation=%s job=%s balance=%d reserved=%d metering=%s", missingReservationStatus, missingJobStatus, balance, reserved, meteringStatus) } }