diff options
Diffstat (limited to '')
| -rw-r--r-- | internal/billing/auto_topup.go | 56 | ||||
| -rw-r--r-- | internal/billing/auto_topup_test.go | 43 | ||||
| -rw-r--r-- | internal/billing/ledger.go | 77 | ||||
| -rw-r--r-- | internal/billing/operations.go | 8 | ||||
| -rw-r--r-- | internal/billing/operations_test.go | 58 | ||||
| -rw-r--r-- | internal/billing/service.go | 115 | ||||
| -rw-r--r-- | internal/billing/service_test.go | 174 | ||||
| -rw-r--r-- | internal/billing/stripe.go | 71 | ||||
| -rw-r--r-- | internal/billing/stripe_preflight.go | 124 | ||||
| -rw-r--r-- | internal/billing/stripe_preflight_test.go | 90 | ||||
| -rw-r--r-- | internal/billing/types.go | 39 |
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, ¤cy, &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"` |
