diff options
Diffstat (limited to '')
| -rw-r--r-- | internal/billing/stripe.go | 207 |
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) +} |
