diff options
| author | Chia <Chia@93.nz> | 2026-08-05 22:01:29 +1200 |
|---|---|---|
| committer | Chia <Chia@93.nz> | 2026-08-05 22:07:50 +1200 |
| commit | eadb2ffe85c43cf6fc741c9823cd28eedb4a844c (patch) | |
| tree | 1aba2536d57360da403aa35c9ced58b615c7064e /internal/billing/service_test.go | |
| parent | cd0dd91ab93653631904f2ea0e574ccde6d60339 (diff) | |
feat: harden prepaid billing and commercial operations
Diffstat (limited to 'internal/billing/service_test.go')
| -rw-r--r-- | internal/billing/service_test.go | 174 |
1 files changed, 171 insertions, 3 deletions
diff --git a/internal/billing/service_test.go b/internal/billing/service_test.go index f675da7..6a1bb49 100644 --- a/internal/billing/service_test.go +++ b/internal/billing/service_test.go @@ -2,6 +2,7 @@ package billing import ( "context" + "crypto/sha256" "encoding/json" "fmt" "net/http" @@ -97,6 +98,26 @@ func TestCheckoutReturnURLPreservesCallbackAndSessionPlaceholder(t *testing.T) { } } +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"}`)) @@ -133,6 +154,7 @@ func TestStripeWebhookCreditsPaidOrderExactlyOncePostgres(t *testing.T) { 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 { @@ -140,6 +162,7 @@ func TestStripeWebhookCreditsPaidOrderExactlyOncePostgres(t *testing.T) { 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}, @@ -177,6 +200,28 @@ func TestStripeWebhookCreditsPaidOrderExactlyOncePostgres(t *testing.T) { 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 { @@ -190,13 +235,136 @@ func TestStripeWebhookCreditsPaidOrderExactlyOncePostgres(t *testing.T) { 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=$1`, eventID).Scan(&webhookCount); err != nil { + 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 != 1 || orderStatus != "paid" { - t.Fatalf("ledger=%d webhook=%d order=%s, want 1/1/paid", ledgerCount, webhookCount, orderStatus) + 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) } } |
