summaryrefslogtreecommitdiff
path: root/internal/billing
diff options
context:
space:
mode:
Diffstat (limited to '')
-rw-r--r--internal/billing/auto_topup.go541
-rw-r--r--internal/billing/auto_topup_test.go155
-rw-r--r--internal/billing/ledger.go8
-rw-r--r--internal/billing/operations.go84
-rw-r--r--internal/billing/profile.go197
-rw-r--r--internal/billing/profile_test.go135
-rw-r--r--internal/billing/service.go66
-rw-r--r--internal/billing/service_test.go67
-rw-r--r--internal/billing/stripe.go25
-rw-r--r--internal/billing/types.go88
10 files changed, 1330 insertions, 36 deletions
diff --git a/internal/billing/auto_topup.go b/internal/billing/auto_topup.go
new file mode 100644
index 0000000..a90405b
--- /dev/null
+++ b/internal/billing/auto_topup.go
@@ -0,0 +1,541 @@
+package billing
+
+import (
+ "context"
+ "errors"
+ "fmt"
+ "net/url"
+ "strings"
+ "time"
+
+ "github.com/jackc/pgx/v5"
+ "github.com/stripe/stripe-go/v86"
+)
+
+type stripeSetupIntentRetriever func(context.Context, string, *stripe.SetupIntentRetrieveParams) (*stripe.SetupIntent, error)
+type stripePaymentIntentCreator func(context.Context, *stripe.PaymentIntentCreateParams) (*stripe.PaymentIntent, error)
+type stripePaymentIntentRetriever func(context.Context, string, *stripe.PaymentIntentRetrieveParams) (*stripe.PaymentIntent, error)
+
+const autoTopUpAction = "auto_topup_setup"
+
+func (s *Service) defaultAutoTopUpAmountMinor() int64 {
+ amount := int64(2000)
+ if amount < s.minTopUpMinor {
+ amount = s.minTopUpMinor
+ }
+ if amount > s.maxTopUpMinor {
+ amount = s.maxTopUpMinor
+ }
+ return amount
+}
+
+func (s *Service) defaultAutoTopUpThresholdMicros() int64 {
+ amount, err := minorToMicros(s.currency, s.defaultAutoTopUpAmountMinor())
+ if err != nil || amount <= 0 {
+ return 0
+ }
+ if amount/4 > 5*microsPerUnit {
+ return 5 * microsPerUnit
+ }
+ return amount / 4
+}
+
+func (s *Service) GetAutoTopUpSettings(ctx context.Context, tenantID string) (AutoTopUpSettings, error) {
+ tenantID = strings.TrimSpace(tenantID)
+ if tenantID == "" {
+ return AutoTopUpSettings{}, ErrBillingAccountNotFound
+ }
+ var result AutoTopUpSettings
+ var paymentMethodID string
+ err := s.db.QueryRow(ctx, `
+ SELECT t.id::text, COALESCE(w.currency,$2), $3::boolean,
+ COALESCE(a.enabled,FALSE), COALESCE(a.threshold_micros,$4),
+ COALESCE(a.topup_amount_minor,$5), COALESCE(a.stripe_payment_method_id,''),
+ COALESCE(a.payment_method_type,''), COALESCE(a.payment_method_brand,''),
+ COALESCE(a.payment_method_last4,''), COALESCE(a.payment_method_exp_month,0),
+ COALESCE(a.payment_method_exp_year,0), COALESCE(a.status,'not_configured'),
+ COALESCE(a.last_error,''), a.last_attempt_at, a.last_succeeded_at,
+ a.next_attempt_at, COALESCE(a.updated_at,t.created_at)
+ FROM tenants t
+ LEFT JOIN tenant_wallets w ON w.tenant_id=t.id
+ LEFT JOIN tenant_auto_topup_settings a ON a.tenant_id=t.id
+ WHERE t.id=$1`, tenantID, s.currency, s.stripeEnabled, s.defaultAutoTopUpThresholdMicros(), s.defaultAutoTopUpAmountMinor()).
+ Scan(&result.TenantID, &result.Currency, &result.StripeEnabled, &result.Enabled,
+ &result.ThresholdMicros, &result.TopUpAmountMinor, &paymentMethodID,
+ &result.PaymentMethodType, &result.PaymentMethodBrand, &result.PaymentMethodLast4,
+ &result.PaymentMethodExpMonth, &result.PaymentMethodExpYear, &result.Status,
+ &result.LastError, &result.LastAttemptAt, &result.LastSucceededAt,
+ &result.NextAttemptAt, &result.UpdatedAt)
+ if errors.Is(err, pgx.ErrNoRows) {
+ return AutoTopUpSettings{}, ErrBillingAccountNotFound
+ }
+ if err != nil {
+ return AutoTopUpSettings{}, fmt.Errorf("query automatic top-up settings: %w", err)
+ }
+ result.PaymentMethodConfigured = paymentMethodID != ""
+ return result, nil
+}
+
+func (s *Service) UpdateAutoTopUp(ctx context.Context, input UpdateAutoTopUpInput) (AutoTopUpSettings, error) {
+ input.TenantID = strings.TrimSpace(input.TenantID)
+ if input.TenantID == "" || input.ThresholdMicros < 0 || input.TopUpAmountMinor < s.minTopUpMinor || input.TopUpAmountMinor > s.maxTopUpMinor {
+ return AutoTopUpSettings{}, ErrInvalidAmount
+ }
+ topUpMicros, err := minorToMicros(s.currency, input.TopUpAmountMinor)
+ if err != nil || topUpMicros <= input.ThresholdMicros {
+ return AutoTopUpSettings{}, fmt.Errorf("%w: automatic top-up amount must exceed the balance threshold", ErrInvalidAmount)
+ }
+ if input.Enabled && !s.stripeEnabled {
+ return AutoTopUpSettings{}, ErrStripeDisabled
+ }
+ tx, err := s.db.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.ReadCommitted})
+ if err != nil {
+ return AutoTopUpSettings{}, err
+ }
+ defer tx.Rollback(ctx)
+ if _, err := tx.Exec(ctx, `INSERT INTO tenant_auto_topup_settings (tenant_id,threshold_micros,topup_amount_minor)
+ VALUES ($1,$2,$3) ON CONFLICT (tenant_id) DO NOTHING`, input.TenantID, input.ThresholdMicros, input.TopUpAmountMinor); err != nil {
+ return AutoTopUpSettings{}, fmt.Errorf("initialize automatic top-up settings: %w", err)
+ }
+ var paymentMethodID, status string
+ if err := tx.QueryRow(ctx, `SELECT COALESCE(stripe_payment_method_id,''),status FROM tenant_auto_topup_settings WHERE tenant_id=$1 FOR UPDATE`, input.TenantID).Scan(&paymentMethodID, &status); err != nil {
+ if errors.Is(err, pgx.ErrNoRows) {
+ return AutoTopUpSettings{}, ErrBillingAccountNotFound
+ }
+ return AutoTopUpSettings{}, err
+ }
+ if input.Enabled && paymentMethodID == "" {
+ return AutoTopUpSettings{}, ErrPaymentMethodRequired
+ }
+ if input.Enabled && status == "action_required" {
+ return AutoTopUpSettings{}, ErrAutoTopUpNeedsAttention
+ }
+ next := any(nil)
+ newStatus := "not_configured"
+ if paymentMethodID != "" {
+ newStatus = "ready"
+ }
+ if input.Enabled {
+ newStatus = "ready"
+ next = time.Now().UTC()
+ }
+ if _, err := tx.Exec(ctx, `UPDATE tenant_auto_topup_settings SET enabled=$2,threshold_micros=$3,topup_amount_minor=$4,
+ status=$5,last_error=CASE WHEN $2 THEN '' ELSE last_error END,
+ next_attempt_at=$6,updated_at=now() WHERE tenant_id=$1`, input.TenantID, input.Enabled,
+ input.ThresholdMicros, input.TopUpAmountMinor, newStatus, next); err != nil {
+ return AutoTopUpSettings{}, fmt.Errorf("update automatic top-up settings: %w", err)
+ }
+ if err := tx.Commit(ctx); err != nil {
+ return AutoTopUpSettings{}, err
+ }
+ return s.GetAutoTopUpSettings(ctx, input.TenantID)
+}
+
+func (s *Service) DisableAutoTopUp(ctx context.Context, tenantID string) (AutoTopUpSettings, error) {
+ tenantID = strings.TrimSpace(tenantID)
+ if tenantID == "" {
+ return AutoTopUpSettings{}, ErrBillingAccountNotFound
+ }
+ if _, err := s.db.Exec(ctx, `UPDATE tenant_auto_topup_settings SET enabled=FALSE,status=CASE WHEN stripe_payment_method_id IS NULL THEN 'not_configured' ELSE 'ready' END,next_attempt_at=NULL,updated_at=now() WHERE tenant_id=$1`, tenantID); err != nil {
+ return AutoTopUpSettings{}, err
+ }
+ return s.GetAutoTopUpSettings(ctx, tenantID)
+}
+
+// CreateAutoTopUpSetupSession opens a Stripe-hosted SetupIntent flow. Stripe
+// owns card collection; this service only receives a PaymentMethod ID after a
+// signed webhook confirms that the setup succeeded.
+func (s *Service) CreateAutoTopUpSetupSession(ctx context.Context, input AutoTopUpSetupInput) (AutoTopUpSetupResult, error) {
+ if !s.stripeEnabled || s.createStripeCheckout == nil {
+ return AutoTopUpSetupResult{}, ErrStripeDisabled
+ }
+ input.TenantID = strings.TrimSpace(input.TenantID)
+ if input.TenantID == "" {
+ return AutoTopUpSetupResult{}, ErrBillingAccountNotFound
+ }
+ if _, err := s.GetAutoTopUpSettings(ctx, input.TenantID); err != nil {
+ return AutoTopUpSetupResult{}, err
+ }
+ customerID, err := s.ensureStripeCustomer(ctx, input.TenantID)
+ if err != nil {
+ return AutoTopUpSetupResult{}, err
+ }
+ params := &stripe.CheckoutSessionCreateParams{
+ Mode: stripe.String(string(stripe.CheckoutSessionModeSetup)),
+ Currency: stripe.String(s.currency),
+ ClientReferenceID: stripe.String(input.TenantID),
+ IntegrationIdentifier: stripe.String(s.integrationIdentifier),
+ SuccessURL: stripe.String(autoTopUpReturnURL(s.stripeSuccessURL, true)),
+ CancelURL: stripe.String(autoTopUpReturnURL(s.stripeCancelURL, false)),
+ Metadata: map[string]string{
+ "aigw_action": autoTopUpAction,
+ "aigw_tenant_id": input.TenantID,
+ },
+ }
+ 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))
+ }
+ }
+ 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
+}
+
+func autoTopUpReturnURL(raw string, success bool) string {
+ parsed, err := url.Parse(raw)
+ if err != nil {
+ return raw
+ }
+ query := parsed.Query()
+ query.Set("autotopup", "setup")
+ if success {
+ query.Set("session_id", "{CHECKOUT_SESSION_ID}")
+ } else {
+ query.Set("autotopup", "cancel")
+ query.Del("session_id")
+ }
+ parsed.RawQuery = strings.ReplaceAll(query.Encode(), url.QueryEscape("{CHECKOUT_SESSION_ID}"), "{CHECKOUT_SESSION_ID}")
+ return parsed.String()
+}
+
+func (s *Service) processAutoTopUpSetupEvent(ctx context.Context, event stripe.Event, session *stripe.CheckoutSession) error {
+ if s.retrieveStripeSetupIntent == nil || session == nil || event.ID == "" || session.ID == "" || session.ClientReferenceID == "" {
+ return ErrInvalidAmount
+ }
+ if event.Type != stripe.EventTypeCheckoutSessionCompleted {
+ return nil
+ }
+ setupIntentID := ""
+ if session.SetupIntent != nil {
+ setupIntentID = session.SetupIntent.ID
+ }
+ if setupIntentID == "" {
+ return ErrPaymentMethodRequired
+ }
+ intent, err := s.retrieveStripeSetupIntent(ctx, setupIntentID, &stripe.SetupIntentRetrieveParams{})
+ if err != nil {
+ return fmt.Errorf("retrieve automatic top-up setup intent: %w", err)
+ }
+ if intent == nil || intent.Status != stripe.SetupIntentStatusSucceeded || intent.PaymentMethod == nil || intent.PaymentMethod.ID == "" {
+ return ErrPaymentMethodRequired
+ }
+ tx, err := s.db.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.ReadCommitted})
+ if err != nil {
+ return err
+ }
+ defer tx.Rollback(ctx)
+ if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended($1,0))`, event.ID); err != nil {
+ return err
+ }
+ if _, err := tx.Exec(ctx, `INSERT INTO stripe_webhook_events (event_id,event_type) VALUES ($1,$2) ON CONFLICT DO NOTHING`, event.ID, string(event.Type)); err != nil {
+ return err
+ }
+ var storedSetupSession string
+ if err := tx.QueryRow(ctx, `SELECT COALESCE(stripe_setup_session_id,'') FROM tenant_auto_topup_settings WHERE tenant_id=$1 FOR UPDATE`, session.ClientReferenceID).Scan(&storedSetupSession); err != nil || storedSetupSession != session.ID {
+ return ErrInvalidAmount
+ }
+ customerID := ""
+ if session.Customer != nil {
+ customerID = session.Customer.ID
+ }
+ if customerID == "" && intent.Customer != nil {
+ customerID = intent.Customer.ID
+ }
+ if customerID == "" {
+ return ErrPaymentMethodRequired
+ }
+ email := session.CustomerEmail
+ if session.CustomerDetails != nil && session.CustomerDetails.Email != "" {
+ email = session.CustomerDetails.Email
+ }
+ if _, err := tx.Exec(ctx, `INSERT INTO stripe_customers (tenant_id,stripe_customer_id,email) VALUES ($1,$2,$3)
+ ON CONFLICT (tenant_id) DO UPDATE SET stripe_customer_id=EXCLUDED.stripe_customer_id,
+ email=CASE WHEN EXCLUDED.email='' THEN stripe_customers.email ELSE EXCLUDED.email END,updated_at=now()`, session.ClientReferenceID, customerID, email); err != nil {
+ return err
+ }
+ methodType, brand, last4 := intent.PaymentMethod.Type, "", ""
+ var expMonth, expYear int64
+ if intent.PaymentMethod.Card != nil {
+ brand, last4 = string(intent.PaymentMethod.Card.Brand), intent.PaymentMethod.Card.Last4
+ expMonth, expYear = intent.PaymentMethod.Card.ExpMonth, intent.PaymentMethod.Card.ExpYear
+ }
+ if _, err := tx.Exec(ctx, `INSERT INTO tenant_auto_topup_settings
+ (tenant_id,threshold_micros,topup_amount_minor,stripe_payment_method_id,payment_method_type,payment_method_brand,payment_method_last4,payment_method_exp_month,payment_method_exp_year,stripe_setup_session_id,status,last_error,failure_count,updated_at)
+ VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,'ready','',0,now())
+ ON CONFLICT (tenant_id) DO UPDATE SET stripe_payment_method_id=EXCLUDED.stripe_payment_method_id,
+ payment_method_type=EXCLUDED.payment_method_type,payment_method_brand=EXCLUDED.payment_method_brand,
+ payment_method_last4=EXCLUDED.payment_method_last4,payment_method_exp_month=EXCLUDED.payment_method_exp_month,
+ payment_method_exp_year=EXCLUDED.payment_method_exp_year,stripe_setup_session_id=EXCLUDED.stripe_setup_session_id,
+ status='ready',last_error='',failure_count=0,next_attempt_at=NULL,updated_at=now()`,
+ session.ClientReferenceID, s.defaultAutoTopUpThresholdMicros(), s.defaultAutoTopUpAmountMinor(), intent.PaymentMethod.ID,
+ methodType, brand, last4, expMonth, expYear, session.ID); err != nil {
+ return err
+ }
+ if _, err := tx.Exec(ctx, `UPDATE stripe_webhook_events SET processed_at=now(),processing_error='' WHERE event_id=$1`, event.ID); err != nil {
+ return err
+ }
+ return tx.Commit(ctx)
+}
+
+// processAutoTopUpOnce claims one eligible tenant before making a Stripe call.
+// The row lock and unique pending-order index make this safe across gateways.
+func (s *Service) processAutoTopUpOnce(ctx context.Context) (bool, error) {
+ if !s.stripeEnabled || s.createStripePaymentIntent == nil {
+ return false, nil
+ }
+ tx, err := s.db.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.ReadCommitted})
+ if err != nil {
+ return false, err
+ }
+ defer tx.Rollback(ctx)
+ var tenantID, customerID, customerEmail, paymentMethodID, currency, orderID string
+ var amountMinor, threshold, balance, reserved int64
+ row := tx.QueryRow(ctx, `
+ SELECT a.tenant_id::text,c.stripe_customer_id,c.email,a.stripe_payment_method_id,w.currency,
+ a.topup_amount_minor,a.threshold_micros,w.balance_micros,w.reserved_micros
+ ,COALESCE(o.id::text,'')
+ FROM tenant_auto_topup_settings a
+ JOIN stripe_customers c ON c.tenant_id=a.tenant_id
+ JOIN tenant_wallets w ON w.tenant_id=a.tenant_id
+ LEFT JOIN LATERAL (SELECT id FROM topup_orders WHERE tenant_id=a.tenant_id AND trigger_type='auto' AND status='pending' ORDER BY created_at DESC LIMIT 1) o ON TRUE
+ WHERE a.enabled AND a.stripe_payment_method_id IS NOT NULL
+ AND a.status IN ('ready','failed','charging')
+ AND (a.next_attempt_at IS NULL OR a.next_attempt_at <= now())
+ AND w.balance_micros-w.reserved_micros <= a.threshold_micros
+ ORDER BY w.balance_micros-w.reserved_micros,a.updated_at
+ FOR UPDATE OF a SKIP LOCKED LIMIT 1`)
+ if err := row.Scan(&tenantID, &customerID, &customerEmail, &paymentMethodID, &currency, &amountMinor, &threshold, &balance, &reserved, &orderID); err != nil {
+ if errors.Is(err, pgx.ErrNoRows) {
+ return false, tx.Commit(ctx)
+ }
+ return false, err
+ }
+ if balance-reserved > threshold {
+ return false, tx.Commit(ctx)
+ }
+ amountMicros, err := minorToMicros(currency, amountMinor)
+ if err != nil {
+ return false, err
+ }
+ if orderID == "" {
+ if err := tx.QueryRow(ctx, `INSERT INTO topup_orders (tenant_id,amount_minor,amount_micros,currency,trigger_type)
+ VALUES ($1,$2,$3,$4,'auto') RETURNING id::text`, tenantID, amountMinor, amountMicros, currency).Scan(&orderID); err != nil {
+ return false, err
+ }
+ }
+ if _, err := tx.Exec(ctx, `UPDATE tenant_auto_topup_settings SET status='charging',last_attempt_at=now(),next_attempt_at=now()+interval '15 minutes',updated_at=now() WHERE tenant_id=$1`, tenantID); err != nil {
+ return false, err
+ }
+ 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)
+ intent, callErr := s.createStripePaymentIntent(ctx, params)
+ if callErr != nil {
+ var stripeErr *stripe.Error
+ if errors.As(callErr, &stripeErr) && stripeErr.PaymentIntent != nil {
+ intent = stripeErr.PaymentIntent
+ if intent.ID != "" {
+ _, _ = s.db.Exec(ctx, `UPDATE topup_orders SET stripe_payment_intent_id=$2,stripe_customer_id=$3 WHERE id=$1`, orderID, intent.ID, customerID)
+ }
+ if intent.Status == stripe.PaymentIntentStatusSucceeded {
+ return true, s.creditAutoTopUpPaymentIntent(ctx, intent)
+ }
+ return true, s.applyAutoTopUpPaymentIntentFailure(ctx, intent)
+ }
+ return true, s.scheduleAutoTopUpRetry(ctx, tenantID, orderID, callErr)
+ }
+ if intent == nil || intent.ID == "" {
+ return true, s.failAutoTopUp(ctx, tenantID, orderID, errors.New("Stripe returned an incomplete automatic top-up PaymentIntent"))
+ }
+ if _, err := s.db.Exec(ctx, `UPDATE topup_orders SET stripe_payment_intent_id=$2,stripe_customer_id=$3 WHERE id=$1`, orderID, intent.ID, customerID); err != nil {
+ return true, err
+ }
+ if intent.Status == stripe.PaymentIntentStatusSucceeded {
+ return true, s.creditAutoTopUpPaymentIntent(ctx, intent)
+ }
+ if intent.Status == stripe.PaymentIntentStatusProcessing {
+ return true, nil
+ }
+ if intent.Status == stripe.PaymentIntentStatusRequiresAction || intent.Status == stripe.PaymentIntentStatusRequiresPaymentMethod || intent.Status == stripe.PaymentIntentStatusCanceled {
+ return true, s.markAutoTopUpAttention(ctx, tenantID, orderID, ErrAutoTopUpNeedsAttention)
+ }
+ return true, s.failAutoTopUp(ctx, tenantID, orderID, fmt.Errorf("automatic top-up PaymentIntent ended in status %s", intent.Status))
+}
+
+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)
+ if err != nil {
+ return errors.Join(cause, err)
+ }
+ _, settingsErr := s.db.Exec(ctx, `UPDATE tenant_auto_topup_settings SET status='failed',failure_count=failure_count+1,last_error=$2,
+ next_attempt_at=now()+interval '1 hour',updated_at=now() WHERE tenant_id=$1`, tenantID, message)
+ if settingsErr != nil {
+ return errors.Join(cause, settingsErr)
+ }
+ return cause
+}
+
+func (s *Service) failAutoTopUp(ctx context.Context, tenantID, orderID string, cause error) error {
+ message := truncateError(cause)
+ _, err := s.db.Exec(ctx, `UPDATE topup_orders SET status='failed',reconciliation_error=$2 WHERE id=$1`, orderID, message)
+ if err != nil {
+ return errors.Join(cause, err)
+ }
+ _, settingsErr := s.db.Exec(ctx, `UPDATE tenant_auto_topup_settings SET status='failed',failure_count=failure_count+1,last_error=$2,
+ next_attempt_at=now()+interval '1 hour',updated_at=now() WHERE tenant_id=$1`, tenantID, message)
+ if settingsErr != nil {
+ return errors.Join(cause, settingsErr)
+ }
+ return cause
+}
+
+func (s *Service) markAutoTopUpAttention(ctx context.Context, tenantID, orderID string, cause error) error {
+ message := truncateError(cause)
+ _, err := s.db.Exec(ctx, `UPDATE topup_orders SET status='failed',reconciliation_error=$2 WHERE id=$1`, orderID, message)
+ if err != nil {
+ return err
+ }
+ _, err = s.db.Exec(ctx, `UPDATE tenant_auto_topup_settings SET enabled=FALSE,status='action_required',last_error=$2,next_attempt_at=NULL,updated_at=now() WHERE tenant_id=$1`, tenantID, message)
+ return err
+}
+
+func (s *Service) creditAutoTopUpPaymentIntent(ctx context.Context, intent *stripe.PaymentIntent) error {
+ if intent == nil || intent.ID == "" {
+ return ErrInvalidAmount
+ }
+ tx, err := s.db.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.ReadCommitted})
+ if err != nil {
+ return err
+ }
+ defer tx.Rollback(ctx)
+ if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended($1,0))`, intent.ID); err != nil {
+ return err
+ }
+ if err := s.applyAutoTopUpPaymentIntentTx(ctx, tx, intent); err != nil {
+ return err
+ }
+ return tx.Commit(ctx)
+}
+
+func (s *Service) applyAutoTopUpPaymentIntentTx(ctx context.Context, tx pgx.Tx, intent *stripe.PaymentIntent) error {
+ if intent.Status != stripe.PaymentIntentStatusSucceeded || intent.Metadata["aigw_action"] != "auto_topup" {
+ return nil
+ }
+ orderID, tenantID := intent.Metadata["aigw_topup_order_id"], intent.Metadata["aigw_tenant_id"]
+ if orderID == "" || tenantID == "" {
+ return ErrInvalidAmount
+ }
+ var amountMinor, amountMicros int64
+ var currency, status, storedPI string
+ if err := tx.QueryRow(ctx, `SELECT amount_minor,amount_micros,currency,status,COALESCE(stripe_payment_intent_id,'') FROM topup_orders WHERE id=$1 AND tenant_id=$2 FOR UPDATE`, orderID, tenantID).Scan(&amountMinor, &amountMicros, &currency, &status, &storedPI); err != nil {
+ return err
+ }
+ if intent.Amount != amountMinor || string(intent.Currency) != currency || (storedPI != "" && storedPI != intent.ID) {
+ return ErrInvalidAmount
+ }
+ if intent.AmountReceived != 0 && intent.AmountReceived != amountMinor {
+ return ErrInvalidAmount
+ }
+ if intent.Customer != nil && intent.Customer.ID != "" {
+ if _, err := tx.Exec(ctx, `UPDATE topup_orders SET stripe_customer_id=$2 WHERE id=$1`, orderID, intent.Customer.ID); err != nil {
+ return err
+ }
+ }
+ if status != "paid" {
+ if _, err := tx.Exec(ctx, `INSERT INTO tenant_wallets (tenant_id,currency) VALUES ($1,$2) ON CONFLICT DO NOTHING`, tenantID, currency); err != nil {
+ return err
+ }
+ var balance, reserved int64
+ if err := tx.QueryRow(ctx, `SELECT balance_micros,reserved_micros FROM tenant_wallets WHERE tenant_id=$1 FOR UPDATE`, tenantID).Scan(&balance, &reserved); err != nil {
+ return err
+ }
+ if balance > int64(^uint64(0)>>1)-amountMicros {
+ return ErrInvalidAmount
+ }
+ newBalance := balance + amountMicros
+ if newBalance < reserved {
+ return ErrInsufficientBalance
+ }
+ var inserted int64
+ if err := tx.QueryRow(ctx, `INSERT INTO billing_ledger (tenant_id,currency,amount_micros,balance_after_micros,kind,source_type,source_id,description)
+ VALUES ($1,$2,$3,$4,'topup','stripe_payment_intent',$5,'Automatic prepaid balance top-up')
+ ON CONFLICT (source_type,source_id) DO NOTHING RETURNING amount_micros`, tenantID, currency, amountMicros, newBalance, intent.ID).Scan(&inserted); err != nil && !errors.Is(err, pgx.ErrNoRows) {
+ return err
+ }
+ if inserted != 0 {
+ if _, err := tx.Exec(ctx, `UPDATE tenant_wallets SET balance_micros=$2,updated_at=now() WHERE tenant_id=$1`, tenantID, newBalance); err != nil {
+ return err
+ }
+ }
+ }
+ if _, err := tx.Exec(ctx, `UPDATE topup_orders SET status='paid',paid_at=COALESCE(paid_at,now()),stripe_payment_intent_id=$2,reconciliation_status='ok',reconciled_at=now(),reconciliation_error='' WHERE id=$1`, orderID, intent.ID); err != nil {
+ return err
+ }
+ _, err := tx.Exec(ctx, `UPDATE tenant_auto_topup_settings SET status='ready',last_error='',failure_count=0,last_succeeded_at=now(),next_attempt_at=NULL,updated_at=now() WHERE tenant_id=$1`, tenantID)
+ return err
+}
+
+func (s *Service) applyAutoTopUpPaymentIntentFailure(ctx context.Context, intent *stripe.PaymentIntent) error {
+ if intent == nil || intent.Metadata["aigw_action"] != "auto_topup" {
+ return nil
+ }
+ tenantID, orderID := intent.Metadata["aigw_tenant_id"], intent.Metadata["aigw_topup_order_id"]
+ message := "automatic top-up payment failed"
+ if intent.LastPaymentError != nil && intent.LastPaymentError.Msg != "" {
+ message = intent.LastPaymentError.Msg
+ }
+ if intent.Status == stripe.PaymentIntentStatusRequiresAction || intent.Status == stripe.PaymentIntentStatusRequiresPaymentMethod || intent.Status == stripe.PaymentIntentStatusCanceled {
+ return s.markAutoTopUpAttention(ctx, tenantID, orderID, errors.New(message))
+ }
+ return s.failAutoTopUp(ctx, tenantID, orderID, errors.New(message))
+}
+
+func (s *Service) applyAutoTopUpPaymentIntentFailureTx(ctx context.Context, tx pgx.Tx, intent *stripe.PaymentIntent) error {
+ if intent == nil || intent.Metadata["aigw_action"] != "auto_topup" {
+ return nil
+ }
+ tenantID, orderID := intent.Metadata["aigw_tenant_id"], intent.Metadata["aigw_topup_order_id"]
+ if tenantID == "" || orderID == "" {
+ return ErrInvalidAmount
+ }
+ message := "automatic top-up payment failed"
+ if intent.LastPaymentError != nil && intent.LastPaymentError.Msg != "" {
+ message = truncateError(errors.New(intent.LastPaymentError.Msg))
+ }
+ if _, err := tx.Exec(ctx, `UPDATE topup_orders SET status='failed',reconciliation_error=$2,stripe_payment_intent_id=COALESCE(NULLIF($3,''),stripe_payment_intent_id) WHERE id=$1 AND status='pending'`, orderID, message, intent.ID); err != nil {
+ return err
+ }
+ status, enabled := "failed", true
+ if intent.Status == stripe.PaymentIntentStatusRequiresAction || intent.Status == stripe.PaymentIntentStatusRequiresPaymentMethod || intent.Status == stripe.PaymentIntentStatusCanceled {
+ status, enabled = "action_required", false
+ }
+ _, err := tx.Exec(ctx, `UPDATE tenant_auto_topup_settings SET enabled=$2,status=$3,last_error=$4,failure_count=failure_count+1,
+ next_attempt_at=CASE WHEN $2 THEN now()+interval '1 hour' ELSE NULL END,updated_at=now() WHERE tenant_id=$1`, tenantID, enabled, status, message)
+ return err
+}
diff --git a/internal/billing/auto_topup_test.go b/internal/billing/auto_topup_test.go
new file mode 100644
index 0000000..85edb3d
--- /dev/null
+++ b/internal/billing/auto_topup_test.go
@@ -0,0 +1,155 @@
+package billing
+
+import (
+ "context"
+ "encoding/json"
+ "fmt"
+ "os"
+ "strings"
+ "testing"
+ "time"
+
+ "aigw/internal/controlplane"
+
+ "github.com/stripe/stripe-go/v86"
+)
+
+func TestAutoTopUpReturnURL(t *testing.T) {
+ success := autoTopUpReturnURL("https://console.example.test/admin/?topup=success", true)
+ if !strings.Contains(success, "autotopup=setup") || !strings.Contains(success, "session_id={CHECKOUT_SESSION_ID}") {
+ t.Fatalf("unexpected setup return URL %q", success)
+ }
+ cancel := autoTopUpReturnURL("https://console.example.test/admin/?topup=cancel&session_id=old", false)
+ if !strings.Contains(cancel, "autotopup=cancel") || strings.Contains(cancel, "session_id=") {
+ t.Fatalf("unexpected setup cancel URL %q", cancel)
+ }
+}
+
+func TestAutomaticTopUpSetupAndCreditAreIdempotentPostgres(t *testing.T) {
+ databaseURL := os.Getenv("AIGW_TEST_DATABASE_URL")
+ if databaseURL == "" {
+ t.Skip("AIGW_TEST_DATABASE_URL is not set")
+ }
+ ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
+ defer cancel()
+ if err := controlplane.MigrateDatabase(ctx, databaseURL); err != nil {
+ t.Fatal(err)
+ }
+ service, err := New(ctx, Options{
+ DatabaseURL: databaseURL, Currency: "usd", MinTopUpMinor: 500, MaxTopUpMinor: 1_000_000,
+ StripeEnabled: true, StripeAPIKey: "rk_test_placeholder", StripeWebhookSecret: "whsec_integration_test",
+ })
+ if err != nil {
+ t.Fatal(err)
+ }
+ t.Cleanup(service.Close)
+
+ suffix := time.Now().UnixNano()
+ eventID := fmt.Sprintf("evt_auto_topup_%d", suffix)
+ var tenantID string
+ if err := service.db.QueryRow(ctx, `INSERT INTO tenants (slug,name) VALUES ($1,'Auto top-up integration') RETURNING id::text`, fmt.Sprintf("auto-topup-%d", suffix)).Scan(&tenantID); err != nil {
+ t.Fatal(err)
+ }
+ if _, err := service.db.Exec(ctx, `INSERT INTO tenant_wallets (tenant_id,currency) VALUES ($1,'usd')`, tenantID); err != nil {
+ t.Fatal(err)
+ }
+ t.Cleanup(func() {
+ cleanupCtx := context.Background()
+ if _, cleanupErr := service.db.Exec(cleanupCtx, `DELETE FROM stripe_webhook_events WHERE event_id=$1`, eventID); cleanupErr != nil {
+ t.Errorf("cleanup automatic top-up webhook: %v", cleanupErr)
+ }
+ for _, query := range []string{
+ `DELETE FROM billing_ledger WHERE tenant_id=$1`,
+ `DELETE FROM topup_orders WHERE tenant_id=$1`,
+ `DELETE FROM tenant_auto_topup_settings WHERE tenant_id=$1`,
+ `DELETE FROM stripe_customers WHERE tenant_id=$1`,
+ `DELETE FROM tenant_wallets WHERE tenant_id=$1`,
+ `DELETE FROM tenants WHERE id=$1`,
+ } {
+ if _, cleanupErr := service.db.Exec(cleanupCtx, query, tenantID); cleanupErr != nil {
+ t.Errorf("cleanup automatic top-up integration data: %v", cleanupErr)
+ }
+ }
+ })
+
+ customerID := fmt.Sprintf("cus_auto_%d", suffix)
+ paymentMethodID := fmt.Sprintf("pm_auto_%d", suffix)
+ setupIntentID := fmt.Sprintf("seti_auto_%d", suffix)
+ setupSessionID := fmt.Sprintf("cs_auto_%d", suffix)
+ service.retrieveStripeSetupIntent = func(context.Context, string, *stripe.SetupIntentRetrieveParams) (*stripe.SetupIntent, error) {
+ return &stripe.SetupIntent{
+ ID: setupIntentID, Status: stripe.SetupIntentStatusSucceeded,
+ Customer: &stripe.Customer{ID: customerID},
+ PaymentMethod: &stripe.PaymentMethod{ID: paymentMethodID, Type: stripe.PaymentMethodTypeCard,
+ Card: &stripe.PaymentMethodCard{Brand: stripe.PaymentMethodCardBrandVisa, Last4: "4242", ExpMonth: 12, ExpYear: 2035}},
+ }, nil
+ }
+ if _, err := service.db.Exec(ctx, `INSERT INTO tenant_auto_topup_settings (tenant_id,threshold_micros,topup_amount_minor,stripe_setup_session_id) VALUES ($1,5000000,2000,$2)`, tenantID, setupSessionID); err != nil {
+ t.Fatal(err)
+ }
+ raw, err := json.Marshal(map[string]any{
+ "id": setupSessionID, "object": "checkout.session", "client_reference_id": tenantID,
+ "customer": customerID, "customer_email": "developer@example.test", "setup_intent": setupIntentID,
+ "metadata": map[string]string{"aigw_action": autoTopUpAction, "aigw_tenant_id": tenantID},
+ })
+ if err != nil {
+ t.Fatal(err)
+ }
+ event := stripe.Event{ID: eventID, Type: stripe.EventTypeCheckoutSessionCompleted, Data: &stripe.EventData{Raw: raw}}
+ if err := service.processStripeEvent(ctx, event); err != nil {
+ t.Fatal(err)
+ }
+ if err := service.processStripeEvent(ctx, event); err != nil {
+ t.Fatalf("replayed setup event: %v", err)
+ }
+ settings, err := service.UpdateAutoTopUp(ctx, UpdateAutoTopUpInput{
+ TenantID: tenantID, Enabled: true, ThresholdMicros: 1_000_000, TopUpAmountMinor: 2000,
+ })
+ if err != nil {
+ t.Fatal(err)
+ }
+ if !settings.Enabled || !settings.PaymentMethodConfigured || settings.PaymentMethodLast4 != "4242" {
+ t.Fatalf("unexpected settings: %+v", settings)
+ }
+
+ var createdIntent *stripe.PaymentIntent
+ stripeCalls := 0
+ service.createStripePaymentIntent = func(_ context.Context, params *stripe.PaymentIntentCreateParams) (*stripe.PaymentIntent, error) {
+ stripeCalls++
+ createdIntent = &stripe.PaymentIntent{
+ ID: fmt.Sprintf("pi_auto_%d", suffix), Status: stripe.PaymentIntentStatusSucceeded,
+ Amount: *params.Amount, AmountReceived: *params.Amount, Currency: stripe.Currency(*params.Currency),
+ Customer: &stripe.Customer{ID: *params.Customer}, PaymentMethod: &stripe.PaymentMethod{ID: *params.PaymentMethod},
+ Metadata: params.Metadata,
+ }
+ return createdIntent, nil
+ }
+ processed, err := service.processAutoTopUpOnce(ctx)
+ if err != nil || !processed {
+ t.Fatalf("process automatic top-up: processed=%v err=%v", processed, err)
+ }
+ processed, err = service.processAutoTopUpOnce(ctx)
+ if err != nil || processed {
+ t.Fatalf("second automatic top-up: processed=%v err=%v", processed, err)
+ }
+ if stripeCalls != 1 {
+ t.Fatalf("Stripe calls = %d, want 1", stripeCalls)
+ }
+ if err := service.creditAutoTopUpPaymentIntent(ctx, createdIntent); err != nil {
+ t.Fatalf("replayed successful PaymentIntent: %v", err)
+ }
+
+ var balance, ledgerCount, paidOrders int64
+ if err := service.db.QueryRow(ctx, `SELECT balance_micros FROM tenant_wallets WHERE tenant_id=$1`, tenantID).Scan(&balance); err != nil {
+ t.Fatal(err)
+ }
+ if err := service.db.QueryRow(ctx, `SELECT count(*) FROM billing_ledger WHERE tenant_id=$1 AND source_type='stripe_payment_intent'`, tenantID).Scan(&ledgerCount); err != nil {
+ t.Fatal(err)
+ }
+ if err := service.db.QueryRow(ctx, `SELECT count(*) FROM topup_orders WHERE tenant_id=$1 AND trigger_type='auto' AND status='paid'`, tenantID).Scan(&paidOrders); err != nil {
+ t.Fatal(err)
+ }
+ if balance != 20_000_000 || ledgerCount != 1 || paidOrders != 1 {
+ t.Fatalf("balance=%d ledger=%d paid_orders=%d", balance, ledgerCount, paidOrders)
+ }
+}
diff --git a/internal/billing/ledger.go b/internal/billing/ledger.go
index 2eb3d87..7a00081 100644
--- a/internal/billing/ledger.go
+++ b/internal/billing/ledger.go
@@ -77,7 +77,7 @@ func (s *Service) ListTopUpOrders(ctx context.Context, tenantID string, limit in
if limit < 1 || limit > 200 {
limit = 50
}
- query := `SELECT id::text,tenant_id::text,amount_minor,amount_micros,currency,status,
+ query := `SELECT id::text,tenant_id::text,amount_minor,amount_micros,currency,status,trigger_type,
COALESCE(stripe_session_id,''),COALESCE(checkout_url,''),created_at,paid_at,
COALESCE(stripe_customer_id,''),COALESCE(stripe_payment_intent_id,''),COALESCE(stripe_charge_id,''),
COALESCE(stripe_invoice_id,''),COALESCE(invoice_url,''),COALESCE(invoice_pdf_url,''),COALESCE(receipt_url,''),
@@ -98,7 +98,7 @@ func (s *Service) ListTopUpOrders(ctx context.Context, tenantID string, limit in
result := make([]TopUpOrder, 0)
for rows.Next() {
var item TopUpOrder
- if err := rows.Scan(&item.ID, &item.TenantID, &item.AmountMinor, &item.AmountMicros, &item.Currency, &item.Status,
+ if err := rows.Scan(&item.ID, &item.TenantID, &item.AmountMinor, &item.AmountMicros, &item.Currency, &item.Status, &item.TriggerType,
&item.StripeSessionID, &item.CheckoutURL, &item.CreatedAt, &item.PaidAt, &item.StripeCustomerID,
&item.StripePaymentIntentID, &item.StripeChargeID, &item.StripeInvoiceID, &item.InvoiceURL,
&item.InvoicePDFURL, &item.ReceiptURL, &item.RefundedMicros, &item.DisputedMicros,
@@ -112,7 +112,7 @@ func (s *Service) ListTopUpOrders(ctx context.Context, tenantID string, limit in
func (s *Service) GetTopUpOrder(ctx context.Context, tenantID, orderID string) (TopUpOrder, error) {
var result TopUpOrder
- query := `SELECT id::text,tenant_id::text,amount_minor,amount_micros,currency,status,
+ query := `SELECT id::text,tenant_id::text,amount_minor,amount_micros,currency,status,trigger_type,
COALESCE(stripe_session_id,''),COALESCE(checkout_url,''),created_at,paid_at,
COALESCE(stripe_customer_id,''),COALESCE(stripe_payment_intent_id,''),COALESCE(stripe_charge_id,''),
COALESCE(stripe_invoice_id,''),COALESCE(invoice_url,''),COALESCE(invoice_pdf_url,''),COALESCE(receipt_url,''),
@@ -123,7 +123,7 @@ func (s *Service) GetTopUpOrder(ctx context.Context, tenantID, orderID string) (
args = append(args, tenantID)
}
err := s.db.QueryRow(ctx, query, args...).Scan(&result.ID, &result.TenantID, &result.AmountMinor, &result.AmountMicros,
- &result.Currency, &result.Status, &result.StripeSessionID, &result.CheckoutURL, &result.CreatedAt, &result.PaidAt,
+ &result.Currency, &result.Status, &result.TriggerType, &result.StripeSessionID, &result.CheckoutURL, &result.CreatedAt, &result.PaidAt,
&result.StripeCustomerID, &result.StripePaymentIntentID, &result.StripeChargeID, &result.StripeInvoiceID,
&result.InvoiceURL, &result.InvoicePDFURL, &result.ReceiptURL, &result.RefundedMicros, &result.DisputedMicros,
&result.ReconciliationStatus, &result.ReconciledAt, &result.ReconciliationError)
diff --git a/internal/billing/operations.go b/internal/billing/operations.go
index a6dc653..461c59a 100644
--- a/internal/billing/operations.go
+++ b/internal/billing/operations.go
@@ -19,13 +19,13 @@ func (s *Service) CreatePortalSession(ctx context.Context, tenantID string) (Por
if !s.stripeEnabled || s.stripeClient == nil {
return PortalResult{}, ErrStripeDisabled
}
- var customerID string
- if err := s.db.QueryRow(ctx, `SELECT stripe_customer_id FROM stripe_customers WHERE tenant_id=$1`, tenantID).Scan(&customerID); err != nil {
- if errors.Is(err, pgx.ErrNoRows) {
- return PortalResult{}, errors.New("no Stripe customer exists for this account")
- }
+ customerID, err := s.ensureStripeCustomer(ctx, tenantID)
+ if err != nil {
return PortalResult{}, err
}
+ if customerID == "" {
+ return PortalResult{}, errors.New("no Stripe customer exists for this account")
+ }
session, err := s.stripeClient.V1BillingPortalSessions.Create(ctx, &stripe.BillingPortalSessionCreateParams{
Customer: stripe.String(customerID), ReturnURL: stripe.String(s.stripePortalReturnURL),
})
@@ -353,6 +353,12 @@ func (s *Service) RunStripeOperations(ctx context.Context) {
cancel()
s.refreshOperationalMetrics(ctx)
for {
+ for i := 0; i < 4; i++ {
+ ok, _ := s.processAutoTopUpOnce(ctx)
+ if !ok {
+ break
+ }
+ }
for i := 0; i < 8; i++ {
ok, _ := s.processRefundOperation(ctx)
if !ok {
@@ -793,6 +799,74 @@ func (s *Service) Reconcile(ctx context.Context, limit int) (ReconciliationResul
return s.failReconciliation(ctx, result, fmt.Errorf("record clean reconciliation for order %s: %w", item.id, err))
}
}
+ if s.retrieveStripePaymentIntent == nil {
+ return s.failReconciliation(ctx, result, errors.New("Stripe PaymentIntent retrieval is unavailable"))
+ }
+ piRows, err := s.db.Query(ctx, `SELECT id::text,stripe_payment_intent_id,status,amount_minor,currency FROM topup_orders
+ WHERE trigger_type='auto' AND stripe_payment_intent_id IS NOT NULL ORDER BY created_at DESC LIMIT $1`, limit)
+ if err != nil {
+ return s.failReconciliation(ctx, result, err)
+ }
+ type paymentIntentOrder struct {
+ id, paymentIntent, status, currency string
+ amount int64
+ }
+ var paymentIntentOrders []paymentIntentOrder
+ for piRows.Next() {
+ var item paymentIntentOrder
+ if err := piRows.Scan(&item.id, &item.paymentIntent, &item.status, &item.amount, &item.currency); err != nil {
+ piRows.Close()
+ return s.failReconciliation(ctx, result, err)
+ }
+ paymentIntentOrders = append(paymentIntentOrders, item)
+ }
+ piRows.Close()
+ for _, item := range paymentIntentOrders {
+ intent, retrieveErr := s.retrieveStripePaymentIntent(ctx, item.paymentIntent, &stripe.PaymentIntentRetrieveParams{})
+ result.CheckedOrders++
+ if retrieveErr != nil {
+ message := truncateError(retrieveErr)
+ if err := s.updateOrderReconciliation(ctx, item.id, "mismatch", message); err != nil {
+ return s.failReconciliation(ctx, result, fmt.Errorf("record automatic top-up retrieval failure for order %s: %w", item.id, err))
+ }
+ result.Mismatches = append(result.Mismatches, map[string]any{"order_id": item.id, "type": "payment_intent_retrieve_failed", "error": message})
+ continue
+ }
+ if intent == nil || intent.Amount != item.amount || string(intent.Currency) != item.currency {
+ if err := s.updateOrderReconciliation(ctx, item.id, "mismatch", "amount or currency mismatch"); err != nil {
+ return s.failReconciliation(ctx, result, err)
+ }
+ result.Mismatches = append(result.Mismatches, map[string]any{"order_id": item.id, "type": "payment_intent_mismatch"})
+ continue
+ }
+ expectedPaid := item.status == "paid" || item.status == "partially_refunded" || item.status == "refunded" || item.status == "disputed"
+ stripePaid := intent.Status == stripe.PaymentIntentStatusSucceeded
+ if stripePaid && !expectedPaid {
+ if repairErr := s.creditAutoTopUpPaymentIntent(ctx, intent); repairErr != nil {
+ message := truncateError(repairErr)
+ result.Mismatches = append(result.Mismatches, map[string]any{"order_id": item.id, "type": "payment_intent_repair_failed", "error": message})
+ if err := s.updateOrderReconciliation(ctx, item.id, "mismatch", message); err != nil {
+ return s.failReconciliation(ctx, result, err)
+ }
+ continue
+ }
+ result.Repairs = append(result.Repairs, map[string]any{"order_id": item.id, "type": "credited_paid_payment_intent"})
+ if err := s.updateOrderReconciliation(ctx, item.id, "repaired", ""); err != nil {
+ return s.failReconciliation(ctx, result, err)
+ }
+ continue
+ }
+ if expectedPaid != stripePaid {
+ result.Mismatches = append(result.Mismatches, map[string]any{"order_id": item.id, "type": "payment_intent_state_mismatch", "local_status": item.status, "stripe_status": intent.Status})
+ if err := s.updateOrderReconciliation(ctx, item.id, "mismatch", "payment state mismatch"); err != nil {
+ return s.failReconciliation(ctx, result, err)
+ }
+ continue
+ }
+ if err := s.updateOrderReconciliation(ctx, item.id, "ok", ""); err != nil {
+ return s.failReconciliation(ctx, result, err)
+ }
+ }
result.MismatchCount = int64(len(result.Mismatches))
result.Status = "clean"
if result.MismatchCount > 0 {
diff --git a/internal/billing/profile.go b/internal/billing/profile.go
new file mode 100644
index 0000000..3b08a3c
--- /dev/null
+++ b/internal/billing/profile.go
@@ -0,0 +1,197 @@
+package billing
+
+import (
+ "context"
+ "errors"
+ "fmt"
+ "net/mail"
+ "strings"
+ "unicode"
+
+ "github.com/jackc/pgx/v5"
+ "github.com/stripe/stripe-go/v86"
+)
+
+type stripeCustomerCreator func(context.Context, *stripe.CustomerCreateParams) (*stripe.Customer, error)
+type stripeCustomerUpdater func(context.Context, string, *stripe.CustomerUpdateParams) (*stripe.Customer, error)
+
+func (s *Service) GetBillingProfile(ctx context.Context, tenantID string) (BillingProfile, error) {
+ tenantID = strings.TrimSpace(tenantID)
+ if tenantID == "" {
+ return BillingProfile{}, ErrBillingAccountNotFound
+ }
+ var result BillingProfile
+ err := s.db.QueryRow(ctx, `SELECT t.id::text,
+ COALESCE(p.legal_name,t.name),COALESCE(p.billing_email,sc.email,''),
+ COALESCE(p.address_line1,''),COALESCE(p.address_line2,''),COALESCE(p.city,''),
+ COALESCE(p.region,''),COALESCE(p.postal_code,''),COALESCE(p.country,''),
+ p.tenant_id IS NOT NULL,sc.stripe_customer_id IS NOT NULL,
+ COALESCE(p.stripe_sync_status,CASE WHEN sc.stripe_customer_id IS NOT NULL THEN 'checkout_managed' ELSE 'not_configured' END),
+ p.stripe_synced_at,COALESCE(p.stripe_sync_error,''),p.updated_at
+ FROM tenants t
+ LEFT JOIN tenant_billing_profiles p ON p.tenant_id=t.id
+ LEFT JOIN stripe_customers sc ON sc.tenant_id=t.id
+ WHERE t.id=$1`, tenantID).Scan(&result.TenantID, &result.LegalName, &result.BillingEmail,
+ &result.AddressLine1, &result.AddressLine2, &result.City, &result.Region, &result.PostalCode,
+ &result.Country, &result.Configured, &result.StripeCustomerConfigured, &result.StripeSyncStatus,
+ &result.StripeSyncedAt, &result.StripeSyncError, &result.UpdatedAt)
+ if errors.Is(err, pgx.ErrNoRows) {
+ return BillingProfile{}, ErrBillingAccountNotFound
+ }
+ if err != nil {
+ return BillingProfile{}, fmt.Errorf("get billing profile: %w", err)
+ }
+ return result, nil
+}
+
+func (s *Service) UpdateBillingProfile(ctx context.Context, input UpdateBillingProfileInput) (BillingProfile, error) {
+ normalized, err := normalizeBillingProfile(input)
+ if err != nil {
+ return BillingProfile{}, err
+ }
+ status := "disabled"
+ if s.stripeEnabled {
+ status = "pending"
+ }
+ _, err = s.db.Exec(ctx, `INSERT INTO tenant_billing_profiles
+ (tenant_id,legal_name,billing_email,address_line1,address_line2,city,region,postal_code,country,stripe_sync_status,stripe_synced_at,stripe_sync_error)
+ VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,NULL,'')
+ ON CONFLICT (tenant_id) DO UPDATE SET legal_name=EXCLUDED.legal_name,
+ billing_email=EXCLUDED.billing_email,address_line1=EXCLUDED.address_line1,address_line2=EXCLUDED.address_line2,
+ city=EXCLUDED.city,region=EXCLUDED.region,postal_code=EXCLUDED.postal_code,country=EXCLUDED.country,
+ stripe_sync_status=EXCLUDED.stripe_sync_status,stripe_synced_at=NULL,stripe_sync_error='',updated_at=now()`,
+ normalized.TenantID, normalized.LegalName, normalized.BillingEmail, normalized.AddressLine1,
+ normalized.AddressLine2, normalized.City, normalized.Region, normalized.PostalCode, normalized.Country, status)
+ if err != nil {
+ return BillingProfile{}, fmt.Errorf("save billing profile: %w", err)
+ }
+ if !s.stripeEnabled {
+ return s.GetBillingProfile(ctx, normalized.TenantID)
+ }
+ if _, err := s.ensureStripeCustomer(ctx, normalized.TenantID); err != nil {
+ return BillingProfile{}, err
+ }
+ return s.GetBillingProfile(ctx, normalized.TenantID)
+}
+
+func (s *Service) ensureStripeCustomer(ctx context.Context, tenantID string) (string, error) {
+ var profile BillingProfile
+ err := s.db.QueryRow(ctx, `SELECT tenant_id::text,legal_name,billing_email,address_line1,address_line2,
+ city,region,postal_code,country,TRUE FROM tenant_billing_profiles WHERE tenant_id=$1`, tenantID).
+ Scan(&profile.TenantID, &profile.LegalName, &profile.BillingEmail, &profile.AddressLine1,
+ &profile.AddressLine2, &profile.City, &profile.Region, &profile.PostalCode, &profile.Country, &profile.Configured)
+ if errors.Is(err, pgx.ErrNoRows) {
+ var customerID string
+ err = s.db.QueryRow(ctx, `SELECT stripe_customer_id FROM stripe_customers WHERE tenant_id=$1`, tenantID).Scan(&customerID)
+ if errors.Is(err, pgx.ErrNoRows) {
+ return "", nil
+ }
+ return customerID, err
+ }
+ if err != nil {
+ return "", fmt.Errorf("load billing profile for Stripe: %w", err)
+ }
+ if !s.stripeEnabled || s.createStripeCustomer == nil || s.updateStripeCustomer == nil {
+ return "", ErrStripeDisabled
+ }
+ var customerID string
+ err = s.db.QueryRow(ctx, `SELECT stripe_customer_id FROM stripe_customers WHERE tenant_id=$1`, tenantID).Scan(&customerID)
+ if err != nil && !errors.Is(err, pgx.ErrNoRows) {
+ return "", fmt.Errorf("load Stripe customer: %w", err)
+ }
+ if customerID == "" {
+ params := billingProfileCustomerCreateParams(profile)
+ params.SetIdempotencyKey("aigw_customer_" + tenantID)
+ customer, createErr := s.createStripeCustomer(ctx, params)
+ if createErr != nil || customer == nil || strings.TrimSpace(customer.ID) == "" {
+ if createErr == nil {
+ createErr = errors.New("Stripe returned an incomplete Customer")
+ }
+ return "", s.failBillingProfileSync(ctx, tenantID, createErr)
+ }
+ customerID = customer.ID
+ } else {
+ customer, updateErr := s.updateStripeCustomer(ctx, customerID, billingProfileCustomerUpdateParams(profile))
+ if updateErr != nil || customer == nil || strings.TrimSpace(customer.ID) == "" {
+ if updateErr == nil {
+ updateErr = errors.New("Stripe returned an incomplete Customer")
+ }
+ return "", s.failBillingProfileSync(ctx, tenantID, updateErr)
+ }
+ }
+ if _, err := s.db.Exec(ctx, `INSERT INTO stripe_customers (tenant_id,stripe_customer_id,email) VALUES ($1,$2,$3)
+ ON CONFLICT (tenant_id) DO UPDATE SET stripe_customer_id=EXCLUDED.stripe_customer_id,email=EXCLUDED.email,updated_at=now()`,
+ tenantID, customerID, profile.BillingEmail); err != nil {
+ return "", fmt.Errorf("persist Stripe customer: %w", err)
+ }
+ if _, err := s.db.Exec(ctx, `UPDATE tenant_billing_profiles SET stripe_sync_status='synced',stripe_synced_at=now(),
+ stripe_sync_error='',updated_at=now() WHERE tenant_id=$1`, tenantID); err != nil {
+ return "", fmt.Errorf("record billing profile Stripe synchronization: %w", err)
+ }
+ return customerID, nil
+}
+
+func (s *Service) failBillingProfileSync(ctx context.Context, tenantID string, cause error) error {
+ _, _ = s.db.Exec(ctx, `UPDATE tenant_billing_profiles SET stripe_sync_status='failed',
+ stripe_sync_error='Stripe customer synchronization failed',stripe_synced_at=NULL,updated_at=now() WHERE tenant_id=$1`, tenantID)
+ return fmt.Errorf("%w: %v", ErrBillingProfileSync, cause)
+}
+
+func billingProfileCustomerCreateParams(profile BillingProfile) *stripe.CustomerCreateParams {
+ return &stripe.CustomerCreateParams{
+ Name: stripe.String(profile.LegalName), BusinessName: stripe.String(profile.LegalName),
+ Email: stripe.String(profile.BillingEmail), Address: billingProfileAddress(profile),
+ Metadata: map[string]string{"aigw_tenant_id": profile.TenantID},
+ }
+}
+
+func billingProfileCustomerUpdateParams(profile BillingProfile) *stripe.CustomerUpdateParams {
+ return &stripe.CustomerUpdateParams{
+ Name: stripe.String(profile.LegalName), BusinessName: stripe.String(profile.LegalName),
+ Email: stripe.String(profile.BillingEmail), Address: billingProfileAddress(profile),
+ Metadata: map[string]string{"aigw_tenant_id": profile.TenantID},
+ }
+}
+
+func billingProfileAddress(profile BillingProfile) *stripe.AddressParams {
+ return &stripe.AddressParams{Line1: stripe.String(profile.AddressLine1), Line2: stripe.String(profile.AddressLine2),
+ City: stripe.String(profile.City), State: stripe.String(profile.Region), PostalCode: stripe.String(profile.PostalCode),
+ Country: stripe.String(profile.Country)}
+}
+
+func normalizeBillingProfile(input UpdateBillingProfileInput) (UpdateBillingProfileInput, error) {
+ input.TenantID = strings.TrimSpace(input.TenantID)
+ if input.TenantID == "" {
+ return UpdateBillingProfileInput{}, fmt.Errorf("%w: tenant is required", ErrInvalidBillingProfile)
+ }
+ var err error
+ for _, field := range []struct {
+ value *string
+ name string
+ max int
+ required bool
+ }{
+ {&input.LegalName, "legal name", 150, true}, {&input.BillingEmail, "billing email", 254, true},
+ {&input.AddressLine1, "address line 1", 200, true}, {&input.AddressLine2, "address line 2", 200, false},
+ {&input.City, "city", 100, true}, {&input.Region, "state or region", 100, false},
+ {&input.PostalCode, "postal code", 32, true},
+ } {
+ *field.value = strings.TrimSpace(*field.value)
+ if field.required && *field.value == "" {
+ return UpdateBillingProfileInput{}, fmt.Errorf("%w: %s is required", ErrInvalidBillingProfile, field.name)
+ }
+ if len([]rune(*field.value)) > field.max || strings.IndexFunc(*field.value, unicode.IsControl) >= 0 {
+ return UpdateBillingProfileInput{}, fmt.Errorf("%w: %s is invalid", ErrInvalidBillingProfile, field.name)
+ }
+ }
+ parsedEmail, err := mail.ParseAddress(input.BillingEmail)
+ if err != nil || !strings.EqualFold(parsedEmail.Address, input.BillingEmail) {
+ return UpdateBillingProfileInput{}, fmt.Errorf("%w: billing email is invalid", ErrInvalidBillingProfile)
+ }
+ input.BillingEmail = strings.ToLower(parsedEmail.Address)
+ input.Country = strings.ToUpper(strings.TrimSpace(input.Country))
+ if len(input.Country) != 2 || input.Country[0] < 'A' || input.Country[0] > 'Z' || input.Country[1] < 'A' || input.Country[1] > 'Z' {
+ return UpdateBillingProfileInput{}, fmt.Errorf("%w: country must be a two-letter code", ErrInvalidBillingProfile)
+ }
+ return input, nil
+}
diff --git a/internal/billing/profile_test.go b/internal/billing/profile_test.go
new file mode 100644
index 0000000..e263bb8
--- /dev/null
+++ b/internal/billing/profile_test.go
@@ -0,0 +1,135 @@
+package billing
+
+import (
+ "context"
+ "errors"
+ "fmt"
+ "os"
+ "testing"
+ "time"
+
+ "aigw/internal/controlplane"
+
+ "github.com/stripe/stripe-go/v86"
+)
+
+func TestNormalizeBillingProfile(t *testing.T) {
+ input := UpdateBillingProfileInput{TenantID: " tenant ", LegalName: " Example Limited ", BillingEmail: "BILLING@EXAMPLE.TEST",
+ AddressLine1: " 1 Queen Street ", City: " Auckland ", PostalCode: "1010", Country: "nz"}
+ result, err := normalizeBillingProfile(input)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if result.TenantID != "tenant" || result.LegalName != "Example Limited" || result.BillingEmail != "billing@example.test" || result.Country != "NZ" {
+ t.Fatalf("unexpected normalized profile: %+v", result)
+ }
+ for _, invalid := range []UpdateBillingProfileInput{
+ {TenantID: "tenant", LegalName: "Example", BillingEmail: "not-an-email", AddressLine1: "1 Street", City: "Auckland", PostalCode: "1010", Country: "NZ"},
+ {TenantID: "tenant", LegalName: "Example", BillingEmail: "billing@example.test", AddressLine1: "1 Street", City: "Auckland", PostalCode: "1010", Country: "New Zealand"},
+ {TenantID: "tenant", BillingEmail: "billing@example.test", AddressLine1: "1 Street", City: "Auckland", PostalCode: "1010", Country: "NZ"},
+ } {
+ if _, err := normalizeBillingProfile(invalid); !errors.Is(err, ErrInvalidBillingProfile) {
+ t.Fatalf("error = %v, want invalid billing profile", err)
+ }
+ }
+}
+
+func TestBillingProfileStripeSynchronizationPostgres(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"})
+ 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,'Billing profile integration') RETURNING id::text`, fmt.Sprintf("billing-profile-%d", suffix)).Scan(&tenantID); err != nil {
+ t.Fatal(err)
+ }
+ t.Cleanup(func() {
+ cleanupCtx := context.Background()
+ for _, query := range []string{
+ `DELETE FROM tenant_billing_profiles WHERE tenant_id=$1`,
+ `DELETE FROM stripe_customers WHERE tenant_id=$1`,
+ `DELETE FROM tenants WHERE id=$1`,
+ } {
+ if _, cleanupErr := service.db.Exec(cleanupCtx, query, tenantID); cleanupErr != nil {
+ t.Errorf("cleanup billing profile: %v", cleanupErr)
+ }
+ }
+ })
+
+ customerID := fmt.Sprintf("cus_profile_%d", suffix)
+ createCalls, updateCalls := 0, 0
+ service.createStripeCustomer = func(_ context.Context, params *stripe.CustomerCreateParams) (*stripe.Customer, error) {
+ createCalls++
+ if params.IdempotencyKey == nil || *params.IdempotencyKey != "aigw_customer_"+tenantID || params.Address == nil || *params.Address.Country != "NZ" || *params.Name != "Example Limited" || params.Metadata["aigw_tenant_id"] != tenantID {
+ t.Fatalf("unexpected Stripe create params: %+v", params)
+ }
+ return &stripe.Customer{ID: customerID}, nil
+ }
+ service.updateStripeCustomer = func(_ context.Context, id string, params *stripe.CustomerUpdateParams) (*stripe.Customer, error) {
+ updateCalls++
+ if id != customerID || params.Address == nil || *params.Address.PostalCode != "1010" {
+ t.Fatalf("unexpected Stripe update: id=%s params=%+v", id, params)
+ }
+ return &stripe.Customer{ID: customerID}, nil
+ }
+
+ input := UpdateBillingProfileInput{TenantID: tenantID, LegalName: "Example Limited", BillingEmail: "billing@example.test",
+ AddressLine1: "1 Queen Street", City: "Auckland", Region: "Auckland", PostalCode: "1010", Country: "nz"}
+ profile, err := service.UpdateBillingProfile(ctx, input)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if createCalls != 1 || updateCalls != 0 || !profile.Configured || !profile.StripeCustomerConfigured || profile.StripeSyncStatus != "synced" || profile.StripeSyncedAt == nil {
+ t.Fatalf("unexpected first synchronization: create=%d update=%d profile=%+v", createCalls, updateCalls, profile)
+ }
+
+ input.LegalName = "Example API Limited"
+ profile, err = service.UpdateBillingProfile(ctx, input)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if createCalls != 1 || updateCalls != 1 || profile.LegalName != "Example API Limited" || profile.StripeSyncStatus != "synced" {
+ t.Fatalf("unexpected update synchronization: create=%d update=%d profile=%+v", createCalls, updateCalls, profile)
+ }
+
+ service.updateStripeCustomer = func(context.Context, string, *stripe.CustomerUpdateParams) (*stripe.Customer, error) {
+ return nil, errors.New("temporary Stripe outage")
+ }
+ input.City = "Wellington"
+ if _, err := service.UpdateBillingProfile(ctx, input); !errors.Is(err, ErrBillingProfileSync) {
+ t.Fatalf("error = %v, want Stripe sync error", err)
+ }
+ profile, err = service.GetBillingProfile(ctx, tenantID)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if profile.City != "Wellington" || profile.StripeSyncStatus != "failed" || profile.StripeSyncError == "" {
+ t.Fatalf("failed synchronization did not preserve local profile: %+v", profile)
+ }
+
+ service.updateStripeCustomer = func(context.Context, string, *stripe.CustomerUpdateParams) (*stripe.Customer, error) {
+ return &stripe.Customer{ID: customerID}, nil
+ }
+ if _, err := service.ensureStripeCustomer(ctx, tenantID); err != nil {
+ t.Fatal(err)
+ }
+ profile, err = service.GetBillingProfile(ctx, tenantID)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if profile.StripeSyncStatus != "synced" || profile.StripeSyncError != "" {
+ t.Fatalf("profile did not recover after retry: %+v", profile)
+ }
+}
diff --git a/internal/billing/service.go b/internal/billing/service.go
index 30f0e32..8a91839 100644
--- a/internal/billing/service.go
+++ b/internal/billing/service.go
@@ -25,24 +25,29 @@ import (
const microsPerUnit = int64(1_000_000)
type Service struct {
- db *pgxpool.Pool
- currency string
- defaultMaxOutputTokens int64
- minTopUpMinor int64
- maxTopUpMinor int64
- stripeEnabled bool
- stripeWebhookSecret string
- stripeSuccessURL string
- stripeCancelURL string
- stripePortalReturnURL string
- stripeAutomaticTax bool
- stripeProductTaxCode string
- integrationIdentifier string
- createStripeCheckout stripeCheckoutCreator
- stripeClient *stripe.Client
- settlementSpoolPath string
- metrics OperationalMetrics
- spoolMu sync.Mutex
+ db *pgxpool.Pool
+ currency string
+ defaultMaxOutputTokens int64
+ minTopUpMinor int64
+ maxTopUpMinor int64
+ stripeEnabled bool
+ stripeWebhookSecret string
+ stripeSuccessURL string
+ stripeCancelURL string
+ stripePortalReturnURL string
+ stripeAutomaticTax bool
+ stripeProductTaxCode string
+ integrationIdentifier string
+ createStripeCheckout stripeCheckoutCreator
+ createStripeCustomer stripeCustomerCreator
+ updateStripeCustomer stripeCustomerUpdater
+ retrieveStripeSetupIntent stripeSetupIntentRetriever
+ createStripePaymentIntent stripePaymentIntentCreator
+ retrieveStripePaymentIntent stripePaymentIntentRetriever
+ stripeClient *stripe.Client
+ settlementSpoolPath string
+ metrics OperationalMetrics
+ spoolMu sync.Mutex
}
func New(ctx context.Context, options Options) (*Service, error) {
@@ -68,6 +73,11 @@ 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.createStripeCustomer = service.stripeClient.V1Customers.Create
+ service.updateStripeCustomer = service.stripeClient.V1Customers.Update
+ service.retrieveStripeSetupIntent = service.stripeClient.V1SetupIntents.Retrieve
+ service.createStripePaymentIntent = service.stripeClient.V1PaymentIntents.Create
+ service.retrieveStripePaymentIntent = service.stripeClient.V1PaymentIntents.Retrieve
}
return service, nil
}
@@ -131,6 +141,21 @@ func (s *Service) Authorize(ctx context.Context, input Authorization) error {
return ErrQuotaExceeded
}
}
+ if input.Principal.MonthlySpendMicros > 0 {
+ period := time.Date(time.Now().UTC().Year(), time.Now().UTC().Month(), 1, 0, 0, 0, 0, time.UTC)
+ nextPeriod := period.AddDate(0, 1, 0)
+ 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 monthly spend quota: %w", err)
+ }
+ limit := input.Principal.MonthlySpendMicros
+ if reserved > limit || used > limit-reserved || pending > limit-used-reserved {
+ return ErrQuotaExceeded
+ }
+ }
if balance-held < reserved {
return ErrInsufficientBalance
}
@@ -200,6 +225,11 @@ func (s *Service) settle(ctx context.Context, event domain.UsageEvent) error {
if err != nil {
return fmt.Errorf("load billing reservation: %w", err)
}
+ if _, err := tx.Exec(ctx, `UPDATE api_keys
+ SET last_used_at = GREATEST(COALESCE(last_used_at, $2), $2)
+ WHERE id = $1`, keyID, event.StartedAt); err != nil {
+ return fmt.Errorf("update API key last used time: %w", err)
+ }
if status != "pending" {
return tx.Commit(ctx)
}
diff --git a/internal/billing/service_test.go b/internal/billing/service_test.go
index 6a1bb49..5e21b20 100644
--- a/internal/billing/service_test.go
+++ b/internal/billing/service_test.go
@@ -4,6 +4,7 @@ import (
"context"
"crypto/sha256"
"encoding/json"
+ "errors"
"fmt"
"net/http"
"net/http/httptest"
@@ -19,6 +20,72 @@ import (
"github.com/stripe/stripe-go/v86/webhook"
)
+func TestAuthorizeEnforcesAPIKeyMonthlySpendCapPostgres(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", DefaultMaxOutputTokens: 10})
+ if err != nil {
+ t.Fatal(err)
+ }
+ t.Cleanup(service.Close)
+ slug := fmt.Sprintf("key-budget-%d", time.Now().UnixNano())
+ var tenantID, projectID, keyID string
+ if err := service.db.QueryRow(ctx, `INSERT INTO tenants (slug,name) VALUES ($1,'Key budget') RETURNING id::text`, slug).Scan(&tenantID); err != nil {
+ t.Fatal(err)
+ }
+ if err := service.db.QueryRow(ctx, `INSERT INTO projects (tenant_id,slug,name) VALUES ($1,'default','Default') RETURNING id::text`, tenantID).Scan(&projectID); err != nil {
+ t.Fatal(err)
+ }
+ hash := sha256.Sum256([]byte(slug))
+ if err := service.db.QueryRow(ctx, `INSERT INTO api_keys (tenant_id,project_id,name,key_prefix,key_hash,monthly_spend_micros) VALUES ($1,$2,'limited','sk-test',$3,9) RETURNING id::text`, tenantID, projectID, hash[:]).Scan(&keyID); err != nil {
+ t.Fatal(err)
+ }
+ t.Cleanup(func() {
+ for _, statement := range []struct {
+ query string
+ args []any
+ }{
+ {`DELETE FROM billing_settlement_jobs WHERE request_id LIKE 'req_key_budget_%'`, nil},
+ {`DELETE FROM billing_reservations WHERE tenant_id=$1`, []any{tenantID}},
+ {`DELETE FROM tenant_wallets WHERE tenant_id=$1`, []any{tenantID}},
+ {`DELETE FROM api_keys WHERE id=$1`, []any{keyID}},
+ {`DELETE FROM projects WHERE id=$1`, []any{projectID}},
+ {`DELETE FROM tenants WHERE id=$1`, []any{tenantID}},
+ } {
+ if _, cleanupErr := service.db.Exec(context.Background(), statement.query, statement.args...); cleanupErr != nil {
+ t.Errorf("cleanup API key budget test data: %v", cleanupErr)
+ }
+ }
+ })
+ if _, err := service.db.Exec(ctx, `INSERT INTO tenant_wallets (tenant_id,currency,balance_micros) VALUES ($1,'usd',1000000)`, tenantID); err != nil {
+ t.Fatal(err)
+ }
+ model := domain.Model{ID: "model/budget", PriceCurrency: "usd", OutputPriceMicrosPerMillion: 1_000_000}
+ principal := domain.Principal{TenantID: tenantID, ProjectID: projectID, KeyID: keyID, MonthlySpendMicros: 9}
+ err = service.Authorize(ctx, Authorization{RequestID: "req_key_budget_rejected", Principal: principal, Model: model, Body: []byte(`{"max_tokens":10}`)})
+ if !errors.Is(err, ErrQuotaExceeded) {
+ t.Fatalf("Authorize error = %v, want ErrQuotaExceeded", err)
+ }
+ var rejectedReservations int
+ if err := service.db.QueryRow(ctx, `SELECT count(*) FROM billing_reservations WHERE request_id='req_key_budget_rejected'`).Scan(&rejectedReservations); err != nil {
+ t.Fatal(err)
+ }
+ if rejectedReservations != 0 {
+ t.Fatalf("quota rejection left %d reservation rows", rejectedReservations)
+ }
+ principal.MonthlySpendMicros = 10
+ 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)
+ }
+}
+
func TestUsageCostUsesFixedPointAndRoundsOnce(t *testing.T) {
usage := domain.Usage{InputTokens: 3, OutputTokens: 2, CacheReadInputTokens: 5}
cost, err := usageCost(usage, 150_000, 600_000, 30_000, 0)
diff --git a/internal/billing/stripe.go b/internal/billing/stripe.go
index 59ee9bd..032eb02 100644
--- a/internal/billing/stripe.go
+++ b/internal/billing/stripe.go
@@ -58,8 +58,13 @@ func (s *Service) CreateCheckout(ctx context.Context, input CheckoutInput) (Chec
},
}},
}
- var customerID string
- _ = s.db.QueryRow(ctx, `SELECT stripe_customer_id FROM stripe_customers WHERE tenant_id=$1`, input.TenantID).Scan(&customerID)
+ 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 {
@@ -154,6 +159,7 @@ func (s *Service) processStripeEvent(ctx context.Context, event stripe.Event) er
return s.processCheckoutEvent(ctx, event)
case stripe.EventTypeChargeSucceeded, stripe.EventTypeChargeUpdated, stripe.EventTypeChargeRefunded,
stripe.EventTypeRefundCreated, stripe.EventTypeRefundUpdated, stripe.EventTypeRefundFailed,
+ stripe.EventTypePaymentIntentSucceeded, stripe.EventTypePaymentIntentPaymentFailed, stripe.EventTypePaymentIntentCanceled,
stripe.EventTypeChargeDisputeCreated, stripe.EventTypeChargeDisputeUpdated,
stripe.EventTypeChargeDisputeClosed, stripe.EventTypeChargeDisputeFundsWithdrawn,
stripe.EventTypeChargeDisputeFundsReinstated,
@@ -176,6 +182,9 @@ func (s *Service) processCheckoutEvent(ctx context.Context, event stripe.Event)
if err := json.Unmarshal(event.Data.Raw, &session); err != nil {
return ErrInvalidAmount
}
+ if session.Metadata["aigw_action"] == autoTopUpAction {
+ return s.processAutoTopUpSetupEvent(ctx, event, &session)
+ }
if event.ID == "" || session.ID == "" || session.ClientReferenceID == "" {
return ErrInvalidAmount
}
@@ -330,6 +339,18 @@ func (s *Service) processOperationalStripeEvent(ctx context.Context, event strip
}
}
switch event.Type {
+ case stripe.EventTypePaymentIntentSucceeded, stripe.EventTypePaymentIntentPaymentFailed, stripe.EventTypePaymentIntentCanceled:
+ var intent stripe.PaymentIntent
+ if json.Unmarshal(event.Data.Raw, &intent) != nil || intent.ID == "" {
+ return ErrInvalidAmount
+ }
+ if event.Type == stripe.EventTypePaymentIntentSucceeded {
+ if err := s.applyAutoTopUpPaymentIntentTx(ctx, tx, &intent); err != nil {
+ return err
+ }
+ } else if err := s.applyAutoTopUpPaymentIntentFailureTx(ctx, tx, &intent); err != nil {
+ return err
+ }
case stripe.EventTypeChargeSucceeded, stripe.EventTypeChargeUpdated, stripe.EventTypeChargeRefunded:
var charge stripe.Charge
if json.Unmarshal(event.Data.Raw, &charge) != nil || charge.ID == "" {
diff --git a/internal/billing/types.go b/internal/billing/types.go
index fe5df1e..1633ecc 100644
--- a/internal/billing/types.go
+++ b/internal/billing/types.go
@@ -9,13 +9,18 @@ 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")
+ 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")
)
type Meter interface {
@@ -134,6 +139,36 @@ type CheckoutResult struct {
URL string `json:"url"`
}
+type BillingProfile struct {
+ TenantID string `json:"tenant_id"`
+ LegalName string `json:"legal_name"`
+ BillingEmail string `json:"billing_email"`
+ AddressLine1 string `json:"address_line1"`
+ AddressLine2 string `json:"address_line2"`
+ City string `json:"city"`
+ Region string `json:"region"`
+ PostalCode string `json:"postal_code"`
+ Country string `json:"country"`
+ Configured bool `json:"configured"`
+ StripeCustomerConfigured bool `json:"stripe_customer_configured"`
+ StripeSyncStatus string `json:"stripe_sync_status"`
+ StripeSyncedAt *time.Time `json:"stripe_synced_at,omitempty"`
+ StripeSyncError string `json:"stripe_sync_error,omitempty"`
+ UpdatedAt *time.Time `json:"updated_at,omitempty"`
+}
+
+type UpdateBillingProfileInput struct {
+ TenantID string `json:"tenant_id"`
+ LegalName string `json:"legal_name"`
+ BillingEmail string `json:"billing_email"`
+ AddressLine1 string `json:"address_line1"`
+ AddressLine2 string `json:"address_line2"`
+ City string `json:"city"`
+ Region string `json:"region"`
+ PostalCode string `json:"postal_code"`
+ Country string `json:"country"`
+}
+
type TopUpOrder struct {
ID string `json:"id"`
TenantID string `json:"tenant_id"`
@@ -141,6 +176,7 @@ type TopUpOrder struct {
AmountMicros int64 `json:"amount_micros"`
Currency string `json:"currency"`
Status string `json:"status"`
+ TriggerType string `json:"trigger_type"`
StripeSessionID string `json:"stripe_session_id,omitempty"`
CheckoutURL string `json:"checkout_url,omitempty"`
CreatedAt time.Time `json:"created_at"`
@@ -159,6 +195,44 @@ type TopUpOrder struct {
ReconciliationError string `json:"reconciliation_error,omitempty"`
}
+type AutoTopUpSettings struct {
+ TenantID string `json:"tenant_id"`
+ Currency string `json:"currency"`
+ StripeEnabled bool `json:"stripe_enabled"`
+ Enabled bool `json:"enabled"`
+ ThresholdMicros int64 `json:"threshold_micros"`
+ TopUpAmountMinor int64 `json:"topup_amount_minor"`
+ PaymentMethodConfigured bool `json:"payment_method_configured"`
+ PaymentMethodType string `json:"payment_method_type,omitempty"`
+ PaymentMethodBrand string `json:"payment_method_brand,omitempty"`
+ PaymentMethodLast4 string `json:"payment_method_last4,omitempty"`
+ PaymentMethodExpMonth int64 `json:"payment_method_exp_month,omitempty"`
+ PaymentMethodExpYear int64 `json:"payment_method_exp_year,omitempty"`
+ Status string `json:"status"`
+ LastError string `json:"last_error,omitempty"`
+ LastAttemptAt *time.Time `json:"last_attempt_at,omitempty"`
+ LastSucceededAt *time.Time `json:"last_succeeded_at,omitempty"`
+ NextAttemptAt *time.Time `json:"next_attempt_at,omitempty"`
+ UpdatedAt time.Time `json:"updated_at"`
+}
+
+type UpdateAutoTopUpInput struct {
+ TenantID string `json:"tenant_id"`
+ Enabled bool `json:"enabled"`
+ ThresholdMicros int64 `json:"threshold_micros"`
+ TopUpAmountMinor int64 `json:"topup_amount_minor"`
+}
+
+type AutoTopUpSetupInput struct {
+ TenantID string `json:"tenant_id"`
+ CustomerEmail string `json:"-"`
+}
+
+type AutoTopUpSetupResult struct {
+ SessionID string `json:"session_id"`
+ URL string `json:"url"`
+}
+
type ResolveMissingTopUpInput struct {
Reason string `json:"reason"`
}