summaryrefslogtreecommitdiff
path: root/internal/billing/service_test.go
diff options
context:
space:
mode:
authorChia <Chia@93.nz>2026-08-05 22:01:29 +1200
committerChia <Chia@93.nz>2026-08-05 22:07:50 +1200
commiteadb2ffe85c43cf6fc741c9823cd28eedb4a844c (patch)
tree1aba2536d57360da403aa35c9ced58b615c7064e /internal/billing/service_test.go
parentcd0dd91ab93653631904f2ea0e574ccde6d60339 (diff)
feat: harden prepaid billing and commercial operations
Diffstat (limited to 'internal/billing/service_test.go')
-rw-r--r--internal/billing/service_test.go174
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)
}
}