diff options
| author | Chia <Chia@93.nz> | 2026-08-06 09:29:41 +1200 |
|---|---|---|
| committer | Chia <Chia@93.nz> | 2026-08-06 09:32:46 +1200 |
| commit | 41e322c53d7b4b796eb377d0df9c29ecd10ba431 (patch) | |
| tree | c730526150e55e39b822d5197e4a20318ecaa449 /internal/billing | |
| parent | eadb2ffe85c43cf6fc741c9823cd28eedb4a844c (diff) | |
feat: complete commercial control plane, billing, auth, and model catalog
- add PostgreSQL control-plane persistence with Redis-degraded hot reload
- implement prepaid balance, usage ledger, Stripe top-up and reconciliation
- add registration, email verification, password reset, invitations and RBAC
- support TOTP, Passkey MFA, device sessions, quotas and rate limits
- add tenant billing profiles, audit logs and operational readiness checks
- build authenticated admin console, Quickstart, Playground and usage analytics
- add public model catalog with pricing, filtering and cost estimation
- support OpenAI Responses providers and provider health failover
- validate real upstream usage reporting and balance settlement
Diffstat (limited to 'internal/billing')
| -rw-r--r-- | internal/billing/auto_topup.go | 541 | ||||
| -rw-r--r-- | internal/billing/auto_topup_test.go | 155 | ||||
| -rw-r--r-- | internal/billing/ledger.go | 8 | ||||
| -rw-r--r-- | internal/billing/operations.go | 84 | ||||
| -rw-r--r-- | internal/billing/profile.go | 197 | ||||
| -rw-r--r-- | internal/billing/profile_test.go | 135 | ||||
| -rw-r--r-- | internal/billing/service.go | 66 | ||||
| -rw-r--r-- | internal/billing/service_test.go | 67 | ||||
| -rw-r--r-- | internal/billing/stripe.go | 25 | ||||
| -rw-r--r-- | internal/billing/types.go | 88 |
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, ¤cy, &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, ¤cy, &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"` } |
