summaryrefslogtreecommitdiff
path: root/internal/billing
diff options
context:
space:
mode:
Diffstat (limited to '')
-rw-r--r--internal/billing/auto_topup.go56
-rw-r--r--internal/billing/auto_topup_test.go43
-rw-r--r--internal/billing/ledger.go77
-rw-r--r--internal/billing/operations.go8
-rw-r--r--internal/billing/operations_test.go58
-rw-r--r--internal/billing/service.go115
-rw-r--r--internal/billing/service_test.go174
-rw-r--r--internal/billing/stripe.go71
-rw-r--r--internal/billing/stripe_preflight.go124
-rw-r--r--internal/billing/stripe_preflight_test.go90
-rw-r--r--internal/billing/types.go39
11 files changed, 756 insertions, 99 deletions
diff --git a/internal/billing/auto_topup.go b/internal/billing/auto_topup.go
index a90405b..7fa4e59 100644
--- a/internal/billing/auto_topup.go
+++ b/internal/billing/auto_topup.go
@@ -160,6 +160,23 @@ func (s *Service) CreateAutoTopUpSetupSession(ctx context.Context, input AutoTop
if err != nil {
return AutoTopUpSetupResult{}, err
}
+ params := s.autoTopUpSetupSessionParams(input, customerID)
+ session, err := s.createStripeCheckout(ctx, params)
+ if err != nil {
+ return AutoTopUpSetupResult{}, fmt.Errorf("create automatic top-up setup session: %w", err)
+ }
+ if session == nil || session.ID == "" || session.URL == "" {
+ return AutoTopUpSetupResult{}, errors.New("Stripe returned an incomplete setup session")
+ }
+ if _, err := s.db.Exec(ctx, `INSERT INTO tenant_auto_topup_settings (tenant_id,threshold_micros,topup_amount_minor,stripe_setup_session_id)
+ VALUES ($1,$2,$3,$4) ON CONFLICT (tenant_id) DO UPDATE SET stripe_setup_session_id=EXCLUDED.stripe_setup_session_id,updated_at=now()`,
+ input.TenantID, s.defaultAutoTopUpThresholdMicros(), s.defaultAutoTopUpAmountMinor(), session.ID); err != nil {
+ return AutoTopUpSetupResult{}, fmt.Errorf("persist automatic top-up setup session: %w", err)
+ }
+ return AutoTopUpSetupResult{SessionID: session.ID, URL: session.URL}, nil
+}
+
+func (s *Service) autoTopUpSetupSessionParams(input AutoTopUpSetupInput, customerID string) *stripe.CheckoutSessionCreateParams {
params := &stripe.CheckoutSessionCreateParams{
Mode: stripe.String(string(stripe.CheckoutSessionModeSetup)),
Currency: stripe.String(s.currency),
@@ -181,19 +198,7 @@ func (s *Service) CreateAutoTopUpSetupSession(ctx context.Context, input AutoTop
}
}
params.SetIdempotencyKey("aigw_autotopup_setup_" + randomHex(16))
- session, err := s.createStripeCheckout(ctx, params)
- if err != nil {
- return AutoTopUpSetupResult{}, fmt.Errorf("create automatic top-up setup session: %w", err)
- }
- if session.ID == "" || session.URL == "" {
- return AutoTopUpSetupResult{}, errors.New("Stripe returned an incomplete setup session")
- }
- if _, err := s.db.Exec(ctx, `INSERT INTO tenant_auto_topup_settings (tenant_id,threshold_micros,topup_amount_minor,stripe_setup_session_id)
- VALUES ($1,$2,$3,$4) ON CONFLICT (tenant_id) DO UPDATE SET stripe_setup_session_id=EXCLUDED.stripe_setup_session_id,updated_at=now()`,
- input.TenantID, s.defaultAutoTopUpThresholdMicros(), s.defaultAutoTopUpAmountMinor(), session.ID); err != nil {
- return AutoTopUpSetupResult{}, fmt.Errorf("persist automatic top-up setup session: %w", err)
- }
- return AutoTopUpSetupResult{SessionID: session.ID, URL: session.URL}, nil
+ return params
}
func autoTopUpReturnURL(raw string, success bool) string {
@@ -344,16 +349,7 @@ func (s *Service) processAutoTopUpOnce(ctx context.Context) (bool, error) {
if err := tx.Commit(ctx); err != nil {
return false, err
}
- params := &stripe.PaymentIntentCreateParams{
- Amount: stripe.Int64(amountMinor), Currency: stripe.String(currency), Customer: stripe.String(customerID),
- PaymentMethod: stripe.String(paymentMethodID), Confirm: stripe.Bool(true), OffSession: stripe.Bool(true),
- ErrorOnRequiresAction: stripe.Bool(true), Description: stripe.String("AIGW automatic prepaid balance top-up"),
- Metadata: map[string]string{"aigw_action": "auto_topup", "aigw_tenant_id": tenantID, "aigw_topup_order_id": orderID},
- }
- if strings.TrimSpace(customerEmail) != "" {
- params.ReceiptEmail = stripe.String(strings.TrimSpace(customerEmail))
- }
- params.SetIdempotencyKey("aigw_autotopup_" + orderID)
+ params := s.autoTopUpPaymentIntentParams(tenantID, orderID, customerID, customerEmail, paymentMethodID, currency, amountMinor)
intent, callErr := s.createStripePaymentIntent(ctx, params)
if callErr != nil {
var stripeErr *stripe.Error
@@ -387,6 +383,20 @@ func (s *Service) processAutoTopUpOnce(ctx context.Context) (bool, error) {
return true, s.failAutoTopUp(ctx, tenantID, orderID, fmt.Errorf("automatic top-up PaymentIntent ended in status %s", intent.Status))
}
+func (s *Service) autoTopUpPaymentIntentParams(tenantID, orderID, customerID, customerEmail, paymentMethodID, currency string, amountMinor int64) *stripe.PaymentIntentCreateParams {
+ params := &stripe.PaymentIntentCreateParams{
+ Amount: stripe.Int64(amountMinor), Currency: stripe.String(currency), Customer: stripe.String(customerID),
+ PaymentMethod: stripe.String(paymentMethodID), Confirm: stripe.Bool(true), OffSession: stripe.Bool(true),
+ ErrorOnRequiresAction: stripe.Bool(true), Description: stripe.String("AIGW automatic prepaid balance top-up"),
+ Metadata: map[string]string{"aigw_action": "auto_topup", "aigw_tenant_id": tenantID, "aigw_topup_order_id": orderID},
+ }
+ if strings.TrimSpace(customerEmail) != "" {
+ params.ReceiptEmail = stripe.String(strings.TrimSpace(customerEmail))
+ }
+ params.SetIdempotencyKey("aigw_autotopup_" + orderID)
+ return params
+}
+
func (s *Service) scheduleAutoTopUpRetry(ctx context.Context, tenantID, orderID string, cause error) error {
message := truncateError(cause)
_, err := s.db.Exec(ctx, `UPDATE topup_orders SET reconciliation_error=$2 WHERE id=$1 AND status='pending'`, orderID, message)
diff --git a/internal/billing/auto_topup_test.go b/internal/billing/auto_topup_test.go
index 85edb3d..55dc22d 100644
--- a/internal/billing/auto_topup_test.go
+++ b/internal/billing/auto_topup_test.go
@@ -25,6 +25,49 @@ func TestAutoTopUpReturnURL(t *testing.T) {
}
}
+func TestAutoTopUpStripeContracts(t *testing.T) {
+ service := &Service{
+ currency: "usd",
+ stripeSuccessURL: "https://console.example.test/billing?topup=success",
+ stripeCancelURL: "https://console.example.test/billing?topup=cancel",
+ integrationIdentifier: "aigw_balance_abcdefgh",
+ }
+ setup := service.autoTopUpSetupSessionParams(AutoTopUpSetupInput{
+ TenantID: "tenant-123", CustomerEmail: " billing@example.test ",
+ }, "")
+ if setup.Mode == nil || *setup.Mode != string(stripe.CheckoutSessionModeSetup) || setup.Currency == nil || *setup.Currency != "usd" {
+ t.Fatalf("unexpected setup contract %+v", setup)
+ }
+ if len(setup.PaymentMethodTypes) != 0 || setup.CustomerCreation == nil || *setup.CustomerCreation != string(stripe.CheckoutSessionCustomerCreationAlways) {
+ t.Fatal("setup Checkout must create a customer and use Dashboard-managed payment methods")
+ }
+ if setup.CustomerEmail == nil || *setup.CustomerEmail != "billing@example.test" || setup.Metadata["aigw_action"] != autoTopUpAction {
+ t.Fatal("setup Checkout customer or metadata contract is incomplete")
+ }
+ if setup.IdempotencyKey == nil || !strings.HasPrefix(*setup.IdempotencyKey, "aigw_autotopup_setup_") {
+ t.Fatalf("setup idempotency key = %v", setup.IdempotencyKey)
+ }
+
+ payment := service.autoTopUpPaymentIntentParams(
+ "tenant-123", "order-123", "cus_123", " billing@example.test ", "pm_123", "usd", 2000,
+ )
+ if payment.Amount == nil || *payment.Amount != 2000 || payment.Currency == nil || *payment.Currency != "usd" ||
+ payment.Customer == nil || *payment.Customer != "cus_123" || payment.PaymentMethod == nil || *payment.PaymentMethod != "pm_123" {
+ t.Fatalf("unexpected automatic top-up PaymentIntent %+v", payment)
+ }
+ if payment.Confirm == nil || !*payment.Confirm || payment.OffSession == nil || !*payment.OffSession ||
+ payment.ErrorOnRequiresAction == nil || !*payment.ErrorOnRequiresAction {
+ t.Fatal("automatic top-up must be confirmed off-session and stop on required customer action")
+ }
+ if payment.ReceiptEmail == nil || *payment.ReceiptEmail != "billing@example.test" ||
+ payment.Metadata["aigw_topup_order_id"] != "order-123" {
+ t.Fatal("automatic top-up receipt or reconciliation metadata is incomplete")
+ }
+ if payment.IdempotencyKey == nil || *payment.IdempotencyKey != "aigw_autotopup_order-123" {
+ t.Fatalf("payment idempotency key = %v", payment.IdempotencyKey)
+ }
+}
+
func TestAutomaticTopUpSetupAndCreditAreIdempotentPostgres(t *testing.T) {
databaseURL := os.Getenv("AIGW_TEST_DATABASE_URL")
if databaseURL == "" {
diff --git a/internal/billing/ledger.go b/internal/billing/ledger.go
index 7a00081..64f258c 100644
--- a/internal/billing/ledger.go
+++ b/internal/billing/ledger.go
@@ -190,6 +190,83 @@ func (s *Service) AdjustBalance(ctx context.Context, input AdjustmentInput) (Led
return result, nil
}
+// ReleaseUnmeteredReservation is an audited operational escape hatch for a
+// fail-closed success that cannot be reconciled. It never invents usage or
+// changes wallet balance; it only returns the existing hold to availability.
+func (s *Service) ReleaseUnmeteredReservation(ctx context.Context, tenantID, requestID string, input ReleaseReservationInput, actor ResolutionActor) (ReservationRelease, error) {
+ tenantID = strings.TrimSpace(tenantID)
+ requestID = strings.TrimSpace(requestID)
+ reason := normalizeDescription(input.Reason)
+ if tenantID == "" || requestID == "" || reason == "" || actor.ID == "" || actor.Type == "" {
+ return ReservationRelease{}, fmt.Errorf("%w: tenant, request, reason, and actor are required", ErrReservationNotReleasable)
+ }
+ tx, err := s.db.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.ReadCommitted})
+ if err != nil {
+ return ReservationRelease{}, fmt.Errorf("begin reservation release: %w", err)
+ }
+ defer tx.Rollback(ctx)
+ var result ReservationRelease
+ var projectID, currency, status string
+ if err := tx.QueryRow(ctx, `SELECT request_id,tenant_id::text,project_id::text,currency,reserved_micros,status
+ FROM billing_reservations WHERE request_id=$1 AND tenant_id=$2 FOR UPDATE`, requestID, tenantID).Scan(
+ &result.RequestID, &result.TenantID, &projectID, &currency, &result.ReservedMicros, &status); errors.Is(err, pgx.ErrNoRows) {
+ return ReservationRelease{}, ErrReservationNotReleasable
+ } else if err != nil {
+ return ReservationRelease{}, fmt.Errorf("lock reservation for release: %w", err)
+ }
+ if status == "released" {
+ var evidenceExists bool
+ if err := tx.QueryRow(ctx, `SELECT EXISTS(SELECT 1 FROM billing_ledger
+ WHERE source_type='unmetered_reservation' AND source_id=$1)`, requestID).Scan(&evidenceExists); err != nil {
+ return ReservationRelease{}, fmt.Errorf("read reservation release evidence: %w", err)
+ }
+ if !evidenceExists {
+ return ReservationRelease{}, fmt.Errorf("%w: reservation was released by normal settlement", ErrReservationNotReleasable)
+ }
+ if err := tx.QueryRow(ctx, `SELECT COALESCE(settled_at,created_at) FROM billing_reservations WHERE request_id=$1`, requestID).Scan(&result.ReleasedAt); err != nil {
+ return ReservationRelease{}, err
+ }
+ result.Status = status
+ return result, tx.Commit(ctx)
+ }
+ if status != "metering_failed" {
+ return ReservationRelease{}, fmt.Errorf("%w: reservation status is %s", ErrReservationNotReleasable, status)
+ }
+ var balance, held int64
+ if err := tx.QueryRow(ctx, `SELECT balance_micros,reserved_micros FROM tenant_wallets WHERE tenant_id=$1 FOR UPDATE`, tenantID).Scan(&balance, &held); err != nil {
+ return ReservationRelease{}, fmt.Errorf("lock wallet for reservation release: %w", err)
+ }
+ if result.ReservedMicros > held {
+ return ReservationRelease{}, errors.New("wallet reservation invariant violated during release")
+ }
+ if _, err := tx.Exec(ctx, `UPDATE tenant_wallets SET reserved_micros=reserved_micros-$2,updated_at=now() WHERE tenant_id=$1`, tenantID, result.ReservedMicros); err != nil {
+ return ReservationRelease{}, fmt.Errorf("release wallet hold: %w", err)
+ }
+ if err := tx.QueryRow(ctx, `UPDATE billing_reservations SET status='released',settled_at=now()
+ WHERE request_id=$1 RETURNING status,settled_at`, requestID).Scan(&result.Status, &result.ReleasedAt); err != nil {
+ return ReservationRelease{}, fmt.Errorf("mark reservation released: %w", err)
+ }
+ command, err := tx.Exec(ctx, `UPDATE usage_events SET metering_status='released_unmetered'
+ WHERE request_id=$1 AND tenant_id=$2 AND metering_status='missing' AND usage_reported=FALSE`, requestID, tenantID)
+ if err != nil {
+ return ReservationRelease{}, fmt.Errorf("mark unmetered usage resolved: %w", err)
+ }
+ if command.RowsAffected() != 1 {
+ return ReservationRelease{}, errors.New("unmetered usage invariant violated during release")
+ }
+ description := fmt.Sprintf("Unmetered reservation released by %s %s: %s", actor.Type, actor.ID, reason)
+ if _, err := tx.Exec(ctx, `INSERT INTO billing_ledger
+ (tenant_id,project_id,currency,amount_micros,balance_after_micros,kind,source_type,source_id,description)
+ VALUES ($1,$2,$3,0,$4,'release','unmetered_reservation',$5,$6)
+ ON CONFLICT (source_type,source_id) DO NOTHING`, tenantID, projectID, currency, balance, requestID, description); err != nil {
+ return ReservationRelease{}, fmt.Errorf("write reservation release evidence: %w", err)
+ }
+ if err := tx.Commit(ctx); err != nil {
+ return ReservationRelease{}, fmt.Errorf("commit reservation release: %w", err)
+ }
+ return result, nil
+}
+
func (s *Service) createTopUpOrder(ctx context.Context, input CheckoutInput) (string, int64, error) {
if !s.stripeEnabled {
return "", 0, ErrStripeDisabled
diff --git a/internal/billing/operations.go b/internal/billing/operations.go
index 461c59a..508ebd5 100644
--- a/internal/billing/operations.go
+++ b/internal/billing/operations.go
@@ -15,8 +15,10 @@ import (
"github.com/stripe/stripe-go/v86"
)
+type stripePortalSessionCreator func(context.Context, *stripe.BillingPortalSessionCreateParams) (*stripe.BillingPortalSession, error)
+
func (s *Service) CreatePortalSession(ctx context.Context, tenantID string) (PortalResult, error) {
- if !s.stripeEnabled || s.stripeClient == nil {
+ if !s.stripeEnabled || s.createStripePortalSession == nil {
return PortalResult{}, ErrStripeDisabled
}
customerID, err := s.ensureStripeCustomer(ctx, tenantID)
@@ -26,13 +28,13 @@ func (s *Service) CreatePortalSession(ctx context.Context, tenantID string) (Por
if customerID == "" {
return PortalResult{}, errors.New("no Stripe customer exists for this account")
}
- session, err := s.stripeClient.V1BillingPortalSessions.Create(ctx, &stripe.BillingPortalSessionCreateParams{
+ session, err := s.createStripePortalSession(ctx, &stripe.BillingPortalSessionCreateParams{
Customer: stripe.String(customerID), ReturnURL: stripe.String(s.stripePortalReturnURL),
})
if err != nil {
return PortalResult{}, fmt.Errorf("create Stripe customer portal session: %w", err)
}
- if session.URL == "" {
+ if session == nil || session.URL == "" {
return PortalResult{}, errors.New("Stripe returned an incomplete portal session")
}
return PortalResult{URL: session.URL}, nil
diff --git a/internal/billing/operations_test.go b/internal/billing/operations_test.go
index 46bedeb..4d90383 100644
--- a/internal/billing/operations_test.go
+++ b/internal/billing/operations_test.go
@@ -1,8 +1,15 @@
package billing
import (
+ "context"
+ "fmt"
+ "os"
"testing"
"time"
+
+ "aigw/internal/controlplane"
+
+ "github.com/stripe/stripe-go/v86"
)
func TestOperationalStatusReadiness(t *testing.T) {
@@ -34,3 +41,54 @@ func TestOperationalStatusReadiness(t *testing.T) {
t.Fatal("unmetered success must fail readiness")
}
}
+
+func TestCustomerPortalSessionContractPostgres(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", StripeEnabled: true, StripeAPIKey: "rk_test_placeholder",
+ StripePortalReturnURL: "https://console.example.test/billing",
+ })
+ if err != nil {
+ t.Fatal(err)
+ }
+ t.Cleanup(service.Close)
+
+ suffix := time.Now().UnixNano()
+ var tenantID string
+ if err := service.db.QueryRow(ctx, `INSERT INTO tenants (slug,name) VALUES ($1,'Portal contract') RETURNING id::text`, fmt.Sprintf("portal-%d", suffix)).Scan(&tenantID); err != nil {
+ t.Fatal(err)
+ }
+ t.Cleanup(func() {
+ if _, cleanupErr := service.db.Exec(context.Background(), `DELETE FROM tenants WHERE id=$1`, tenantID); cleanupErr != nil {
+ t.Errorf("cleanup portal contract tenant: %v", cleanupErr)
+ }
+ })
+ customerID := fmt.Sprintf("cus_portal_%d", suffix)
+ if _, err := service.db.Exec(ctx, `INSERT INTO stripe_customers (tenant_id,stripe_customer_id,email) VALUES ($1,$2,'billing@example.test')`, tenantID, customerID); err != nil {
+ t.Fatal(err)
+ }
+ service.createStripePortalSession = func(_ context.Context, params *stripe.BillingPortalSessionCreateParams) (*stripe.BillingPortalSession, error) {
+ if params.Customer == nil || *params.Customer != customerID || params.ReturnURL == nil || *params.ReturnURL != service.stripePortalReturnURL {
+ t.Fatalf("unexpected Portal params %+v", params)
+ }
+ return &stripe.BillingPortalSession{URL: "https://billing.stripe.test/session"}, nil
+ }
+ result, err := service.CreatePortalSession(ctx, tenantID)
+ if err != nil || result.URL != "https://billing.stripe.test/session" {
+ t.Fatalf("CreatePortalSession result=%+v err=%v", result, err)
+ }
+ service.createStripePortalSession = func(context.Context, *stripe.BillingPortalSessionCreateParams) (*stripe.BillingPortalSession, error) {
+ return nil, nil
+ }
+ if _, err := service.CreatePortalSession(ctx, tenantID); err == nil {
+ t.Fatal("incomplete Stripe Portal response was accepted")
+ }
+}
diff --git a/internal/billing/service.go b/internal/billing/service.go
index 8a91839..8a5f266 100644
--- a/internal/billing/service.go
+++ b/internal/billing/service.go
@@ -39,6 +39,7 @@ type Service struct {
stripeProductTaxCode string
integrationIdentifier string
createStripeCheckout stripeCheckoutCreator
+ createStripePortalSession stripePortalSessionCreator
createStripeCustomer stripeCustomerCreator
updateStripeCustomer stripeCustomerUpdater
retrieveStripeSetupIntent stripeSetupIntentRetriever
@@ -73,6 +74,7 @@ func New(ctx context.Context, options Options) (*Service, error) {
if options.StripeEnabled {
service.stripeClient = stripe.NewClient(options.StripeAPIKey)
service.createStripeCheckout = service.stripeClient.V1CheckoutSessions.Create
+ service.createStripePortalSession = service.stripeClient.V1BillingPortalSessions.Create
service.createStripeCustomer = service.stripeClient.V1Customers.Create
service.updateStripeCustomer = service.stripeClient.V1Customers.Update
service.retrieveStripeSetupIntent = service.stripeClient.V1SetupIntents.Retrieve
@@ -103,7 +105,7 @@ func (s *Service) Authorize(ctx context.Context, input Authorization) error {
if input.Model.PriceCurrency != "" && input.Model.PriceCurrency != s.currency {
return fmt.Errorf("model price currency %s does not match wallet currency %s", input.Model.PriceCurrency, s.currency)
}
- reserved, err := reservationCost(input.Model, input.Body, s.defaultMaxOutputTokens)
+ reserved, err := reservationCost(input.Model, input.Body, s.defaultMaxOutputTokens, input.Protocol)
if err != nil {
return err
}
@@ -156,6 +158,22 @@ func (s *Service) Authorize(ctx context.Context, input Authorization) error {
return ErrQuotaExceeded
}
}
+ if input.Principal.DailySpendMicros > 0 {
+ now := time.Now().UTC()
+ period := time.Date(now.Year(), now.Month(), now.Day(), 0, 0, 0, 0, time.UTC)
+ nextPeriod := period.AddDate(0, 0, 1)
+ var used, pending int64
+ if err := tx.QueryRow(ctx, `SELECT
+ COALESCE((SELECT sum(cost_micros) FROM usage_events WHERE key_id=$1 AND started_at >= $2 AND started_at < $3),0),
+ COALESCE((SELECT sum(reserved_micros) FROM billing_reservations WHERE key_id=$1 AND status IN ('pending','metering_failed') AND created_at >= $2 AND created_at < $3),0)`,
+ input.Principal.KeyID, period, nextPeriod).Scan(&used, &pending); err != nil {
+ return fmt.Errorf("read API key daily spend quota: %w", err)
+ }
+ limit := input.Principal.DailySpendMicros
+ if reserved > limit || used > limit-reserved || pending > limit-used-reserved {
+ return ErrDailyQuotaExceeded
+ }
+ }
if balance-held < reserved {
return ErrInsufficientBalance
}
@@ -250,15 +268,15 @@ func (s *Service) settle(ctx context.Context, event domain.UsageEvent) error {
}
if _, err := tx.Exec(ctx, `INSERT INTO usage_events (
request_id,tenant_id,project_id,key_id,public_model,provider_id,upstream_model,protocol,stream,
- status_code,success,error_type,attempts,started_at,duration_ms,input_tokens,output_tokens,total_tokens,
+ status_code,success,error_type,attempts,started_at,duration_ms,ttft_ms,input_tokens,output_tokens,total_tokens,
cache_creation_input_tokens,cache_read_input_tokens,cost_micros,charged_micros,uncollected_micros,
usage_reported,metering_status)
VALUES ($1,$2,$3,$4,$5,NULLIF($6,''),NULLIF($7,''),$8,$9,$10,$11,'usage_not_reported',$12,$13,$14,
- $15,$16,$17,$18,$19,0,0,0,false,'missing')
- ON CONFLICT (request_id) DO UPDATE SET error_type='usage_not_reported',usage_reported=false,metering_status='missing'`,
+ $15,$16,$17,$18,$19,$20,0,0,0,false,'missing')
+ ON CONFLICT (request_id) DO UPDATE SET error_type='usage_not_reported',ttft_ms=EXCLUDED.ttft_ms,usage_reported=false,metering_status='missing'`,
event.RequestID, tenantID, projectID, keyID, modelID, event.ProviderID, event.UpstreamModel,
string(event.Protocol), event.Stream, event.StatusCode, event.Success, event.Attempts, event.StartedAt,
- event.DurationMS, event.Usage.InputTokens, event.Usage.OutputTokens, event.Usage.TotalTokens,
+ event.DurationMS, event.TTFTMS, event.Usage.InputTokens, event.Usage.OutputTokens, event.Usage.TotalTokens,
event.Usage.CacheCreationInputTokens, event.Usage.CacheReadInputTokens); err != nil {
return fmt.Errorf("persist unmetered usage event: %w", err)
}
@@ -316,21 +334,21 @@ func (s *Service) settle(ctx context.Context, event domain.UsageEvent) error {
}
if usageAlreadyRecorded {
if _, err := tx.Exec(ctx, `UPDATE usage_events SET cost_micros=$2, charged_micros=$3, uncollected_micros=$4,
- usage_reported=$5,metering_status=$6 WHERE request_id=$1`,
- event.RequestID, actualCost, charged, uncollected, event.UsageReported, meteringStatus(event)); err != nil {
+ ttft_ms=GREATEST(ttft_ms,$5),usage_reported=$6,metering_status=$7 WHERE request_id=$1`,
+ event.RequestID, actualCost, charged, uncollected, event.TTFTMS, event.UsageReported, meteringStatus(event)); err != nil {
return fmt.Errorf("apply usage charge: %w", err)
}
} else if _, err := tx.Exec(ctx, `
INSERT INTO usage_events (
request_id, tenant_id, project_id, key_id, public_model, provider_id, upstream_model,
protocol, stream, status_code, success, error_type, attempts, started_at, duration_ms,
- input_tokens, output_tokens, total_tokens, cache_creation_input_tokens, cache_read_input_tokens,
+ ttft_ms, input_tokens, output_tokens, total_tokens, cache_creation_input_tokens, cache_read_input_tokens,
cost_micros, charged_micros, uncollected_micros, usage_reported, metering_status)
- VALUES ($1,$2,$3,$4,$5,NULLIF($6,''),NULLIF($7,''),$8,$9,$10,$11,$12,$13,$14,$15,$16,$17,$18,$19,$20,$21,$22,$23,$24,$25)
+ VALUES ($1,$2,$3,$4,$5,NULLIF($6,''),NULLIF($7,''),$8,$9,$10,$11,$12,$13,$14,$15,$16,$17,$18,$19,$20,$21,$22,$23,$24,$25,$26)
ON CONFLICT (request_id) DO NOTHING`,
event.RequestID, tenantID, projectID, keyID, modelID, event.ProviderID, event.UpstreamModel,
string(event.Protocol), event.Stream, event.StatusCode, event.Success, event.ErrorType, event.Attempts,
- event.StartedAt, event.DurationMS, event.Usage.InputTokens, event.Usage.OutputTokens,
+ event.StartedAt, event.DurationMS, event.TTFTMS, event.Usage.InputTokens, event.Usage.OutputTokens,
event.Usage.TotalTokens, event.Usage.CacheCreationInputTokens, event.Usage.CacheReadInputTokens,
actualCost, charged, uncollected, event.UsageReported, meteringStatus(event)); err != nil {
return fmt.Errorf("persist usage event: %w", err)
@@ -582,25 +600,32 @@ func meteringStatus(event domain.UsageEvent) string {
return "missing"
}
-func reservationCost(model domain.Model, body []byte, defaultMaxOutput int64) (int64, error) {
- maxOutput := defaultMaxOutput
- if model.MaxOutputTokens > 0 && (maxOutput == 0 || model.MaxOutputTokens < maxOutput) {
- maxOutput = model.MaxOutputTokens
+func reservationCost(model domain.Model, body []byte, defaultMaxOutput int64, protocols ...domain.Protocol) (int64, error) {
+ protocol := domain.ProtocolOpenAI
+ if len(protocols) > 0 && protocols[0] != "" {
+ protocol = protocols[0]
}
+ maxOutput := int64(0)
var limits struct {
MaxTokens int64 `json:"max_tokens"`
MaxCompletionTokens int64 `json:"max_completion_tokens"`
MaxOutputTokens int64 `json:"max_output_tokens"`
}
- if json.Unmarshal(body, &limits) == nil {
- explicitMax := int64(0)
- for _, value := range []int64{limits.MaxTokens, limits.MaxCompletionTokens, limits.MaxOutputTokens} {
- if value > explicitMax {
- explicitMax = value
- }
+ if protocol != domain.ProtocolOpenAIEmbeddings {
+ maxOutput = defaultMaxOutput
+ if model.MaxOutputTokens > 0 && (maxOutput == 0 || model.MaxOutputTokens < maxOutput) {
+ maxOutput = model.MaxOutputTokens
}
- if explicitMax > 0 {
- maxOutput = explicitMax
+ if json.Unmarshal(body, &limits) == nil {
+ explicitMax := int64(0)
+ for _, value := range []int64{limits.MaxTokens, limits.MaxCompletionTokens, limits.MaxOutputTokens} {
+ if value > explicitMax {
+ explicitMax = value
+ }
+ }
+ if explicitMax > 0 {
+ maxOutput = explicitMax
+ }
}
}
cacheReservePrice := model.CacheReadPriceMicrosPerMillion
@@ -618,21 +643,51 @@ func usageCost(usage domain.Usage, inputPrice, outputPrice, cacheReadPrice, cach
}
func calculateCost(input, output, cacheRead, cacheWrite, inputPrice, outputPrice, cacheReadPrice, cacheWritePrice int64) (int64, error) {
- values := []int64{input, output, cacheRead, cacheWrite, inputPrice, outputPrice, cacheReadPrice, cacheWritePrice}
- for _, value := range values {
- if value < 0 {
- return 0, errors.New("billing values cannot be negative")
+ return calculateMeteredCost([]meteredCharge{
+ {Unit: domain.MeteringUnitToken, Quantity: input, PriceMicros: inputPrice, PerQuantity: microsPerUnit},
+ {Unit: domain.MeteringUnitToken, Quantity: output, PriceMicros: outputPrice, PerQuantity: microsPerUnit},
+ {Unit: domain.MeteringUnitToken, Quantity: cacheRead, PriceMicros: cacheReadPrice, PerQuantity: microsPerUnit},
+ {Unit: domain.MeteringUnitToken, Quantity: cacheWrite, PriceMicros: cacheWritePrice, PerQuantity: microsPerUnit},
+ })
+}
+
+type meteredCharge struct {
+ Unit domain.MeteringUnit
+ Quantity int64
+ PriceMicros int64
+ PerQuantity int64
+}
+
+// calculateMeteredCost is the common fixed-point primitive for token, image,
+// and duration pricing. Token rates use PerQuantity=1_000_000; image and second
+// rates can use PerQuantity=1 without changing wallet or ledger arithmetic.
+func calculateMeteredCost(charges []meteredCharge) (int64, error) {
+ byScale := make(map[int64]*big.Int)
+ for _, charge := range charges {
+ if charge.Unit != domain.MeteringUnitToken && charge.Unit != domain.MeteringUnitImage && charge.Unit != domain.MeteringUnitSecond {
+ return 0, fmt.Errorf("unsupported metering unit %q", charge.Unit)
+ }
+ if charge.Quantity < 0 || charge.PriceMicros < 0 || charge.PerQuantity <= 0 {
+ return 0, errors.New("metering quantity, price, or scale is invalid")
+ }
+ if charge.Quantity == 0 || charge.PriceMicros == 0 {
+ continue
+ }
+ component := new(big.Int).Mul(big.NewInt(charge.Quantity), big.NewInt(charge.PriceMicros))
+ if byScale[charge.PerQuantity] == nil {
+ byScale[charge.PerQuantity] = new(big.Int)
}
+ byScale[charge.PerQuantity].Add(byScale[charge.PerQuantity], component)
}
total := new(big.Int)
- for _, pair := range [][2]int64{{input, inputPrice}, {output, outputPrice}, {cacheRead, cacheReadPrice}, {cacheWrite, cacheWritePrice}} {
- total.Add(total, new(big.Int).Mul(big.NewInt(pair[0]), big.NewInt(pair[1])))
+ for scale, numerator := range byScale {
+ numerator.Add(numerator, big.NewInt(scale-1))
+ numerator.Div(numerator, big.NewInt(scale))
+ total.Add(total, numerator)
}
if total.Sign() == 0 {
return 0, nil
}
- total.Add(total, big.NewInt(microsPerUnit-1))
- total.Div(total, big.NewInt(microsPerUnit))
if !total.IsInt64() {
return 0, errors.New("calculated charge exceeds supported range")
}
diff --git a/internal/billing/service_test.go b/internal/billing/service_test.go
index 5e21b20..a75e4e0 100644
--- a/internal/billing/service_test.go
+++ b/internal/billing/service_test.go
@@ -84,6 +84,16 @@ func TestAuthorizeEnforcesAPIKeyMonthlySpendCapPostgres(t *testing.T) {
if err := service.Authorize(ctx, Authorization{RequestID: "req_key_budget_allowed", Principal: principal, Model: model, Body: []byte(`{"max_tokens":10}`)}); err != nil {
t.Fatalf("Authorize at exact cap: %v", err)
}
+ principal.MonthlySpendMicros = 0
+ principal.DailySpendMicros = 19
+ err = service.Authorize(ctx, Authorization{RequestID: "req_key_budget_daily_rejected", Principal: principal, Model: model, Body: []byte(`{"max_tokens":10}`)})
+ if !errors.Is(err, ErrDailyQuotaExceeded) {
+ t.Fatalf("Authorize daily error = %v, want ErrDailyQuotaExceeded", err)
+ }
+ principal.DailySpendMicros = 20
+ if err := service.Authorize(ctx, Authorization{RequestID: "req_key_budget_daily_allowed", Principal: principal, Model: model, Body: []byte(`{"max_tokens":10}`)}); err != nil {
+ t.Fatalf("Authorize at exact daily cap including pending reservation: %v", err)
+ }
}
func TestUsageCostUsesFixedPointAndRoundsOnce(t *testing.T) {
@@ -111,6 +121,35 @@ func TestReservationUsesExplicitOutputLimit(t *testing.T) {
}
}
+func TestEmbeddingReservationDoesNotReserveOutputTokens(t *testing.T) {
+ model := domain.Model{InputPriceMicrosPerMillion: 1_000_000, OutputPriceMicrosPerMillion: 50_000_000, MaxOutputTokens: 8192}
+ body := []byte(`{"model":"embedding","input":"hello","max_tokens":99999}`)
+ cost, err := reservationCost(model, body, 4096, domain.ProtocolOpenAIEmbeddings)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if cost != int64(len(body)) {
+ t.Fatalf("embedding reservation = %d, want conservative input-only %d", cost, len(body))
+ }
+}
+
+func TestMeteredCostSupportsTokenImageAndSecondUnits(t *testing.T) {
+ cost, err := calculateMeteredCost([]meteredCharge{
+ {Unit: domain.MeteringUnitToken, Quantity: 500_000, PriceMicros: 2_000_000, PerQuantity: 1_000_000},
+ {Unit: domain.MeteringUnitImage, Quantity: 2, PriceMicros: 40_000, PerQuantity: 1},
+ {Unit: domain.MeteringUnitSecond, Quantity: 3, PriceMicros: 500, PerQuantity: 1},
+ })
+ if err != nil {
+ t.Fatal(err)
+ }
+ if cost != 1_081_500 {
+ t.Fatalf("metered cost = %d, want 1081500", cost)
+ }
+ if _, err := calculateMeteredCost([]meteredCharge{{Unit: "byte", Quantity: 1, PriceMicros: 1, PerQuantity: 1}}); err == nil {
+ t.Fatal("unsupported metering unit was accepted")
+ }
+}
+
func TestMinorToMicrosSupportsCurrencyExponents(t *testing.T) {
tests := []struct {
currency string
@@ -153,6 +192,57 @@ func TestIntegrationIdentifierSuffixUsesLetters(t *testing.T) {
}
}
+func TestStripeSDKVersionAndCheckoutContract(t *testing.T) {
+ if stripe.APIVersion != "2026-07-29.dahlia" {
+ t.Fatalf("Stripe API version = %q; review the integration before changing the pinned version", stripe.APIVersion)
+ }
+ service := &Service{
+ currency: "usd",
+ stripeSuccessURL: "https://console.example.test/billing?topup=success",
+ stripeCancelURL: "https://console.example.test/billing?topup=cancel",
+ integrationIdentifier: "aigw_balance_abcdefgh",
+ }
+ params := service.checkoutSessionParams("order-123", CheckoutInput{TenantID: "tenant-123", AmountMinor: 2500})
+ if params.Mode == nil || *params.Mode != string(stripe.CheckoutSessionModePayment) {
+ t.Fatalf("mode = %v", params.Mode)
+ }
+ if params.IntegrationIdentifier == nil || *params.IntegrationIdentifier != "aigw_balance_abcdefgh" {
+ t.Fatalf("integration identifier = %v", params.IntegrationIdentifier)
+ }
+ if len(params.PaymentMethodTypes) != 0 || len(params.ExcludedPaymentMethodTypes) != 0 {
+ t.Fatal("Checkout must use Dashboard-managed dynamic payment methods")
+ }
+ if params.AutomaticTax != nil || params.TaxIDCollection != nil {
+ t.Fatal("Stripe Tax must remain disabled unless registration is explicitly confirmed")
+ }
+ if params.InvoiceCreation == nil || params.InvoiceCreation.Enabled == nil || !*params.InvoiceCreation.Enabled {
+ t.Fatal("one-time prepaid top-up invoice creation is not enabled")
+ }
+ if len(params.LineItems) != 1 || params.LineItems[0].PriceData == nil || params.LineItems[0].PriceData.UnitAmount == nil || *params.LineItems[0].PriceData.UnitAmount != 2500 {
+ t.Fatalf("unexpected line item %+v", params.LineItems)
+ }
+ if params.IdempotencyKey == nil || *params.IdempotencyKey != "aigw_topup_order-123" {
+ t.Fatalf("idempotency key = %v", params.IdempotencyKey)
+ }
+}
+
+func TestCheckoutContractEnablesTaxOnlyWhenExplicitlyConfigured(t *testing.T) {
+ service := &Service{
+ currency: "usd",
+ stripeAutomaticTax: true,
+ stripeProductTaxCode: "txcd_10103000",
+ integrationIdentifier: "aigw_balance_abcdefgh",
+ }
+ params := service.checkoutSessionParams("order-tax", CheckoutInput{TenantID: "tenant-tax", AmountMinor: 1000})
+ if params.AutomaticTax == nil || params.AutomaticTax.Enabled == nil || !*params.AutomaticTax.Enabled ||
+ params.TaxIDCollection == nil || params.TaxIDCollection.Enabled == nil || !*params.TaxIDCollection.Enabled {
+ t.Fatal("explicit Stripe Tax configuration was not applied")
+ }
+ if params.LineItems[0].PriceData.ProductData.TaxCode == nil || *params.LineItems[0].PriceData.ProductData.TaxCode != "txcd_10103000" {
+ t.Fatal("canonical Stripe product tax code was not applied")
+ }
+}
+
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") ||
@@ -196,6 +286,62 @@ func TestWebhookRejectsInvalidSignatureBeforeProcessing(t *testing.T) {
}
}
+func TestIncompleteStripeCheckoutMarksOrderFailedPostgres(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",
+ StripeSuccessURL: "https://console.example.test/billing?topup=success",
+ StripeCancelURL: "https://console.example.test/billing?topup=cancel",
+ })
+ if err != nil {
+ t.Fatal(err)
+ }
+ t.Cleanup(service.Close)
+
+ suffix := time.Now().UnixNano()
+ var tenantID string
+ if err := service.db.QueryRow(ctx, `INSERT INTO tenants (slug,name) VALUES ($1,'Incomplete Stripe Checkout') RETURNING id::text`, fmt.Sprintf("incomplete-checkout-%d", suffix)).Scan(&tenantID); err != nil {
+ t.Fatal(err)
+ }
+ t.Cleanup(func() {
+ if _, cleanupErr := service.db.Exec(context.Background(), `DELETE FROM topup_orders WHERE tenant_id=$1`, tenantID); cleanupErr != nil {
+ t.Errorf("cleanup incomplete Checkout orders: %v", cleanupErr)
+ }
+ if _, cleanupErr := service.db.Exec(context.Background(), `DELETE FROM tenants WHERE id=$1`, tenantID); cleanupErr != nil {
+ t.Errorf("cleanup incomplete Checkout tenant: %v", cleanupErr)
+ }
+ })
+ if _, err := service.db.Exec(ctx, `INSERT INTO stripe_customers (tenant_id,stripe_customer_id,email) VALUES ($1,$2,'billing@example.test')`, tenantID, fmt.Sprintf("cus_incomplete_%d", suffix)); err != nil {
+ t.Fatal(err)
+ }
+ service.createStripeCheckout = func(context.Context, *stripe.CheckoutSessionCreateParams) (*stripe.CheckoutSession, error) {
+ return nil, nil
+ }
+
+ _, err = service.CreateCheckout(ctx, CheckoutInput{TenantID: tenantID, AmountMinor: 500})
+ if err == nil || !strings.Contains(err.Error(), "incomplete Checkout Session") {
+ t.Fatalf("CreateCheckout error = %v", err)
+ }
+ var status, reconciliationStatus, reconciliationError string
+ if err := service.db.QueryRow(ctx, `SELECT status,reconciliation_status,reconciliation_error
+ FROM topup_orders WHERE tenant_id=$1 ORDER BY created_at DESC LIMIT 1`, tenantID).
+ Scan(&status, &reconciliationStatus, &reconciliationError); err != nil {
+ t.Fatal(err)
+ }
+ if status != "failed" || reconciliationStatus != "unknown" || !strings.Contains(reconciliationError, "incomplete Checkout Session") {
+ t.Fatalf("order state status=%q reconciliation=%q error=%q", status, reconciliationStatus, reconciliationError)
+ }
+}
+
func TestStripeWebhookCreditsPaidOrderExactlyOncePostgres(t *testing.T) {
databaseURL := os.Getenv("AIGW_TEST_DATABASE_URL")
if databaseURL == "" {
@@ -434,4 +580,32 @@ func TestSettlementWorkerPersistsUsageAndReleasesReservationPostgres(t *testing.
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)
}
+ release, err := service.ReleaseUnmeteredReservation(ctx, tenantID, missingRequestID,
+ ReleaseReservationInput{Reason: "provider returned a non-meterable success"}, ResolutionActor{ID: "test-operator", Type: "integration"})
+ if err != nil {
+ t.Fatal(err)
+ }
+ if release.Status != "released" || release.ReservedMicros != 30 {
+ t.Fatalf("unexpected reservation release: %+v", release)
+ }
+ 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)
+ }
+ var releaseLedgerCount int
+ if err := service.db.QueryRow(ctx, `SELECT count(*) FROM billing_ledger WHERE source_type='unmetered_reservation' AND source_id=$1 AND kind='release' AND amount_micros=0`, missingRequestID).Scan(&releaseLedgerCount); err != nil {
+ t.Fatal(err)
+ }
+ if balance != 990 || reserved != 0 || releaseLedgerCount != 1 {
+ t.Fatalf("released wallet balance=%d reserved=%d ledger=%d", balance, reserved, releaseLedgerCount)
+ }
+ 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 meteringStatus != "released_unmetered" {
+ t.Fatalf("released metering status=%s", meteringStatus)
+ }
+ if _, err := service.ReleaseUnmeteredReservation(ctx, tenantID, missingRequestID,
+ ReleaseReservationInput{Reason: "idempotent retry"}, ResolutionActor{ID: "test-operator", Type: "integration"}); err != nil {
+ t.Fatalf("idempotent release retry: %v", err)
+ }
}
diff --git a/internal/billing/stripe.go b/internal/billing/stripe.go
index 032eb02..ea4b7ea 100644
--- a/internal/billing/stripe.go
+++ b/internal/billing/stripe.go
@@ -29,6 +29,35 @@ func (s *Service) CreateCheckout(ctx context.Context, input CheckoutInput) (Chec
if err != nil {
return CheckoutResult{}, err
}
+ params := s.checkoutSessionParams(orderID, input)
+ customerID, err := s.ensureStripeCustomer(ctx, input.TenantID)
+ if err != nil {
+ return CheckoutResult{}, s.failCheckoutCreation(ctx, orderID, err)
+ }
+ if customerID != "" {
+ params.Customer = stripe.String(customerID)
+ } else {
+ params.CustomerCreation = stripe.String(string(stripe.CheckoutSessionCustomerCreationAlways))
+ if strings.TrimSpace(input.CustomerEmail) != "" {
+ params.CustomerEmail = stripe.String(strings.TrimSpace(input.CustomerEmail))
+ }
+ }
+ session, err := s.createStripeCheckout(ctx, params)
+ if err != nil {
+ return CheckoutResult{}, s.failCheckoutCreation(ctx, orderID, fmt.Errorf("create Stripe Checkout Session: %w", err))
+ }
+ if session == nil || session.ID == "" || session.URL == "" {
+ return CheckoutResult{}, s.failCheckoutCreation(ctx, orderID, errors.New("Stripe returned an incomplete Checkout Session"))
+ }
+ if _, err := s.db.Exec(ctx, `
+ UPDATE topup_orders SET stripe_session_id = $2, checkout_url = $3
+ WHERE id = $1`, orderID, session.ID, session.URL); err != nil {
+ return CheckoutResult{}, s.failCheckoutCreation(ctx, orderID, fmt.Errorf("persist Stripe Checkout Session: %w", err))
+ }
+ return CheckoutResult{OrderID: orderID, SessionID: session.ID, URL: session.URL}, nil
+}
+
+func (s *Service) checkoutSessionParams(orderID string, input CheckoutInput) *stripe.CheckoutSessionCreateParams {
params := &stripe.CheckoutSessionCreateParams{
Mode: stripe.String("payment"),
ClientReferenceID: stripe.String(orderID),
@@ -58,43 +87,23 @@ func (s *Service) CreateCheckout(ctx context.Context, input CheckoutInput) (Chec
},
}},
}
- customerID, err := s.ensureStripeCustomer(ctx, input.TenantID)
- if err != nil {
- if _, updateErr := s.db.Exec(ctx, `UPDATE topup_orders SET status = 'failed' WHERE id = $1 AND status = 'pending'`, orderID); updateErr != nil {
- return CheckoutResult{}, errors.Join(err, fmt.Errorf("mark top-up order failed: %w", updateErr))
- }
- return CheckoutResult{}, err
- }
- if customerID != "" {
- params.Customer = stripe.String(customerID)
- } else {
- params.CustomerCreation = stripe.String(string(stripe.CheckoutSessionCustomerCreationAlways))
- if strings.TrimSpace(input.CustomerEmail) != "" {
- params.CustomerEmail = stripe.String(strings.TrimSpace(input.CustomerEmail))
- }
- }
if s.stripeAutomaticTax {
params.AutomaticTax = &stripe.CheckoutSessionCreateAutomaticTaxParams{Enabled: stripe.Bool(true)}
params.TaxIDCollection = &stripe.CheckoutSessionCreateTaxIDCollectionParams{Enabled: stripe.Bool(true)}
params.LineItems[0].PriceData.ProductData.TaxCode = stripe.String(s.stripeProductTaxCode)
}
params.SetIdempotencyKey("aigw_topup_" + orderID)
- session, err := s.createStripeCheckout(ctx, params)
- if err != nil {
- if _, updateErr := s.db.Exec(ctx, `UPDATE topup_orders SET status = 'failed' WHERE id = $1 AND status = 'pending'`, orderID); updateErr != nil {
- return CheckoutResult{}, errors.Join(fmt.Errorf("create Stripe Checkout Session: %w", err), fmt.Errorf("mark top-up order failed: %w", updateErr))
- }
- return CheckoutResult{}, fmt.Errorf("create Stripe Checkout Session: %w", err)
- }
- if session.ID == "" || session.URL == "" {
- return CheckoutResult{}, errors.New("Stripe returned an incomplete Checkout Session")
- }
- if _, err := s.db.Exec(ctx, `
- UPDATE topup_orders SET stripe_session_id = $2, checkout_url = $3
- WHERE id = $1`, orderID, session.ID, session.URL); err != nil {
- return CheckoutResult{}, fmt.Errorf("persist Stripe Checkout Session: %w", err)
- }
- return CheckoutResult{OrderID: orderID, SessionID: session.ID, URL: session.URL}, nil
+ return params
+}
+
+func (s *Service) failCheckoutCreation(ctx context.Context, orderID string, cause error) error {
+ _, updateErr := s.db.Exec(ctx, `UPDATE topup_orders
+ SET status='failed', reconciliation_status='unknown', reconciliation_error=$2
+ WHERE id=$1 AND status='pending'`, orderID, truncateError(cause))
+ if updateErr != nil {
+ return errors.Join(cause, fmt.Errorf("mark top-up order failed: %w", updateErr))
+ }
+ return cause
}
func checkoutReturnURL(raw, orderID string, includeStripeSession bool) string {
diff --git a/internal/billing/stripe_preflight.go b/internal/billing/stripe_preflight.go
new file mode 100644
index 0000000..4f9250c
--- /dev/null
+++ b/internal/billing/stripe_preflight.go
@@ -0,0 +1,124 @@
+package billing
+
+import (
+ "context"
+ "encoding/json"
+ "errors"
+ "fmt"
+ "io"
+ "net/http"
+ "net/url"
+ "strings"
+
+ "github.com/stripe/stripe-go/v86"
+)
+
+const (
+ stripeAPIBaseURL = "https://api.stripe.com"
+ maxStripeErrorBodySize = 64 << 10
+)
+
+var ErrLiveStripeKey = errors.New("Stripe permission preflight only accepts test-mode keys")
+
+type StripePermissionCheck struct {
+ Name string `json:"name"`
+ OK bool `json:"ok"`
+ StatusCode int `json:"status_code"`
+ ErrorCode string `json:"error_code,omitempty"`
+ ErrorType string `json:"error_type,omitempty"`
+}
+
+type StripePreflightResult struct {
+ APIVersion string `json:"api_version"`
+ TestMode bool `json:"test_mode"`
+ Ready bool `json:"ready"`
+ Checks []StripePermissionCheck `json:"checks"`
+}
+
+type stripeReadCheck struct {
+ name string
+ path string
+}
+
+var stripeRequiredReadChecks = []stripeReadCheck{
+ {name: "customers_read", path: "/v1/customers"},
+ {name: "checkout_sessions_read", path: "/v1/checkout/sessions"},
+ {name: "setup_intents_read", path: "/v1/setup_intents"},
+ {name: "payment_intents_read", path: "/v1/payment_intents"},
+ {name: "refunds_read", path: "/v1/refunds"},
+ {name: "charges_read", path: "/v1/charges"},
+ {name: "disputes_read", path: "/v1/disputes"},
+ {name: "invoices_read", path: "/v1/invoices"},
+ {name: "billing_portal_configurations_read", path: "/v1/billing_portal/configurations"},
+}
+
+// CheckStripePermissions validates the read side of the restricted-key contract
+// without creating Stripe objects. Write permissions are exercised by the
+// sandbox Checkout, Portal, automatic top-up, refund, and reconciliation flows.
+func CheckStripePermissions(ctx context.Context, apiKey string) (StripePreflightResult, error) {
+ return checkStripePermissions(ctx, apiKey, stripeAPIBaseURL, http.DefaultClient)
+}
+
+func checkStripePermissions(ctx context.Context, apiKey, baseURL string, client *http.Client) (StripePreflightResult, error) {
+ apiKey = strings.TrimSpace(apiKey)
+ result := StripePreflightResult{
+ APIVersion: stripe.APIVersion,
+ TestMode: isStripeTestKey(apiKey),
+ Checks: make([]StripePermissionCheck, 0, len(stripeRequiredReadChecks)),
+ }
+ if !result.TestMode {
+ return result, ErrLiveStripeKey
+ }
+ result.Ready = true
+ if client == nil {
+ client = http.DefaultClient
+ }
+ for _, check := range stripeRequiredReadChecks {
+ item := runStripeReadCheck(ctx, client, apiKey, baseURL, check)
+ result.Checks = append(result.Checks, item)
+ result.Ready = result.Ready && item.OK
+ }
+ return result, nil
+}
+
+func runStripeReadCheck(ctx context.Context, client *http.Client, apiKey, baseURL string, check stripeReadCheck) StripePermissionCheck {
+ endpoint, err := url.JoinPath(baseURL, check.path)
+ if err != nil {
+ return StripePermissionCheck{Name: check.name, ErrorType: "configuration_error"}
+ }
+ query := url.Values{"limit": []string{"1"}}
+ request, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint+"?"+query.Encode(), nil)
+ if err != nil {
+ return StripePermissionCheck{Name: check.name, ErrorType: "configuration_error"}
+ }
+ request.Header.Set("Authorization", "Bearer "+apiKey)
+ request.Header.Set("Stripe-Version", stripe.APIVersion)
+ response, err := client.Do(request)
+ if err != nil {
+ return StripePermissionCheck{Name: check.name, ErrorType: "network_error"}
+ }
+ defer response.Body.Close()
+ item := StripePermissionCheck{Name: check.name, OK: response.StatusCode >= 200 && response.StatusCode < 300, StatusCode: response.StatusCode}
+ if item.OK {
+ _, _ = io.Copy(io.Discard, response.Body)
+ return item
+ }
+ var envelope struct {
+ Error struct {
+ Code string `json:"code"`
+ Type string `json:"type"`
+ } `json:"error"`
+ }
+ if err := json.NewDecoder(io.LimitReader(response.Body, maxStripeErrorBodySize)).Decode(&envelope); err == nil {
+ item.ErrorCode = envelope.Error.Code
+ item.ErrorType = envelope.Error.Type
+ }
+ if item.ErrorType == "" {
+ item.ErrorType = fmt.Sprintf("http_%d", response.StatusCode)
+ }
+ return item
+}
+
+func isStripeTestKey(value string) bool {
+ return strings.HasPrefix(value, "rk_test_") || strings.HasPrefix(value, "sk_test_")
+}
diff --git a/internal/billing/stripe_preflight_test.go b/internal/billing/stripe_preflight_test.go
new file mode 100644
index 0000000..e8f1841
--- /dev/null
+++ b/internal/billing/stripe_preflight_test.go
@@ -0,0 +1,90 @@
+package billing
+
+import (
+ "context"
+ "io"
+ "net/http"
+ "strings"
+ "testing"
+
+ "github.com/stripe/stripe-go/v86"
+)
+
+func TestStripePermissionPreflightChecksRequiredResourcesWithoutLeakingKey(t *testing.T) {
+ const key = "rk_test_do_not_log_this_value"
+ seen := make(map[string]bool)
+ client := &http.Client{Transport: stripeRoundTripFunc(func(r *http.Request) (*http.Response, error) {
+ if r.Method != http.MethodGet || r.URL.Query().Get("limit") != "1" {
+ t.Errorf("unexpected request %s %s", r.Method, r.URL.String())
+ }
+ if r.Header.Get("Authorization") != "Bearer "+key {
+ t.Errorf("missing Stripe bearer authentication")
+ }
+ if r.Header.Get("Stripe-Version") != stripe.APIVersion {
+ t.Errorf("Stripe-Version = %q", r.Header.Get("Stripe-Version"))
+ }
+ seen[r.URL.Path] = true
+ return stripeTestResponse(http.StatusOK, `{"object":"list","data":[]}`), nil
+ })}
+
+ result, err := checkStripePermissions(context.Background(), key, "https://stripe.test", client)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if !result.Ready || !result.TestMode || result.APIVersion != stripe.APIVersion {
+ t.Fatalf("unexpected result %+v", result)
+ }
+ if len(result.Checks) != len(stripeRequiredReadChecks) {
+ t.Fatalf("checks = %d, want %d", len(result.Checks), len(stripeRequiredReadChecks))
+ }
+ for _, check := range stripeRequiredReadChecks {
+ if !seen[check.path] {
+ t.Errorf("endpoint %s was not checked", check.path)
+ }
+ }
+}
+
+func TestStripePermissionPreflightReportsSanitizedStripeError(t *testing.T) {
+ client := &http.Client{Transport: stripeRoundTripFunc(func(*http.Request) (*http.Response, error) {
+ return stripeTestResponse(http.StatusForbidden, `{"error":{"type":"invalid_request_error","code":"permission_denied","message":"secret details"}}`), nil
+ })}
+
+ result, err := checkStripePermissions(context.Background(), "rk_test_placeholder", "https://stripe.test", client)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if result.Ready || len(result.Checks) == 0 {
+ t.Fatalf("unexpected result %+v", result)
+ }
+ for _, check := range result.Checks {
+ if check.OK || check.StatusCode != http.StatusForbidden || check.ErrorCode != "permission_denied" || check.ErrorType != "invalid_request_error" {
+ t.Fatalf("unexpected check %+v", check)
+ }
+ }
+}
+
+type stripeRoundTripFunc func(*http.Request) (*http.Response, error)
+
+func (fn stripeRoundTripFunc) RoundTrip(request *http.Request) (*http.Response, error) {
+ return fn(request)
+}
+
+func stripeTestResponse(status int, body string) *http.Response {
+ return &http.Response{
+ StatusCode: status,
+ Header: make(http.Header),
+ Body: io.NopCloser(strings.NewReader(body)),
+ }
+}
+
+func TestStripePermissionPreflightRejectsLiveAndMalformedKeys(t *testing.T) {
+ for _, key := range []string{"", "rk_live_forbidden", "not-a-stripe-key"} {
+ result, err := checkStripePermissions(context.Background(), key, "http://unused", nil)
+ if err != ErrLiveStripeKey {
+ t.Fatalf("key %q: error = %v, want ErrLiveStripeKey", key, err)
+ }
+ if result.Ready || result.Checks == nil || len(result.Checks) != 0 {
+ t.Fatalf("key %q: unexpected rejected result %+v", key, result)
+ }
+ }
+}
diff --git a/internal/billing/types.go b/internal/billing/types.go
index 1633ecc..338eeb2 100644
--- a/internal/billing/types.go
+++ b/internal/billing/types.go
@@ -9,18 +9,20 @@ import (
)
var (
- ErrInsufficientBalance = errors.New("insufficient balance")
- ErrStripeDisabled = errors.New("Stripe top-ups are disabled")
- ErrInvalidAmount = errors.New("invalid amount")
- ErrQuotaExceeded = errors.New("monthly spend quota exceeded")
- ErrTopUpOrderNotFound = errors.New("top-up order not found")
- ErrUsageNotReported = errors.New("billable successful response did not report usage")
- ErrCannotResolveTopUp = errors.New("top-up order cannot be resolved as missing")
- ErrBillingAccountNotFound = errors.New("billing account not found")
- ErrPaymentMethodRequired = errors.New("a saved payment method is required")
- ErrAutoTopUpNeedsAttention = errors.New("automatic top-up payment method requires attention")
- ErrInvalidBillingProfile = errors.New("invalid billing profile")
- ErrBillingProfileSync = errors.New("billing profile Stripe synchronization failed")
+ ErrInsufficientBalance = errors.New("insufficient balance")
+ ErrStripeDisabled = errors.New("Stripe top-ups are disabled")
+ ErrInvalidAmount = errors.New("invalid amount")
+ ErrQuotaExceeded = errors.New("monthly spend quota exceeded")
+ ErrDailyQuotaExceeded = errors.New("daily spend quota exceeded")
+ ErrTopUpOrderNotFound = errors.New("top-up order not found")
+ ErrUsageNotReported = errors.New("billable successful response did not report usage")
+ ErrCannotResolveTopUp = errors.New("top-up order cannot be resolved as missing")
+ ErrBillingAccountNotFound = errors.New("billing account not found")
+ ErrPaymentMethodRequired = errors.New("a saved payment method is required")
+ ErrAutoTopUpNeedsAttention = errors.New("automatic top-up payment method requires attention")
+ ErrInvalidBillingProfile = errors.New("invalid billing profile")
+ ErrBillingProfileSync = errors.New("billing profile Stripe synchronization failed")
+ ErrReservationNotReleasable = errors.New("billing reservation is not releasable")
)
type Meter interface {
@@ -32,6 +34,7 @@ type Authorization struct {
RequestID string
Principal domain.Principal
Model domain.Model
+ Protocol domain.Protocol
Body []byte
Policy domain.LimitPolicy
}
@@ -127,6 +130,18 @@ type AdjustmentInput struct {
Description string `json:"description"`
}
+type ReleaseReservationInput struct {
+ Reason string `json:"reason"`
+}
+
+type ReservationRelease struct {
+ RequestID string `json:"request_id"`
+ TenantID string `json:"tenant_id"`
+ ReservedMicros int64 `json:"reserved_micros"`
+ Status string `json:"status"`
+ ReleasedAt time.Time `json:"released_at"`
+}
+
type CheckoutInput struct {
TenantID string `json:"tenant_id"`
AmountMinor int64 `json:"amount_minor"`