package billing import ( "context" "encoding/json" "errors" "fmt" "io" "net/http" "net/url" "strings" "github.com/jackc/pgx/v5" "github.com/stripe/stripe-go/v86" "github.com/stripe/stripe-go/v86/webhook" ) const maxWebhookBodyBytes = 1 << 20 type stripeCheckoutCreator func(context.Context, *stripe.CheckoutSessionCreateParams) (*stripe.CheckoutSession, error) func newStripeCheckoutCreator(apiKey string) stripeCheckoutCreator { client := stripe.NewClient(apiKey) return client.V1CheckoutSessions.Create } func (s *Service) CreateCheckout(ctx context.Context, input CheckoutInput) (CheckoutResult, error) { orderID, _, err := s.createTopUpOrder(ctx, input) if err != nil { return CheckoutResult{}, err } params := s.checkoutSessionParams(orderID, input) customerID, err := s.ensureStripeCustomer(ctx, input.TenantID) if err != nil { return CheckoutResult{}, s.failCheckoutCreation(ctx, orderID, err) } if customerID != "" { params.Customer = stripe.String(customerID) } else { params.CustomerCreation = stripe.String(string(stripe.CheckoutSessionCustomerCreationAlways)) if strings.TrimSpace(input.CustomerEmail) != "" { params.CustomerEmail = stripe.String(strings.TrimSpace(input.CustomerEmail)) } } session, err := s.createStripeCheckout(ctx, params) if err != nil { return CheckoutResult{}, s.failCheckoutCreation(ctx, orderID, fmt.Errorf("create Stripe Checkout Session: %w", err)) } if session == nil || session.ID == "" || session.URL == "" { return CheckoutResult{}, s.failCheckoutCreation(ctx, orderID, errors.New("Stripe returned an incomplete Checkout Session")) } if _, err := s.db.Exec(ctx, ` UPDATE topup_orders SET stripe_session_id = $2, checkout_url = $3 WHERE id = $1`, orderID, session.ID, session.URL); err != nil { return CheckoutResult{}, s.failCheckoutCreation(ctx, orderID, fmt.Errorf("persist Stripe Checkout Session: %w", err)) } return CheckoutResult{OrderID: orderID, SessionID: session.ID, URL: session.URL}, nil } func (s *Service) checkoutSessionParams(orderID string, input CheckoutInput) *stripe.CheckoutSessionCreateParams { params := &stripe.CheckoutSessionCreateParams{ Mode: stripe.String("payment"), ClientReferenceID: stripe.String(orderID), IntegrationIdentifier: stripe.String(s.integrationIdentifier), SuccessURL: stripe.String(checkoutReturnURL(s.stripeSuccessURL, orderID, true)), CancelURL: stripe.String(checkoutReturnURL(s.stripeCancelURL, orderID, false)), Metadata: map[string]string{ "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{ Currency: stripe.String(s.currency), UnitAmount: stripe.Int64(input.AmountMinor), ProductData: &stripe.CheckoutSessionCreateLineItemPriceDataProductDataParams{ Name: stripe.String("AIGW prepaid balance"), Description: stripe.String("Prepaid API usage credit"), }, }, }}, } 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) return params } func (s *Service) failCheckoutCreation(ctx context.Context, orderID string, cause error) error { _, updateErr := s.db.Exec(ctx, `UPDATE topup_orders SET status='failed', reconciliation_status='unknown', reconciliation_error=$2 WHERE id=$1 AND status='pending'`, orderID, truncateError(cause)) if updateErr != nil { return errors.Join(cause, fmt.Errorf("mark top-up order failed: %w", updateErr)) } return cause } func checkoutReturnURL(raw, orderID string, includeStripeSession bool) string { parsed, err := url.Parse(raw) if err != nil { return raw } query := parsed.Query() query.Set("order_id", orderID) if includeStripeSession { query.Set("session_id", "{CHECKOUT_SESSION_ID}") } else { query.Del("session_id") } encoded := query.Encode() encoded = strings.ReplaceAll(encoded, url.QueryEscape("{CHECKOUT_SESSION_ID}"), "{CHECKOUT_SESSION_ID}") parsed.RawQuery = encoded return parsed.String() } func (s *Service) WebhookHandler() http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if r.Method != http.MethodPost { w.Header().Set("Allow", http.MethodPost) http.Error(w, "method not allowed", http.StatusMethodNotAllowed) return } body, err := io.ReadAll(io.LimitReader(r.Body, maxWebhookBodyBytes+1)) if err != nil || len(body) > maxWebhookBodyBytes { http.Error(w, "invalid webhook body", http.StatusBadRequest) return } event, err := webhook.ConstructEvent(body, r.Header.Get("Stripe-Signature"), s.stripeWebhookSecret) if err != nil { 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 } http.Error(w, "webhook processing failed", http.StatusInternalServerError) return } w.Header().Set("Content-Type", "application/json") _, _ = io.WriteString(w, `{"received":true}`+"\n") }) } func (s *Service) processStripeEvent(ctx context.Context, event stripe.Event) error { 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.EventTypePaymentIntentSucceeded, stripe.EventTypePaymentIntentPaymentFailed, stripe.EventTypePaymentIntentCanceled, stripe.EventTypeChargeDisputeCreated, stripe.EventTypeChargeDisputeUpdated, stripe.EventTypeChargeDisputeClosed, stripe.EventTypeChargeDisputeFundsWithdrawn, stripe.EventTypeChargeDisputeFundsReinstated, stripe.EventTypeInvoiceCreated, stripe.EventTypeInvoiceFinalized, stripe.EventTypeInvoicePaid, stripe.EventTypeInvoicePaymentFailed: return s.processOperationalStripeEvent(ctx, event) default: _, 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 } var session stripe.CheckoutSession 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 } 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 (event_id) DO NOTHING`, event.ID, typeName) if err != nil { return fmt.Errorf("record Stripe event: %w", 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) } } var tenantID, currency, status string var amountMinor, amountMicros int64 var storedSessionID *string err = tx.QueryRow(ctx, ` SELECT tenant_id::text, amount_minor, amount_micros, currency, status, stripe_session_id FROM topup_orders WHERE id = $1 FOR UPDATE`, session.ClientReferenceID, ).Scan(&tenantID, &amountMinor, &amountMicros, ¤cy, &status, &storedSessionID) if err != nil { return err } 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 { orderStatus = "expired" } 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(),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(),processing_error='' WHERE event_id=$1`, event.ID) if err != nil { return err } return tx.Commit(ctx) } 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 { return err } var walletCurrency string var balance int64 if err := tx.QueryRow(ctx, `SELECT currency, balance_micros FROM tenant_wallets WHERE tenant_id = $1 FOR UPDATE`, tenantID).Scan(&walletCurrency, &balance); err != nil { return err } if walletCurrency != currency || amountMicros <= 0 || balance > int64(^uint64(0)>>1)-amountMicros { return ErrInvalidAmount } newBalance := balance + amountMicros 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, ` 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_checkout',$5,'Stripe balance top-up') 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=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) } } 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 == "" { 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) }