summaryrefslogtreecommitdiff
path: root/internal/billing/stripe.go
diff options
context:
space:
mode:
authorChia <Chia@93.nz>2026-08-05 22:01:29 +1200
committerChia <Chia@93.nz>2026-08-05 22:07:50 +1200
commiteadb2ffe85c43cf6fc741c9823cd28eedb4a844c (patch)
tree1aba2536d57360da403aa35c9ced58b615c7064e /internal/billing/stripe.go
parentcd0dd91ab93653631904f2ea0e574ccde6d60339 (diff)
feat: harden prepaid billing and commercial operations
Diffstat (limited to 'internal/billing/stripe.go')
-rw-r--r--internal/billing/stripe.go207
1 files changed, 196 insertions, 11 deletions
diff --git a/internal/billing/stripe.go b/internal/billing/stripe.go
index b887035..59ee9bd 100644
--- a/internal/billing/stripe.go
+++ b/internal/billing/stripe.go
@@ -39,6 +39,13 @@ func (s *Service) CreateCheckout(ctx context.Context, input CheckoutInput) (Chec
"aigw_topup_order_id": orderID,
"aigw_tenant_id": strings.TrimSpace(input.TenantID),
},
+ InvoiceCreation: &stripe.CheckoutSessionCreateInvoiceCreationParams{
+ Enabled: stripe.Bool(true),
+ InvoiceData: &stripe.CheckoutSessionCreateInvoiceCreationInvoiceDataParams{
+ Description: stripe.String("AIGW prepaid API usage credit"),
+ Metadata: map[string]string{"aigw_topup_order_id": orderID, "aigw_tenant_id": strings.TrimSpace(input.TenantID)},
+ },
+ },
LineItems: []*stripe.CheckoutSessionCreateLineItemParams{{
Quantity: stripe.Int64(1),
PriceData: &stripe.CheckoutSessionCreateLineItemPriceDataParams{
@@ -51,10 +58,27 @@ 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)
+ if customerID != "" {
+ params.Customer = stripe.String(customerID)
+ } else {
+ params.CustomerCreation = stripe.String(string(stripe.CheckoutSessionCustomerCreationAlways))
+ if strings.TrimSpace(input.CustomerEmail) != "" {
+ params.CustomerEmail = stripe.String(strings.TrimSpace(input.CustomerEmail))
+ }
+ }
+ if s.stripeAutomaticTax {
+ params.AutomaticTax = &stripe.CheckoutSessionCreateAutomaticTaxParams{Enabled: stripe.Bool(true)}
+ params.TaxIDCollection = &stripe.CheckoutSessionCreateTaxIDCollectionParams{Enabled: stripe.Bool(true)}
+ params.LineItems[0].PriceData.ProductData.TaxCode = stripe.String(s.stripeProductTaxCode)
+ }
params.SetIdempotencyKey("aigw_topup_" + orderID)
session, err := s.createStripeCheckout(ctx, params)
if err != nil {
- _, _ = s.db.Exec(ctx, `UPDATE topup_orders SET status = 'failed' WHERE id = $1 AND status = 'pending'`, orderID)
+ if _, updateErr := s.db.Exec(ctx, `UPDATE topup_orders SET status = 'failed' WHERE id = $1 AND status = 'pending'`, orderID); updateErr != nil {
+ return CheckoutResult{}, errors.Join(fmt.Errorf("create Stripe Checkout Session: %w", err), fmt.Errorf("mark top-up order failed: %w", updateErr))
+ }
return CheckoutResult{}, fmt.Errorf("create Stripe Checkout Session: %w", err)
}
if session.ID == "" || session.URL == "" {
@@ -103,7 +127,12 @@ func (s *Service) WebhookHandler() http.Handler {
http.Error(w, "invalid webhook signature", http.StatusBadRequest)
return
}
+ if err := s.recordWebhookAttempt(r.Context(), event); err != nil {
+ http.Error(w, "webhook persistence failed", http.StatusInternalServerError)
+ return
+ }
if err := s.processStripeEvent(r.Context(), event); err != nil {
+ s.recordWebhookFailure(r.Context(), event.ID, err)
if errors.Is(err, ErrInvalidAmount) || isNotFound(err) {
http.Error(w, "invalid checkout event", http.StatusBadRequest)
return
@@ -117,15 +146,29 @@ func (s *Service) WebhookHandler() http.Handler {
}
func (s *Service) processStripeEvent(ctx context.Context, event stripe.Event) error {
- typeName := string(event.Type)
switch event.Type {
case stripe.EventTypeCheckoutSessionCompleted,
stripe.EventTypeCheckoutSessionAsyncPaymentSucceeded,
stripe.EventTypeCheckoutSessionAsyncPaymentFailed,
stripe.EventTypeCheckoutSessionExpired:
+ return s.processCheckoutEvent(ctx, event)
+ case stripe.EventTypeChargeSucceeded, stripe.EventTypeChargeUpdated, stripe.EventTypeChargeRefunded,
+ stripe.EventTypeRefundCreated, stripe.EventTypeRefundUpdated, stripe.EventTypeRefundFailed,
+ stripe.EventTypeChargeDisputeCreated, stripe.EventTypeChargeDisputeUpdated,
+ stripe.EventTypeChargeDisputeClosed, stripe.EventTypeChargeDisputeFundsWithdrawn,
+ stripe.EventTypeChargeDisputeFundsReinstated,
+ stripe.EventTypeInvoiceCreated, stripe.EventTypeInvoiceFinalized,
+ stripe.EventTypeInvoicePaid, stripe.EventTypeInvoicePaymentFailed:
+ return s.processOperationalStripeEvent(ctx, event)
default:
- return nil
+ _, err := s.db.Exec(ctx, `UPDATE stripe_webhook_events SET processed_at=now(),processing_error='ignored event type'
+ WHERE event_id=$1 AND processed_at IS NULL`, event.ID)
+ return err
}
+}
+
+func (s *Service) processCheckoutEvent(ctx context.Context, event stripe.Event) error {
+ typeName := string(event.Type)
if event.Data == nil {
return ErrInvalidAmount
}
@@ -141,6 +184,9 @@ func (s *Service) processStripeEvent(ctx context.Context, event stripe.Event) er
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
+ }
tag, err := tx.Exec(ctx, `
INSERT INTO stripe_webhook_events (event_id, event_type) VALUES ($1,$2)
ON CONFLICT (event_id) DO NOTHING`, event.ID, typeName)
@@ -148,7 +194,13 @@ func (s *Service) processStripeEvent(ctx context.Context, event stripe.Event) er
return fmt.Errorf("record Stripe event: %w", err)
}
if tag.RowsAffected() == 0 {
- return tx.Commit(ctx)
+ var processed bool
+ if err := tx.QueryRow(ctx, `SELECT processed_at IS NOT NULL FROM stripe_webhook_events WHERE event_id=$1`, event.ID).Scan(&processed); err != nil {
+ return err
+ }
+ if processed {
+ return tx.Commit(ctx)
+ }
}
var tenantID, currency, status string
@@ -164,6 +216,32 @@ func (s *Service) processStripeEvent(ctx context.Context, event stripe.Event) er
if (storedSessionID != nil && *storedSessionID != session.ID) || amountMinor != session.AmountTotal || currency != string(session.Currency) {
return ErrInvalidAmount
}
+ customerID, paymentIntentID, invoiceID := "", "", ""
+ if session.Customer != nil {
+ customerID = session.Customer.ID
+ }
+ if session.PaymentIntent != nil {
+ paymentIntentID = session.PaymentIntent.ID
+ }
+ if session.Invoice != nil {
+ invoiceID = session.Invoice.ID
+ }
+ if customerID != "" {
+ 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()`, tenantID, customerID, email); err != nil {
+ return err
+ }
+ }
+ if _, err := tx.Exec(ctx, `UPDATE topup_orders SET stripe_customer_id=COALESCE(NULLIF($2,''),stripe_customer_id),
+ stripe_payment_intent_id=COALESCE(NULLIF($3,''),stripe_payment_intent_id),
+ stripe_invoice_id=COALESCE(NULLIF($4,''),stripe_invoice_id) WHERE id=$1`, session.ClientReferenceID, customerID, paymentIntentID, invoiceID); err != nil {
+ return err
+ }
if event.Type == stripe.EventTypeCheckoutSessionAsyncPaymentFailed || event.Type == stripe.EventTypeCheckoutSessionExpired {
orderStatus := "failed"
if event.Type == stripe.EventTypeCheckoutSessionExpired {
@@ -172,20 +250,24 @@ func (s *Service) processStripeEvent(ctx context.Context, event stripe.Event) er
if _, err := tx.Exec(ctx, `UPDATE topup_orders SET status = $2, stripe_session_id = COALESCE(stripe_session_id, $3) WHERE id = $1 AND status = 'pending'`, session.ClientReferenceID, orderStatus, session.ID); err != nil {
return err
}
- _, err = tx.Exec(ctx, `UPDATE stripe_webhook_events SET processed_at = now() WHERE event_id = $1`, event.ID)
+ _, err = tx.Exec(ctx, `UPDATE stripe_webhook_events SET processed_at=now(),processing_error='' WHERE event_id=$1`, event.ID)
if err != nil {
return err
}
return tx.Commit(ctx)
}
if session.PaymentStatus != stripe.CheckoutSessionPaymentStatusPaid {
- _, err = tx.Exec(ctx, `UPDATE stripe_webhook_events SET processed_at = now() WHERE event_id = $1`, event.ID)
+ _, err = tx.Exec(ctx, `UPDATE stripe_webhook_events SET processed_at=now(),processing_error='' WHERE event_id=$1`, event.ID)
if err != nil {
return err
}
return tx.Commit(ctx)
}
- if status != "paid" {
+ var alreadyCredited bool
+ if err := tx.QueryRow(ctx, `SELECT EXISTS(SELECT 1 FROM billing_ledger WHERE source_type='stripe_checkout' AND source_id=$1)`, session.ID).Scan(&alreadyCredited); err != nil {
+ return err
+ }
+ if !alreadyCredited {
if _, err := tx.Exec(ctx, `
INSERT INTO tenant_wallets (tenant_id, currency) VALUES ($1,$2)
ON CONFLICT (tenant_id) DO NOTHING`, tenantID, currency); err != nil {
@@ -209,14 +291,117 @@ func (s *Service) processStripeEvent(ctx context.Context, event stripe.Event) er
ON CONFLICT (source_type, source_id) DO NOTHING`, tenantID, currency, amountMicros, newBalance, session.ID); err != nil {
return err
}
- if _, err := tx.Exec(ctx, `
- UPDATE topup_orders SET status = 'paid', stripe_session_id = COALESCE(stripe_session_id, $2), paid_at = now()
- WHERE id = $1`, session.ClientReferenceID, session.ID); err != nil {
+ }
+ if _, err := tx.Exec(ctx, `
+ UPDATE topup_orders SET status=CASE WHEN status IN ('pending','failed','expired') THEN 'paid' ELSE status END,
+ stripe_session_id=COALESCE(stripe_session_id,$2), paid_at=COALESCE(paid_at,now())
+ WHERE id=$1`, session.ClientReferenceID, 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)
+}
+
+func (s *Service) processOperationalStripeEvent(ctx context.Context, event stripe.Event) error {
+ if event.ID == "" || event.Data == nil {
+ 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))`, event.ID); err != nil {
+ return err
+ }
+ tag, 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))
+ if err != nil {
+ return err
+ }
+ if tag.RowsAffected() == 0 {
+ var processed bool
+ if err := tx.QueryRow(ctx, `SELECT processed_at IS NOT NULL FROM stripe_webhook_events WHERE event_id=$1`, event.ID).Scan(&processed); err != nil {
return err
}
+ if processed {
+ return tx.Commit(ctx)
+ }
}
- if _, err := tx.Exec(ctx, `UPDATE stripe_webhook_events SET processed_at = now() WHERE event_id = $1`, event.ID); err != nil {
+ switch event.Type {
+ case stripe.EventTypeChargeSucceeded, stripe.EventTypeChargeUpdated, stripe.EventTypeChargeRefunded:
+ var charge stripe.Charge
+ if json.Unmarshal(event.Data.Raw, &charge) != nil || charge.ID == "" {
+ return ErrInvalidAmount
+ }
+ paymentIntentID := ""
+ if charge.PaymentIntent != nil {
+ paymentIntentID = charge.PaymentIntent.ID
+ }
+ if paymentIntentID != "" {
+ if _, err := tx.Exec(ctx, `UPDATE topup_orders SET stripe_charge_id=$2,receipt_url=COALESCE(NULLIF($3,''),receipt_url) WHERE stripe_payment_intent_id=$1`, paymentIntentID, charge.ID, charge.ReceiptURL); err != nil {
+ return err
+ }
+ }
+ if event.Type == stripe.EventTypeChargeRefunded && charge.Refunds != nil {
+ for _, refund := range charge.Refunds.Data {
+ if err := s.applyRefundTx(ctx, tx, refund); err != nil {
+ return err
+ }
+ }
+ }
+ case stripe.EventTypeRefundCreated, stripe.EventTypeRefundUpdated, stripe.EventTypeRefundFailed:
+ var refund stripe.Refund
+ if json.Unmarshal(event.Data.Raw, &refund) != nil || refund.ID == "" {
+ return ErrInvalidAmount
+ }
+ if err := s.applyRefundTx(ctx, tx, &refund); err != nil {
+ return err
+ }
+ case stripe.EventTypeChargeDisputeCreated, stripe.EventTypeChargeDisputeUpdated,
+ stripe.EventTypeChargeDisputeClosed, stripe.EventTypeChargeDisputeFundsWithdrawn, stripe.EventTypeChargeDisputeFundsReinstated:
+ var dispute stripe.Dispute
+ if json.Unmarshal(event.Data.Raw, &dispute) != nil || dispute.ID == "" {
+ return ErrInvalidAmount
+ }
+ if err := s.applyDisputeTx(ctx, tx, &dispute, event.Type); err != nil {
+ return err
+ }
+ case stripe.EventTypeInvoiceCreated, stripe.EventTypeInvoiceFinalized, stripe.EventTypeInvoicePaid, stripe.EventTypeInvoicePaymentFailed:
+ var invoice stripe.Invoice
+ if json.Unmarshal(event.Data.Raw, &invoice) != nil || invoice.ID == "" {
+ return ErrInvalidAmount
+ }
+ if err := s.applyInvoiceTx(ctx, tx, &invoice); 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)
}
+
+func (s *Service) recordWebhookAttempt(ctx context.Context, event stripe.Event) error {
+ if event.ID == "" {
+ return ErrInvalidAmount
+ }
+ _, err := s.db.Exec(ctx, `INSERT INTO stripe_webhook_events
+ (event_id,event_type,attempts,last_attempt_at) VALUES ($1,$2,1,now())
+ ON CONFLICT (event_id) DO UPDATE SET attempts=stripe_webhook_events.attempts+1,last_attempt_at=now()`,
+ event.ID, string(event.Type))
+ return err
+}
+
+func (s *Service) recordWebhookFailure(ctx context.Context, eventID string, cause error) {
+ message := "webhook processing failed"
+ if cause != nil {
+ message = cause.Error()
+ }
+ if len(message) > 1000 {
+ message = message[:1000]
+ }
+ _, _ = s.db.Exec(ctx, `UPDATE stripe_webhook_events SET processing_error=$2,last_attempt_at=now()
+ WHERE event_id=$1 AND processed_at IS NULL`, eventID, message)
+}