summaryrefslogtreecommitdiff
path: root/internal/billing
diff options
context:
space:
mode:
Diffstat (limited to '')
-rw-r--r--internal/billing/ledger.go41
-rw-r--r--internal/billing/operations.go881
-rw-r--r--internal/billing/operations_test.go36
-rw-r--r--internal/billing/service.go319
-rw-r--r--internal/billing/service_test.go174
-rw-r--r--internal/billing/stripe.go207
-rw-r--r--internal/billing/types.go162
7 files changed, 1774 insertions, 46 deletions
diff --git a/internal/billing/ledger.go b/internal/billing/ledger.go
index c408cfe..2eb3d87 100644
--- a/internal/billing/ledger.go
+++ b/internal/billing/ledger.go
@@ -77,9 +77,20 @@ func (s *Service) ListTopUpOrders(ctx context.Context, tenantID string, limit in
if limit < 1 || limit > 200 {
limit = 50
}
- rows, err := s.db.Query(ctx, `SELECT id::text,tenant_id::text,amount_minor,amount_micros,currency,status,
- COALESCE(stripe_session_id,''),COALESCE(checkout_url,''),created_at,paid_at FROM topup_orders
- WHERE tenant_id=$1 ORDER BY created_at DESC LIMIT $2`, tenantID, limit)
+ query := `SELECT id::text,tenant_id::text,amount_minor,amount_micros,currency,status,
+ 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,''),
+ refunded_micros,disputed_micros,reconciliation_status,reconciled_at,reconciliation_error FROM topup_orders`
+ args := []any{}
+ if strings.TrimSpace(tenantID) != "" {
+ query += ` WHERE tenant_id=$1 ORDER BY created_at DESC LIMIT $2`
+ args = []any{tenantID, limit}
+ } else {
+ query += ` ORDER BY created_at DESC LIMIT $1`
+ args = []any{limit}
+ }
+ rows, err := s.db.Query(ctx, query, args...)
if err != nil {
return nil, fmt.Errorf("query top-up orders: %w", err)
}
@@ -88,7 +99,10 @@ func (s *Service) ListTopUpOrders(ctx context.Context, tenantID string, limit in
for rows.Next() {
var item TopUpOrder
if err := rows.Scan(&item.ID, &item.TenantID, &item.AmountMinor, &item.AmountMicros, &item.Currency, &item.Status,
- &item.StripeSessionID, &item.CheckoutURL, &item.CreatedAt, &item.PaidAt); err != nil {
+ &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,
+ &item.ReconciliationStatus, &item.ReconciledAt, &item.ReconciliationError); err != nil {
return nil, err
}
result = append(result, item)
@@ -98,10 +112,21 @@ 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
- err := s.db.QueryRow(ctx, `SELECT id::text,tenant_id::text,amount_minor,amount_micros,currency,status,
- COALESCE(stripe_session_id,''),COALESCE(checkout_url,''),created_at,paid_at FROM topup_orders
- WHERE id=$1 AND tenant_id=$2`, orderID, tenantID).Scan(&result.ID, &result.TenantID, &result.AmountMinor, &result.AmountMicros,
- &result.Currency, &result.Status, &result.StripeSessionID, &result.CheckoutURL, &result.CreatedAt, &result.PaidAt)
+ query := `SELECT id::text,tenant_id::text,amount_minor,amount_micros,currency,status,
+ 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,''),
+ refunded_micros,disputed_micros,reconciliation_status,reconciled_at,reconciliation_error FROM topup_orders WHERE id=$1`
+ args := []any{orderID}
+ if strings.TrimSpace(tenantID) != "" {
+ query += ` AND tenant_id=$2`
+ 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.StripeCustomerID, &result.StripePaymentIntentID, &result.StripeChargeID, &result.StripeInvoiceID,
+ &result.InvoiceURL, &result.InvoicePDFURL, &result.ReceiptURL, &result.RefundedMicros, &result.DisputedMicros,
+ &result.ReconciliationStatus, &result.ReconciledAt, &result.ReconciliationError)
if errors.Is(err, pgx.ErrNoRows) {
return TopUpOrder{}, ErrTopUpOrderNotFound
}
diff --git a/internal/billing/operations.go b/internal/billing/operations.go
new file mode 100644
index 0000000..a6dc653
--- /dev/null
+++ b/internal/billing/operations.go
@@ -0,0 +1,881 @@
+package billing
+
+import (
+ "context"
+ "encoding/csv"
+ "encoding/json"
+ "errors"
+ "fmt"
+ "io"
+ "strconv"
+ "strings"
+ "time"
+
+ "github.com/jackc/pgx/v5"
+ "github.com/stripe/stripe-go/v86"
+)
+
+func (s *Service) CreatePortalSession(ctx context.Context, tenantID string) (PortalResult, error) {
+ 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")
+ }
+ return PortalResult{}, err
+ }
+ session, err := s.stripeClient.V1BillingPortalSessions.Create(ctx, &stripe.BillingPortalSessionCreateParams{
+ Customer: stripe.String(customerID), ReturnURL: stripe.String(s.stripePortalReturnURL),
+ })
+ if err != nil {
+ return PortalResult{}, fmt.Errorf("create Stripe customer portal session: %w", err)
+ }
+ if session.URL == "" {
+ return PortalResult{}, errors.New("Stripe returned an incomplete portal session")
+ }
+ return PortalResult{URL: session.URL}, nil
+}
+
+func (s *Service) RetryCheckout(ctx context.Context, tenantID, orderID, email string) (CheckoutResult, error) {
+ var amount int64
+ var status string
+ if err := s.db.QueryRow(ctx, `SELECT amount_minor,status FROM topup_orders WHERE id=$1 AND tenant_id=$2`, orderID, tenantID).Scan(&amount, &status); err != nil {
+ if errors.Is(err, pgx.ErrNoRows) {
+ return CheckoutResult{}, ErrTopUpOrderNotFound
+ }
+ return CheckoutResult{}, err
+ }
+ if status != "failed" && status != "expired" {
+ return CheckoutResult{}, errors.New("only failed or expired top-ups can be retried")
+ }
+ return s.CreateCheckout(ctx, CheckoutInput{TenantID: tenantID, AmountMinor: amount, CustomerEmail: email})
+}
+
+// ResolveMissingTopUp closes an uncredited local order only after reconciliation
+// proved that its Checkout Session does not exist in the configured Stripe account.
+func (s *Service) ResolveMissingTopUp(ctx context.Context, tenantID, orderID string, input ResolveMissingTopUpInput, actor ResolutionActor) (TopUpOrder, error) {
+ reason := normalizeDescription(input.Reason)
+ if reason == "" {
+ return TopUpOrder{}, fmt.Errorf("%w: a resolution reason is required", ErrCannotResolveTopUp)
+ }
+ tx, err := s.db.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.ReadCommitted})
+ if err != nil {
+ return TopUpOrder{}, err
+ }
+ defer tx.Rollback(ctx)
+ var status, reconciliationStatus, paymentIntentID string
+ if err := tx.QueryRow(ctx, `SELECT status,reconciliation_status,COALESCE(stripe_payment_intent_id,'')
+ FROM topup_orders WHERE id=$1 AND tenant_id=$2 FOR UPDATE`, orderID, tenantID).
+ Scan(&status, &reconciliationStatus, &paymentIntentID); err != nil {
+ if errors.Is(err, pgx.ErrNoRows) {
+ return TopUpOrder{}, ErrTopUpOrderNotFound
+ }
+ return TopUpOrder{}, err
+ }
+ if status != "pending" || reconciliationStatus != "missing" || paymentIntentID != "" {
+ return TopUpOrder{}, fmt.Errorf("%w: only an uncredited pending order confirmed missing by reconciliation can be resolved", ErrCannotResolveTopUp)
+ }
+ var credits int64
+ if err := tx.QueryRow(ctx, `SELECT count(*) FROM billing_ledger
+ WHERE source_type='stripe_checkout' AND source_id=(SELECT stripe_session_id FROM topup_orders WHERE id=$1)`, orderID).Scan(&credits); err != nil {
+ return TopUpOrder{}, err
+ }
+ if credits != 0 {
+ return TopUpOrder{}, fmt.Errorf("%w: a credited top-up cannot be resolved as missing", ErrCannotResolveTopUp)
+ }
+ if actor.Type != "console_user" && actor.Type != "bootstrap" && actor.Type != "maintenance" {
+ return TopUpOrder{}, fmt.Errorf("%w: a valid resolution actor is required", ErrCannotResolveTopUp)
+ }
+ if _, err := tx.Exec(ctx, `INSERT INTO billing_reconciliation_resolutions
+ (topup_order_id,tenant_id,actor_id,actor_type,reason) VALUES ($1,$2,$3,$4,$5)`,
+ orderID, tenantID, actor.ID, actor.Type, reason); err != nil {
+ return TopUpOrder{}, err
+ }
+ if _, err := tx.Exec(ctx, `UPDATE topup_orders SET status='failed',reconciliation_status='resolved',
+ reconciled_at=now(),reconciliation_error=$2 WHERE id=$1`, orderID, reason); err != nil {
+ return TopUpOrder{}, err
+ }
+ if err := tx.Commit(ctx); err != nil {
+ return TopUpOrder{}, err
+ }
+ return s.GetTopUpOrder(ctx, tenantID, orderID)
+}
+
+// ReverseMissingTopUpCredit preserves the original credit and adds an equal
+// negative ledger entry when the configured Stripe account cannot prove the
+// payment. It refuses to consume funds reserved for in-flight requests.
+func (s *Service) ReverseMissingTopUpCredit(ctx context.Context, tenantID, orderID string, input ResolveMissingTopUpInput, actor ResolutionActor) (TopUpOrder, error) {
+ reason := normalizeDescription(input.Reason)
+ if reason == "" || (actor.Type != "console_user" && actor.Type != "bootstrap" && actor.Type != "maintenance") {
+ return TopUpOrder{}, fmt.Errorf("%w: a reason and valid resolution actor are required", ErrCannotResolveTopUp)
+ }
+ tx, err := s.db.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.ReadCommitted})
+ if err != nil {
+ return TopUpOrder{}, err
+ }
+ defer tx.Rollback(ctx)
+ var status, reconciliationStatus, paymentIntentID, currency, sessionID string
+ var amount int64
+ if err := tx.QueryRow(ctx, `SELECT status,reconciliation_status,COALESCE(stripe_payment_intent_id,''),
+ currency,amount_micros,COALESCE(stripe_session_id,'') FROM topup_orders
+ WHERE id=$1 AND tenant_id=$2 FOR UPDATE`, orderID, tenantID).
+ Scan(&status, &reconciliationStatus, &paymentIntentID, &currency, &amount, &sessionID); err != nil {
+ if errors.Is(err, pgx.ErrNoRows) {
+ return TopUpOrder{}, ErrTopUpOrderNotFound
+ }
+ return TopUpOrder{}, err
+ }
+ if status != "paid" || reconciliationStatus != "missing" || paymentIntentID != "" || sessionID == "" {
+ return TopUpOrder{}, fmt.Errorf("%w: only a paid credit confirmed missing with no PaymentIntent can be reversed", ErrCannotResolveTopUp)
+ }
+ var originalCredits int64
+ if err := tx.QueryRow(ctx, `SELECT COALESCE(sum(amount_micros),0) FROM billing_ledger
+ WHERE tenant_id=$1 AND source_type='stripe_checkout' AND source_id=$2`, tenantID, sessionID).Scan(&originalCredits); err != nil {
+ return TopUpOrder{}, err
+ }
+ if originalCredits != amount {
+ return TopUpOrder{}, fmt.Errorf("%w: original Stripe credit does not match the order", ErrCannotResolveTopUp)
+ }
+ 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 TopUpOrder{}, err
+ }
+ if balance-reserved < amount {
+ return TopUpOrder{}, fmt.Errorf("%w: available balance is insufficient to reverse the orphaned credit", ErrCannotResolveTopUp)
+ }
+ newBalance := balance - amount
+ if _, err := tx.Exec(ctx, `UPDATE tenant_wallets SET balance_micros=$2,updated_at=now() WHERE tenant_id=$1`, tenantID, newBalance); err != nil {
+ return TopUpOrder{}, 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,'adjustment','stripe_reconciliation',$5,$6)`,
+ tenantID, currency, -amount, newBalance, orderID, reason); err != nil {
+ return TopUpOrder{}, err
+ }
+ if _, err := tx.Exec(ctx, `INSERT INTO billing_reconciliation_resolutions
+ (topup_order_id,tenant_id,actor_id,actor_type,reason) VALUES ($1,$2,$3,$4,$5)`,
+ orderID, tenantID, actor.ID, actor.Type, reason); err != nil {
+ return TopUpOrder{}, err
+ }
+ if _, err := tx.Exec(ctx, `UPDATE topup_orders SET status='reversed',reconciliation_status='resolved',
+ reconciled_at=now(),reconciliation_error=$2 WHERE id=$1`, orderID, reason); err != nil {
+ return TopUpOrder{}, err
+ }
+ if err := tx.Commit(ctx); err != nil {
+ return TopUpOrder{}, err
+ }
+ return s.GetTopUpOrder(ctx, tenantID, orderID)
+}
+
+func (s *Service) CreateRefund(ctx context.Context, tenantID, orderID string, input RefundInput) (Refund, error) {
+ if !s.stripeEnabled {
+ return Refund{}, ErrStripeDisabled
+ }
+ input.Reason = strings.TrimSpace(input.Reason)
+ if input.Reason == "" {
+ input.Reason = "requested_by_customer"
+ }
+ if input.Reason != "requested_by_customer" && input.Reason != "duplicate" && input.Reason != "fraudulent" {
+ return Refund{}, errors.New("refund reason must be requested_by_customer, duplicate, or fraudulent")
+ }
+ tx, err := s.db.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.ReadCommitted})
+ if err != nil {
+ return Refund{}, err
+ }
+ defer tx.Rollback(ctx)
+ var orderTenant, currency, status, paymentIntentID string
+ var orderAmount, refunded int64
+ if err := tx.QueryRow(ctx, `SELECT tenant_id::text,currency,status,COALESCE(stripe_payment_intent_id,''),amount_micros,refunded_micros
+ FROM topup_orders WHERE id=$1 AND tenant_id=$2 FOR UPDATE`, orderID, tenantID).Scan(&orderTenant, &currency, &status, &paymentIntentID, &orderAmount, &refunded); err != nil {
+ if errors.Is(err, pgx.ErrNoRows) {
+ return Refund{}, ErrTopUpOrderNotFound
+ }
+ return Refund{}, err
+ }
+ if status != "paid" && status != "partially_refunded" {
+ return Refund{}, errors.New("only a paid top-up can be refunded")
+ }
+ if paymentIntentID == "" {
+ return Refund{}, errors.New("top-up has no Stripe PaymentIntent")
+ }
+ amountMicros, err := minorToMicros(currency, input.AmountMinor)
+ if err != nil {
+ return Refund{}, ErrInvalidAmount
+ }
+ var pendingRefunds int64
+ if err := tx.QueryRow(ctx, `SELECT COALESCE(sum(amount_micros),0) FROM stripe_refunds
+ WHERE topup_order_id=$1 AND status IN ('queued','submitting','pending','requires_action','succeeded')`, orderID).Scan(&pendingRefunds); err != nil {
+ return Refund{}, err
+ }
+ // Legacy successful refunds are already reflected in refunded_micros. New
+ // rows are included in the sum, so use the larger value without double-counting.
+ committedRefunds := refunded
+ if pendingRefunds > committedRefunds {
+ committedRefunds = pendingRefunds
+ }
+ if amountMicros > orderAmount-committedRefunds {
+ return Refund{}, ErrInvalidAmount
+ }
+ var balance, held int64
+ if err := tx.QueryRow(ctx, `SELECT balance_micros,reserved_micros FROM tenant_wallets WHERE tenant_id=$1 FOR UPDATE`, tenantID).Scan(&balance, &held); err != nil {
+ return Refund{}, err
+ }
+ available := balance - held
+ if available < 0 {
+ available = 0
+ }
+ refundHold := amountMicros
+ if refundHold > available {
+ refundHold = available
+ }
+ if _, err := tx.Exec(ctx, `UPDATE tenant_wallets SET reserved_micros=reserved_micros+$2,updated_at=now() WHERE tenant_id=$1`, tenantID, refundHold); err != nil {
+ return Refund{}, err
+ }
+ var result Refund
+ err = tx.QueryRow(ctx, `INSERT INTO stripe_refunds (tenant_id,topup_order_id,amount_minor,amount_micros,held_micros,currency,reason)
+ VALUES ($1,$2,$3,$4,$5,$6,$7) RETURNING id::text,tenant_id::text,topup_order_id::text,amount_minor,amount_micros,currency,reason,status,last_error,created_at,completed_at`,
+ tenantID, orderID, input.AmountMinor, amountMicros, refundHold, currency, input.Reason).Scan(&result.ID, &result.TenantID, &result.TopUpOrderID,
+ &result.AmountMinor, &result.AmountMicros, &result.Currency, &result.Reason, &result.Status, &result.LastError, &result.CreatedAt, &result.CompletedAt)
+ if err != nil {
+ return Refund{}, err
+ }
+ if err := tx.Commit(ctx); err != nil {
+ return Refund{}, err
+ }
+ return result, nil
+}
+
+func (s *Service) ListRefunds(ctx context.Context, tenantID string, limit int) ([]Refund, error) {
+ if limit < 1 || limit > 500 {
+ limit = 100
+ }
+ query := `SELECT id::text,tenant_id::text,topup_order_id::text,COALESCE(stripe_refund_id,''),amount_minor,amount_micros,
+ currency,reason,status,last_error,created_at,completed_at FROM stripe_refunds`
+ args := []any{}
+ if tenantID != "" {
+ query += ` WHERE tenant_id=$1 ORDER BY created_at DESC LIMIT $2`
+ args = []any{tenantID, limit}
+ } else {
+ query += ` ORDER BY created_at DESC LIMIT $1`
+ args = []any{limit}
+ }
+ rows, err := s.db.Query(ctx, query, args...)
+ if err != nil {
+ return nil, err
+ }
+ defer rows.Close()
+ result := make([]Refund, 0)
+ for rows.Next() {
+ var item Refund
+ if err := rows.Scan(&item.ID, &item.TenantID, &item.TopUpOrderID, &item.StripeRefundID, &item.AmountMinor, &item.AmountMicros, &item.Currency, &item.Reason, &item.Status, &item.LastError, &item.CreatedAt, &item.CompletedAt); err != nil {
+ return nil, err
+ }
+ result = append(result, item)
+ }
+ return result, rows.Err()
+}
+
+func (s *Service) ListDisputes(ctx context.Context, tenantID string, limit int) ([]PaymentDispute, error) {
+ if limit < 1 || limit > 500 {
+ limit = 100
+ }
+ query := `SELECT stripe_dispute_id,COALESCE(tenant_id::text,''),COALESCE(topup_order_id::text,''),
+ amount_minor,amount_micros,currency,status,reason,debited_micros,uncollected_micros,due_by,updated_at FROM stripe_disputes`
+ args := []any{}
+ if tenantID != "" {
+ query += ` WHERE tenant_id=$1 ORDER BY updated_at DESC LIMIT $2`
+ args = []any{tenantID, limit}
+ } else {
+ query += ` ORDER BY updated_at DESC LIMIT $1`
+ args = []any{limit}
+ }
+ rows, err := s.db.Query(ctx, query, args...)
+ if err != nil {
+ return nil, err
+ }
+ defer rows.Close()
+ result := []PaymentDispute{}
+ for rows.Next() {
+ var item PaymentDispute
+ if err := rows.Scan(&item.ID, &item.TenantID, &item.TopUpOrderID, &item.AmountMinor, &item.AmountMicros, &item.Currency, &item.Status, &item.Reason, &item.DebitedMicros, &item.UncollectedMicros, &item.DueBy, &item.UpdatedAt); err != nil {
+ return nil, err
+ }
+ result = append(result, item)
+ }
+ return result, rows.Err()
+}
+
+func (s *Service) ListInvoices(ctx context.Context, tenantID string, limit int) ([]Invoice, error) {
+ if limit < 1 || limit > 500 {
+ limit = 100
+ }
+ query := `SELECT stripe_invoice_id,COALESCE(tenant_id::text,''),COALESCE(topup_order_id::text,''),status,currency,
+ amount_due_minor,amount_paid_minor,attempt_count,next_payment_attempt,hosted_invoice_url,invoice_pdf_url,last_failure,updated_at FROM stripe_invoices`
+ args := []any{}
+ if tenantID != "" {
+ query += ` WHERE tenant_id=$1 ORDER BY updated_at DESC LIMIT $2`
+ args = []any{tenantID, limit}
+ } else {
+ query += ` ORDER BY updated_at DESC LIMIT $1`
+ args = []any{limit}
+ }
+ rows, err := s.db.Query(ctx, query, args...)
+ if err != nil {
+ return nil, err
+ }
+ defer rows.Close()
+ result := []Invoice{}
+ for rows.Next() {
+ var item Invoice
+ if err := rows.Scan(&item.ID, &item.TenantID, &item.TopUpOrderID, &item.Status, &item.Currency, &item.AmountDueMinor, &item.AmountPaidMinor, &item.AttemptCount, &item.NextPaymentAttempt, &item.HostedInvoiceURL, &item.InvoicePDFURL, &item.LastFailure, &item.UpdatedAt); err != nil {
+ return nil, err
+ }
+ result = append(result, item)
+ }
+ return result, rows.Err()
+}
+
+func (s *Service) RunStripeOperations(ctx context.Context) {
+ if !s.stripeEnabled || s.stripeClient == nil {
+ return
+ }
+ ticker := time.NewTicker(2 * time.Second)
+ reconcile := time.NewTicker(30 * time.Minute)
+ metrics := time.NewTicker(15 * time.Second)
+ defer ticker.Stop()
+ defer reconcile.Stop()
+ defer metrics.Stop()
+ initialCtx, cancel := context.WithTimeout(ctx, 2*time.Minute)
+ _, _ = s.Reconcile(initialCtx, 100)
+ cancel()
+ s.refreshOperationalMetrics(ctx)
+ for {
+ for i := 0; i < 8; i++ {
+ ok, _ := s.processRefundOperation(ctx)
+ if !ok {
+ break
+ }
+ }
+ _, _ = s.pollPendingRefund(ctx)
+ select {
+ case <-ctx.Done():
+ return
+ case <-ticker.C:
+ case <-reconcile.C:
+ _, _ = s.Reconcile(ctx, 100)
+ case <-metrics.C:
+ s.refreshOperationalMetrics(ctx)
+ }
+ }
+}
+
+func (s *Service) refreshOperationalMetrics(ctx context.Context) {
+ if s.metrics == nil {
+ return
+ }
+ status, err := s.OperationalStatus(ctx)
+ if err == nil {
+ s.metrics.SetStripeOperations(status)
+ }
+}
+
+func (s *Service) OperationalStatus(ctx context.Context) (OperationalStatus, error) {
+ status := OperationalStatus{StripeEnabled: s.stripeEnabled, ReconciliationStatus: "disabled"}
+ if err := s.db.QueryRow(ctx, `SELECT
+ count(*) FILTER (WHERE status IN ('queued','submitting','pending','requires_action')),
+ min(created_at) FILTER (WHERE status IN ('queued','submitting','pending','requires_action')),
+ COALESCE((SELECT sum(uncollected_micros) FROM stripe_refunds),0)+
+ COALESCE((SELECT sum(uncollected_micros) FROM stripe_disputes),0)+
+ COALESCE((SELECT sum(uncollected_micros) FROM usage_events),0)
+ FROM stripe_refunds`).Scan(&status.RefundBacklog, &status.OldestRefund, &status.UncollectedMicros); err != nil {
+ return status, err
+ }
+ if err := s.db.QueryRow(ctx, `SELECT count(*),min(created_at) FROM stripe_webhook_events
+ WHERE processed_at IS NULL`).Scan(&status.UnprocessedWebhooks, &status.OldestUnprocessedWebhook); err != nil {
+ return status, err
+ }
+ if err := s.db.QueryRow(ctx, `SELECT count(*) FROM usage_events WHERE metering_status='missing'`).Scan(&status.UnmeteredSuccesses); err != nil {
+ return status, err
+ }
+ if !s.stripeEnabled {
+ return status, nil
+ }
+ err := s.db.QueryRow(ctx, `SELECT status,mismatch_count,completed_at,error FROM billing_reconciliation_runs
+ WHERE status<>'running' ORDER BY started_at DESC LIMIT 1`).Scan(&status.ReconciliationStatus,
+ &status.ReconciliationMismatches, &status.ReconciliationCompletedAt, &status.ReconciliationError)
+ if errors.Is(err, pgx.ErrNoRows) {
+ status.ReconciliationStatus = "never_run"
+ return status, nil
+ }
+ return status, err
+}
+
+func (s *Service) processRefundOperation(ctx context.Context) (bool, error) {
+ tx, err := s.db.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.ReadCommitted})
+ if err != nil {
+ return false, err
+ }
+ defer tx.Rollback(ctx)
+ var id, orderID, paymentIntentID, reason string
+ var amount int64
+ err = tx.QueryRow(ctx, `WITH selected AS (
+ SELECT r.id FROM stripe_refunds r WHERE r.available_at<=now() AND
+ (r.status='queued' OR (r.status='submitting' AND r.updated_at<now()-interval '5 minutes'))
+ ORDER BY r.available_at,r.created_at FOR UPDATE SKIP LOCKED LIMIT 1)
+ UPDATE stripe_refunds r SET status='submitting',attempts=attempts+1,updated_at=now()
+ FROM selected,topup_orders o WHERE r.id=selected.id AND o.id=r.topup_order_id
+ RETURNING r.id::text,r.topup_order_id::text,o.stripe_payment_intent_id,r.amount_minor,r.reason`).Scan(&id, &orderID, &paymentIntentID, &amount, &reason)
+ if errors.Is(err, pgx.ErrNoRows) {
+ return false, tx.Commit(ctx)
+ }
+ if err != nil {
+ return false, err
+ }
+ if err := tx.Commit(ctx); err != nil {
+ return false, err
+ }
+ params := &stripe.RefundCreateParams{Amount: stripe.Int64(amount), PaymentIntent: stripe.String(paymentIntentID), Reason: stripe.String(reason),
+ Metadata: map[string]string{"aigw_refund_id": id, "aigw_topup_order_id": orderID}}
+ params.SetIdempotencyKey("aigw_refund_" + id)
+ refund, err := s.stripeClient.V1Refunds.Create(ctx, params)
+ if err != nil {
+ message := err.Error()
+ if len(message) > 1000 {
+ message = message[:1000]
+ }
+ _, _ = s.db.Exec(ctx, `UPDATE stripe_refunds SET status='queued',last_error=$2,
+ available_at=now()+make_interval(secs=>LEAST(1800,power(2,LEAST(attempts,10))::int)),updated_at=now() WHERE id=$1`, id, message)
+ return true, err
+ }
+ return true, s.applyRefund(ctx, refund)
+}
+
+func (s *Service) pollPendingRefund(ctx context.Context) (bool, error) {
+ var id string
+ err := s.db.QueryRow(ctx, `UPDATE stripe_refunds SET available_at=now()+interval '1 minute',updated_at=now()
+ WHERE id=(SELECT id FROM stripe_refunds WHERE status IN ('pending','requires_action')
+ AND stripe_refund_id IS NOT NULL AND available_at<=now() ORDER BY available_at LIMIT 1 FOR UPDATE SKIP LOCKED)
+ RETURNING stripe_refund_id`).Scan(&id)
+ if errors.Is(err, pgx.ErrNoRows) {
+ return false, nil
+ }
+ if err != nil {
+ return false, err
+ }
+ refund, err := s.stripeClient.V1Refunds.Retrieve(ctx, id, &stripe.RefundRetrieveParams{})
+ if err != nil {
+ return true, err
+ }
+ return true, s.applyRefund(ctx, refund)
+}
+
+func (s *Service) applyRefund(ctx context.Context, refund *stripe.Refund) error {
+ if refund == nil || refund.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 := s.applyRefundTx(ctx, tx, refund); err != nil {
+ return err
+ }
+ return tx.Commit(ctx)
+}
+
+func (s *Service) applyRefundTx(ctx context.Context, tx pgx.Tx, refund *stripe.Refund) error {
+ if refund == nil || refund.ID == "" {
+ return ErrInvalidAmount
+ }
+ localID := refund.Metadata["aigw_refund_id"]
+ var id, tenantID, orderID, currentStatus string
+ var amountMicros, held int64
+ query := `SELECT id::text,tenant_id::text,topup_order_id::text,amount_micros,held_micros,status FROM stripe_refunds WHERE `
+ arg := refund.ID
+ if localID != "" {
+ query += `id=$1 FOR UPDATE`
+ arg = localID
+ } else {
+ query += `stripe_refund_id=$1 FOR UPDATE`
+ }
+ err := tx.QueryRow(ctx, query, arg).Scan(&id, &tenantID, &orderID, &amountMicros, &held, &currentStatus)
+ if errors.Is(err, pgx.ErrNoRows) {
+ paymentIntentID := ""
+ if refund.PaymentIntent != nil {
+ paymentIntentID = refund.PaymentIntent.ID
+ }
+ if paymentIntentID == "" {
+ return ErrInvalidAmount
+ }
+ var currency string
+ if err := tx.QueryRow(ctx, `SELECT id::text,tenant_id::text,currency FROM topup_orders WHERE stripe_payment_intent_id=$1 FOR UPDATE`, paymentIntentID).Scan(&orderID, &tenantID, &currency); err != nil {
+ return err
+ }
+ amountMicros, err = minorToMicros(currency, refund.Amount)
+ if err != nil {
+ return err
+ }
+ err = tx.QueryRow(ctx, `INSERT INTO stripe_refunds (tenant_id,topup_order_id,stripe_refund_id,amount_minor,amount_micros,currency,reason,status)
+ VALUES ($1,$2,$3,$4,$5,$6,$7,$8) RETURNING id::text`, tenantID, orderID, refund.ID, refund.Amount, amountMicros, currency, string(refund.Reason), string(refund.Status)).Scan(&id)
+ if err != nil {
+ return err
+ }
+ currentStatus = ""
+ } else if err != nil {
+ return err
+ }
+ if _, err := tx.Exec(ctx, `UPDATE stripe_refunds SET stripe_refund_id=$2,status=$3,last_error=$4,updated_at=now(),
+ completed_at=CASE WHEN $3 IN ('succeeded','failed','canceled') THEN now() ELSE completed_at END WHERE id=$1`, id, refund.ID, string(refund.Status), string(refund.FailureReason)); err != nil {
+ return err
+ }
+ if refund.Status == stripe.RefundStatusSucceeded && currentStatus != "succeeded" {
+ 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 held > 0 {
+ if held > reserved {
+ return errors.New("refund hold invariant violated")
+ }
+ reserved -= held
+ }
+ debit := amountMicros
+ if debit > balance-reserved {
+ debit = balance - reserved
+ }
+ if debit < 0 {
+ debit = 0
+ }
+ newBalance := balance - debit
+ if _, err := tx.Exec(ctx, `UPDATE tenant_wallets SET balance_micros=$2,reserved_micros=$3,updated_at=now() WHERE tenant_id=$1`, tenantID, newBalance, reserved); err != nil {
+ return err
+ }
+ if _, err := tx.Exec(ctx, `UPDATE stripe_refunds SET held_micros=0,uncollected_micros=$2 WHERE id=$1`, id, amountMicros-debit); 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)
+ SELECT $1,currency,$2,$3,'refund','stripe_refund',$4,'Stripe top-up refund' FROM topup_orders WHERE id=$5
+ ON CONFLICT (source_type,source_id) DO NOTHING`, tenantID, -debit, newBalance, refund.ID, orderID); err != nil {
+ return err
+ }
+ if _, err := tx.Exec(ctx, `UPDATE topup_orders SET refunded_micros=LEAST(amount_micros,refunded_micros+$2),
+ status=CASE WHEN refunded_micros+$2>=amount_micros THEN 'refunded' ELSE 'partially_refunded' END WHERE id=$1`, orderID, amountMicros); err != nil {
+ return err
+ }
+ } else if (refund.Status == stripe.RefundStatusFailed || refund.Status == stripe.RefundStatusCanceled) && held > 0 {
+ if _, err := tx.Exec(ctx, `UPDATE tenant_wallets SET reserved_micros=reserved_micros-$2,updated_at=now() WHERE tenant_id=$1`, tenantID, held); err != nil {
+ return err
+ }
+ if _, err := tx.Exec(ctx, `UPDATE stripe_refunds SET held_micros=0 WHERE id=$1`, id); err != nil {
+ return err
+ }
+ }
+ return nil
+}
+
+func (s *Service) applyDisputeTx(ctx context.Context, tx pgx.Tx, dispute *stripe.Dispute, eventType stripe.EventType) error {
+ paymentIntentID := ""
+ if dispute.PaymentIntent != nil {
+ paymentIntentID = dispute.PaymentIntent.ID
+ }
+ if paymentIntentID == "" && dispute.Charge != nil && dispute.Charge.PaymentIntent != nil {
+ paymentIntentID = dispute.Charge.PaymentIntent.ID
+ }
+ if paymentIntentID == "" {
+ return ErrInvalidAmount
+ }
+ var orderID, tenantID, currency string
+ if err := tx.QueryRow(ctx, `SELECT id::text,tenant_id::text,currency FROM topup_orders WHERE stripe_payment_intent_id=$1 FOR UPDATE`, paymentIntentID).Scan(&orderID, &tenantID, &currency); err != nil {
+ return err
+ }
+ amountMicros, err := minorToMicros(currency, dispute.Amount)
+ if err != nil {
+ return err
+ }
+ dueBy := (*time.Time)(nil)
+ if dispute.EvidenceDetails != nil && dispute.EvidenceDetails.DueBy > 0 {
+ value := time.Unix(dispute.EvidenceDetails.DueBy, 0).UTC()
+ dueBy = &value
+ }
+ var previousDebited int64
+ queryErr := tx.QueryRow(ctx, `SELECT debited_micros FROM stripe_disputes WHERE stripe_dispute_id=$1 FOR UPDATE`, dispute.ID).Scan(&previousDebited)
+ if queryErr != nil && !errors.Is(queryErr, pgx.ErrNoRows) {
+ return queryErr
+ }
+ if _, err := tx.Exec(ctx, `INSERT INTO stripe_disputes (stripe_dispute_id,tenant_id,topup_order_id,stripe_payment_intent_id,
+ amount_minor,amount_micros,currency,status,reason,due_by) VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10)
+ ON CONFLICT (stripe_dispute_id) DO UPDATE SET status=EXCLUDED.status,reason=EXCLUDED.reason,
+ due_by=EXCLUDED.due_by,updated_at=now(),closed_at=CASE WHEN EXCLUDED.status IN ('won','lost') THEN now() ELSE stripe_disputes.closed_at END`,
+ dispute.ID, tenantID, orderID, paymentIntentID, dispute.Amount, amountMicros, currency, string(dispute.Status), string(dispute.Reason), dueBy); err != nil {
+ return err
+ }
+ shouldDebit := eventType == stripe.EventTypeChargeDisputeCreated || eventType == stripe.EventTypeChargeDisputeFundsWithdrawn
+ shouldReverse := eventType == stripe.EventTypeChargeDisputeFundsReinstated || dispute.Status == stripe.DisputeStatusWon
+ if shouldDebit && previousDebited == 0 {
+ 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
+ }
+ debit := amountMicros
+ if debit > balance-reserved {
+ debit = balance - reserved
+ }
+ if debit < 0 {
+ debit = 0
+ }
+ newBalance := balance - debit
+ 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 stripe_disputes SET debited_micros=$2,uncollected_micros=$3 WHERE stripe_dispute_id=$1`, dispute.ID, debit, amountMicros-debit); 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,'dispute','stripe_dispute',$5,'Stripe payment dispute') ON CONFLICT DO NOTHING`, tenantID, currency, -debit, newBalance, dispute.ID); err != nil {
+ return err
+ }
+ if _, err := tx.Exec(ctx, `UPDATE topup_orders SET disputed_micros=GREATEST(disputed_micros,$2),status='disputed' WHERE id=$1`, orderID, amountMicros); err != nil {
+ return err
+ }
+ } else if shouldReverse && previousDebited > 0 {
+ var balance int64
+ if err := tx.QueryRow(ctx, `SELECT balance_micros FROM tenant_wallets WHERE tenant_id=$1 FOR UPDATE`, tenantID).Scan(&balance); err != nil {
+ return err
+ }
+ newBalance := balance + previousDebited
+ 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 stripe_disputes SET debited_micros=0,uncollected_micros=0 WHERE stripe_dispute_id=$1`, dispute.ID); 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,'dispute_reversal','stripe_dispute_reversal',$5,'Stripe dispute funds reinstated') ON CONFLICT DO NOTHING`, tenantID, currency, previousDebited, newBalance, dispute.ID); err != nil {
+ return err
+ }
+ if _, err := tx.Exec(ctx, `UPDATE topup_orders SET disputed_micros=0,status=CASE WHEN refunded_micros=0 THEN 'paid' WHEN refunded_micros<amount_micros THEN 'partially_refunded' ELSE 'refunded' END WHERE id=$1`, orderID); err != nil {
+ return err
+ }
+ }
+ return nil
+}
+
+func (s *Service) applyInvoiceTx(ctx context.Context, tx pgx.Tx, invoice *stripe.Invoice) error {
+ orderID := invoice.Metadata["aigw_topup_order_id"]
+ tenantID := invoice.Metadata["aigw_tenant_id"]
+ customerID := ""
+ if invoice.Customer != nil {
+ customerID = invoice.Customer.ID
+ }
+ if tenantID == "" && customerID != "" {
+ _ = tx.QueryRow(ctx, `SELECT tenant_id::text FROM stripe_customers WHERE stripe_customer_id=$1`, customerID).Scan(&tenantID)
+ }
+ if orderID == "" {
+ _ = tx.QueryRow(ctx, `SELECT id::text FROM topup_orders WHERE stripe_invoice_id=$1`, invoice.ID).Scan(&orderID)
+ }
+ failure := ""
+ if invoice.LastFinalizationError != nil {
+ failure = invoice.LastFinalizationError.Msg
+ }
+ var nextAttempt *time.Time
+ if invoice.NextPaymentAttempt > 0 {
+ value := time.Unix(invoice.NextPaymentAttempt, 0).UTC()
+ nextAttempt = &value
+ }
+ if _, err := tx.Exec(ctx, `INSERT INTO stripe_invoices (stripe_invoice_id,tenant_id,topup_order_id,stripe_customer_id,status,currency,
+ amount_due_minor,amount_paid_minor,attempt_count,next_payment_attempt,hosted_invoice_url,invoice_pdf_url,last_failure)
+ VALUES ($1,NULLIF($2,'')::uuid,NULLIF($3,'')::uuid,NULLIF($4,''),$5,$6,$7,$8,$9,$10,$11,$12,$13)
+ ON CONFLICT (stripe_invoice_id) DO UPDATE SET status=EXCLUDED.status,amount_due_minor=EXCLUDED.amount_due_minor,
+ amount_paid_minor=EXCLUDED.amount_paid_minor,attempt_count=EXCLUDED.attempt_count,next_payment_attempt=EXCLUDED.next_payment_attempt,
+ hosted_invoice_url=EXCLUDED.hosted_invoice_url,invoice_pdf_url=EXCLUDED.invoice_pdf_url,last_failure=EXCLUDED.last_failure,updated_at=now()`,
+ invoice.ID, tenantID, orderID, customerID, string(invoice.Status), string(invoice.Currency), invoice.AmountDue, invoice.AmountPaid,
+ invoice.AttemptCount, nextAttempt, invoice.HostedInvoiceURL, invoice.InvoicePDF, failure); err != nil {
+ return err
+ }
+ if orderID != "" {
+ _, err := tx.Exec(ctx, `UPDATE topup_orders SET stripe_invoice_id=$2,invoice_url=COALESCE(NULLIF($3,''),invoice_url),
+ invoice_pdf_url=COALESCE(NULLIF($4,''),invoice_pdf_url) WHERE id=$1`, orderID, invoice.ID, invoice.HostedInvoiceURL, invoice.InvoicePDF)
+ return err
+ }
+ return nil
+}
+
+func (s *Service) Reconcile(ctx context.Context, limit int) (ReconciliationResult, error) {
+ if !s.stripeEnabled || s.stripeClient == nil {
+ return ReconciliationResult{}, ErrStripeDisabled
+ }
+ if limit < 1 || limit > 500 {
+ limit = 100
+ }
+ var result ReconciliationResult
+ if err := s.db.QueryRow(ctx, `INSERT INTO billing_reconciliation_runs (status) VALUES ('running') RETURNING id::text`).Scan(&result.ID); err != nil {
+ return result, err
+ }
+ rows, err := s.db.Query(ctx, `SELECT id::text,stripe_session_id,status,amount_minor,currency,reconciliation_status FROM topup_orders
+ WHERE stripe_session_id IS NOT NULL ORDER BY created_at DESC LIMIT $1`, limit)
+ if err != nil {
+ return s.failReconciliation(ctx, result, err)
+ }
+ type order struct {
+ id, session, status, currency, reconciliationStatus string
+ amount int64
+ }
+ var orders []order
+ for rows.Next() {
+ var item order
+ if err := rows.Scan(&item.id, &item.session, &item.status, &item.amount, &item.currency, &item.reconciliationStatus); err != nil {
+ rows.Close()
+ return s.failReconciliation(ctx, result, err)
+ }
+ orders = append(orders, item)
+ }
+ rows.Close()
+ for _, item := range orders {
+ session, retrieveErr := s.stripeClient.V1CheckoutSessions.Retrieve(ctx, item.session, &stripe.CheckoutSessionRetrieveParams{})
+ result.CheckedOrders++
+ if retrieveErr != nil {
+ if stripeResourceMissing(retrieveErr) && item.reconciliationStatus == "resolved" && (item.status == "failed" || item.status == "reversed") {
+ continue
+ }
+ message := truncateError(retrieveErr)
+ typeName := "stripe_retrieve_failed"
+ state := "mismatch"
+ if stripeResourceMissing(retrieveErr) {
+ typeName = "stripe_session_missing"
+ state = "missing"
+ }
+ if err := s.updateOrderReconciliation(ctx, item.id, state, message); err != nil {
+ return s.failReconciliation(ctx, result, fmt.Errorf("record reconciliation retrieval failure for order %s: %w", item.id, err))
+ }
+ result.Mismatches = append(result.Mismatches, map[string]any{"order_id": item.id, "type": typeName, "error": message})
+ continue
+ }
+ expected := item.status == "paid" || item.status == "partially_refunded" || item.status == "refunded" || item.status == "disputed"
+ stripePaid := session.PaymentStatus == stripe.CheckoutSessionPaymentStatusPaid
+ if session.AmountTotal != item.amount || string(session.Currency) != item.currency {
+ result.Mismatches = append(result.Mismatches, map[string]any{"order_id": item.id, "type": "checkout_mismatch", "local_status": item.status, "stripe_payment_status": session.PaymentStatus, "local_amount": item.amount, "stripe_amount": session.AmountTotal})
+ if err := s.updateOrderReconciliation(ctx, item.id, "mismatch", "amount or currency mismatch"); err != nil {
+ return s.failReconciliation(ctx, result, fmt.Errorf("record checkout mismatch for order %s: %w", item.id, err))
+ }
+ continue
+ }
+ if stripePaid && !expected {
+ raw, marshalErr := json.Marshal(session)
+ if marshalErr != nil {
+ return s.failReconciliation(ctx, result, fmt.Errorf("encode Stripe Checkout Session for order %s: %w", item.id, marshalErr))
+ }
+ if repairErr := s.processStripeEvent(ctx, stripe.Event{ID: "reconcile_" + session.ID + "_paid", Type: stripe.EventTypeCheckoutSessionCompleted, Data: &stripe.EventData{Raw: raw}}); repairErr != nil {
+ message := truncateError(repairErr)
+ result.Mismatches = append(result.Mismatches, map[string]any{"order_id": item.id, "type": "checkout_repair_failed", "error": message})
+ if err := s.updateOrderReconciliation(ctx, item.id, "mismatch", message); err != nil {
+ return s.failReconciliation(ctx, result, fmt.Errorf("record checkout repair failure for order %s: %w", item.id, err))
+ }
+ continue
+ }
+ result.Repairs = append(result.Repairs, map[string]any{"order_id": item.id, "type": "credited_paid_checkout"})
+ if err := s.updateOrderReconciliation(ctx, item.id, "repaired", ""); err != nil {
+ return s.failReconciliation(ctx, result, fmt.Errorf("record checkout repair for order %s: %w", item.id, err))
+ }
+ continue
+ }
+ if expected != stripePaid {
+ result.Mismatches = append(result.Mismatches, map[string]any{"order_id": item.id, "type": "checkout_payment_state_mismatch", "local_status": item.status, "stripe_payment_status": session.PaymentStatus})
+ if err := s.updateOrderReconciliation(ctx, item.id, "mismatch", "payment state mismatch"); err != nil {
+ return s.failReconciliation(ctx, result, fmt.Errorf("record payment-state mismatch for order %s: %w", item.id, err))
+ }
+ continue
+ }
+ if err := s.updateOrderReconciliation(ctx, item.id, "ok", ""); err != nil {
+ return s.failReconciliation(ctx, result, fmt.Errorf("record clean reconciliation for order %s: %w", item.id, err))
+ }
+ }
+ result.MismatchCount = int64(len(result.Mismatches))
+ result.Status = "clean"
+ if result.MismatchCount > 0 {
+ result.Status = "mismatch"
+ }
+ report, marshalErr := json.Marshal(result.Mismatches)
+ if marshalErr != nil {
+ return s.failReconciliation(ctx, result, fmt.Errorf("encode reconciliation report: %w", marshalErr))
+ }
+ _, err = s.db.Exec(ctx, `UPDATE billing_reconciliation_runs SET status=$2,checked_orders=$3,mismatch_count=$4,report=$5,completed_at=now() WHERE id=$1`, result.ID, result.Status, result.CheckedOrders, result.MismatchCount, report)
+ if err != nil {
+ return s.failReconciliation(ctx, result, fmt.Errorf("complete reconciliation run: %w", err))
+ }
+ return result, nil
+}
+
+func (s *Service) updateOrderReconciliation(ctx context.Context, orderID, status, message string) error {
+ command, err := s.db.Exec(ctx, `UPDATE topup_orders SET reconciliation_status=$2,reconciled_at=now(),reconciliation_error=$3 WHERE id=$1`, orderID, status, message)
+ if err != nil {
+ return err
+ }
+ if command.RowsAffected() != 1 {
+ return fmt.Errorf("expected one top-up order, updated %d", command.RowsAffected())
+ }
+ return nil
+}
+
+func stripeResourceMissing(err error) bool {
+ var stripeErr *stripe.Error
+ return errors.As(err, &stripeErr) && (stripeErr.Code == stripe.ErrorCodeResourceMissing || stripeErr.HTTPStatusCode == 404)
+}
+
+func truncateError(err error) string {
+ if err == nil {
+ return ""
+ }
+ message := err.Error()
+ if len(message) > 1000 {
+ return message[:1000]
+ }
+ return message
+}
+
+func (s *Service) failReconciliation(ctx context.Context, result ReconciliationResult, cause error) (ReconciliationResult, error) {
+ if _, err := s.db.Exec(ctx, `UPDATE billing_reconciliation_runs SET status='failed',error=$2,completed_at=now() WHERE id=$1`, result.ID, truncateError(cause)); err != nil {
+ return result, errors.Join(cause, fmt.Errorf("record failed reconciliation run: %w", err))
+ }
+ return result, cause
+}
+
+func (s *Service) WriteFinancialCSV(ctx context.Context, tenantID string, output io.Writer) error {
+ query := `
+ SELECT id::text, tenant_id::text, COALESCE(project_id::text, ''), currency, amount_micros,
+ balance_after_micros, kind, source_type, source_id, description, created_at
+ FROM billing_ledger`
+ args := []any{}
+ if strings.TrimSpace(tenantID) != "" {
+ query += ` WHERE tenant_id=$1`
+ args = append(args, tenantID)
+ }
+ query += ` ORDER BY created_at, id`
+ rows, err := s.db.Query(ctx, query, args...)
+ if err != nil {
+ return err
+ }
+ defer rows.Close()
+ w := csv.NewWriter(output)
+ if err := w.Write([]string{"timestamp", "tenant_id", "project_id", "currency", "kind", "amount_micros", "balance_after_micros", "source_type", "source_id", "description"}); err != nil {
+ return err
+ }
+ for rows.Next() {
+ var e LedgerEntry
+ if err := rows.Scan(&e.ID, &e.TenantID, &e.ProjectID, &e.Currency, &e.AmountMicros,
+ &e.BalanceAfterMicros, &e.Kind, &e.SourceType, &e.SourceID, &e.Description, &e.CreatedAt); err != nil {
+ return err
+ }
+ if err := w.Write([]string{e.CreatedAt.UTC().Format(time.RFC3339Nano), e.TenantID, e.ProjectID, e.Currency, e.Kind, strconv.FormatInt(e.AmountMicros, 10), strconv.FormatInt(e.BalanceAfterMicros, 10), e.SourceType, e.SourceID, e.Description}); err != nil {
+ return err
+ }
+ }
+ if err := rows.Err(); err != nil {
+ return err
+ }
+ w.Flush()
+ return w.Error()
+}
diff --git a/internal/billing/operations_test.go b/internal/billing/operations_test.go
new file mode 100644
index 0000000..46bedeb
--- /dev/null
+++ b/internal/billing/operations_test.go
@@ -0,0 +1,36 @@
+package billing
+
+import (
+ "testing"
+ "time"
+)
+
+func TestOperationalStatusReadiness(t *testing.T) {
+ now := time.Now().UTC()
+ recent := now.Add(-time.Minute)
+ clean := OperationalStatus{
+ StripeEnabled: true, ReconciliationStatus: "clean", ReconciliationCompletedAt: &recent,
+ }
+ if !clean.Ready(now) {
+ t.Fatal("clean operational state should be ready")
+ }
+ oldWebhook := now.Add(-6 * time.Minute)
+ stuckWebhook := clean
+ stuckWebhook.UnprocessedWebhooks = 1
+ stuckWebhook.OldestUnprocessedWebhook = &oldWebhook
+ if stuckWebhook.Ready(now) {
+ t.Fatal("stuck webhook should fail readiness")
+ }
+ freshRefund := now.Add(-time.Minute)
+ processingRefund := clean
+ processingRefund.RefundBacklog = 1
+ processingRefund.OldestRefund = &freshRefund
+ if !processingRefund.Ready(now) {
+ t.Fatal("fresh refund operation should remain ready during its processing window")
+ }
+ unmetered := clean
+ unmetered.UnmeteredSuccesses = 1
+ if unmetered.Ready(now) {
+ t.Fatal("unmetered success must fail readiness")
+ }
+}
diff --git a/internal/billing/service.go b/internal/billing/service.go
index b87a0bd..30f0e32 100644
--- a/internal/billing/service.go
+++ b/internal/billing/service.go
@@ -1,6 +1,7 @@
package billing
import (
+ "bufio"
"context"
"crypto/rand"
"encoding/hex"
@@ -8,13 +9,17 @@ import (
"errors"
"fmt"
"math/big"
+ "os"
+ "path/filepath"
"strings"
+ "sync"
"time"
"aigw/internal/domain"
"github.com/jackc/pgx/v5"
"github.com/jackc/pgx/v5/pgxpool"
+ "github.com/stripe/stripe-go/v86"
)
const microsPerUnit = int64(1_000_000)
@@ -29,8 +34,15 @@ type Service struct {
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
}
func New(ctx context.Context, options Options) (*Service, error) {
@@ -47,10 +59,15 @@ func New(ctx context.Context, options Options) (*Service, error) {
minTopUpMinor: options.MinTopUpMinor, maxTopUpMinor: options.MaxTopUpMinor,
stripeEnabled: options.StripeEnabled, stripeWebhookSecret: options.StripeWebhookSecret,
stripeSuccessURL: options.StripeSuccessURL, stripeCancelURL: options.StripeCancelURL,
+ stripePortalReturnURL: options.StripePortalReturnURL, stripeAutomaticTax: options.StripeAutomaticTax,
+ stripeProductTaxCode: options.StripeProductTaxCode,
integrationIdentifier: "aigw_balance_" + randomLetters(8),
+ settlementSpoolPath: strings.TrimSpace(options.SettlementSpoolPath),
+ metrics: options.Metrics,
}
if options.StripeEnabled {
- service.createStripeCheckout = newStripeCheckoutCreator(options.StripeAPIKey)
+ service.stripeClient = stripe.NewClient(options.StripeAPIKey)
+ service.createStripeCheckout = service.stripeClient.V1CheckoutSessions.Create
}
return service, nil
}
@@ -67,10 +84,15 @@ func (s *Service) Currency() string {
return s.currency
}
+func (s *Service) Ping(ctx context.Context) error { return s.db.Ping(ctx) }
+
func (s *Service) Authorize(ctx context.Context, input Authorization) error {
if input.RequestID == "" || input.Principal.TenantID == "" || input.Principal.ProjectID == "" || input.Principal.KeyID == "" {
return errors.New("billing authorization identity is incomplete")
}
+ if input.Model.PriceCurrency != "" && input.Model.PriceCurrency != s.currency {
+ return fmt.Errorf("model price currency %s does not match wallet currency %s", input.Model.PriceCurrency, s.currency)
+ }
reserved, err := reservationCost(input.Model, input.Body, s.defaultMaxOutputTokens)
if err != nil {
return err
@@ -101,7 +123,7 @@ func (s *Service) Authorize(ctx context.Context, input Authorization) error {
var used, pending int64
if err := tx.QueryRow(ctx, `SELECT
COALESCE((SELECT cost_micros FROM usage_monthly_rollups WHERE project_id=$1 AND period_start=$2),0),
- COALESCE((SELECT sum(reserved_micros) FROM billing_reservations WHERE project_id=$1 AND status='pending' AND created_at >= $2 AND created_at < $3),0)`,
+ COALESCE((SELECT sum(reserved_micros) FROM billing_reservations WHERE project_id=$1 AND status IN ('pending','metering_failed') AND created_at >= $2 AND created_at < $3),0)`,
input.Principal.ProjectID, period, nextPeriod).Scan(&used, &pending); err != nil {
return fmt.Errorf("read monthly spend quota: %w", err)
}
@@ -115,15 +137,19 @@ func (s *Service) Authorize(ctx context.Context, input Authorization) error {
if _, err := tx.Exec(ctx, `
INSERT INTO billing_reservations (
request_id, tenant_id, project_id, key_id, public_model, currency, reserved_micros,
+ price_version_id,
input_price_micros_per_million, output_price_micros_per_million,
cache_read_price_micros_per_million, cache_write_price_micros_per_million)
- VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11)`,
+ VALUES ($1,$2,$3,$4,$5,$6,$7,NULLIF($8,'')::uuid,$9,$10,$11,$12)`,
input.RequestID, input.Principal.TenantID, input.Principal.ProjectID, input.Principal.KeyID,
- input.Model.ID, s.currency, reserved, input.Model.InputPriceMicrosPerMillion,
+ input.Model.ID, s.currency, reserved, input.Model.PriceVersionID, input.Model.InputPriceMicrosPerMillion,
input.Model.OutputPriceMicrosPerMillion, input.Model.CacheReadPriceMicrosPerMillion,
input.Model.CacheWritePriceMicrosPerMillion); err != nil {
return fmt.Errorf("create billing reservation: %w", err)
}
+ if _, err := tx.Exec(ctx, `INSERT INTO billing_settlement_jobs (request_id) VALUES ($1)`, input.RequestID); err != nil {
+ return fmt.Errorf("create settlement job: %w", err)
+ }
if _, err := tx.Exec(ctx, `
UPDATE tenant_wallets SET reserved_micros = reserved_micros + $2, updated_at = now()
WHERE tenant_id = $1`, input.Principal.TenantID, reserved); err != nil {
@@ -135,7 +161,27 @@ func (s *Service) Authorize(ctx context.Context, input Authorization) error {
return nil
}
-func (s *Service) Settle(ctx context.Context, event domain.UsageEvent) error {
+func (s *Service) EnqueueSettlement(ctx context.Context, event domain.UsageEvent) error {
+ payload, err := json.Marshal(event)
+ if err != nil {
+ return fmt.Errorf("encode settlement event: %w", err)
+ }
+ _, err = s.db.Exec(ctx, `
+ INSERT INTO billing_settlement_jobs (request_id, event, status, available_at, updated_at)
+ VALUES ($1,$2,'pending',now(),now())
+ ON CONFLICT (request_id) DO UPDATE SET event=EXCLUDED.event,
+ status=CASE WHEN billing_settlement_jobs.status='done' THEN 'done' ELSE 'pending' END,
+ available_at=now(), locked_at=NULL, last_error='', updated_at=now()`, event.RequestID, payload)
+ if err == nil {
+ return nil
+ }
+ if spoolErr := s.appendSettlementSpool(payload); spoolErr != nil {
+ return fmt.Errorf("enqueue settlement in PostgreSQL: %v; append durable spool: %w", err, spoolErr)
+ }
+ return nil
+}
+
+func (s *Service) settle(ctx context.Context, event domain.UsageEvent) error {
tx, err := s.db.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.ReadCommitted})
if err != nil {
return fmt.Errorf("begin usage settlement: %w", err)
@@ -159,7 +205,51 @@ func (s *Service) Settle(ctx context.Context, event domain.UsageEvent) error {
}
actualCost := int64(0)
- if event.StatusCode >= 200 && event.StatusCode < 300 {
+ billableSuccess := event.StatusCode >= 200 && event.StatusCode < 300 && event.Success
+ if billableSuccess && !event.UsageReported {
+ // Fail closed: keep the authorization hold in place and make the request
+ // visible to reconciliation. Releasing it would turn an unmetered success
+ // into a free request; guessing tokens here could overcharge the customer.
+ if _, err := tx.Exec(ctx, `UPDATE billing_reservations SET status='metering_failed', settled_at=now()
+ WHERE request_id=$1`, event.RequestID); err != nil {
+ return fmt.Errorf("mark unmetered reservation: %w", err)
+ }
+ var usageAlreadyRecorded bool
+ if err := tx.QueryRow(ctx, `SELECT true FROM usage_events WHERE request_id=$1 FOR UPDATE`, event.RequestID).Scan(&usageAlreadyRecorded); err != nil && !errors.Is(err, pgx.ErrNoRows) {
+ return fmt.Errorf("lock unmetered usage event: %w", err)
+ }
+ if _, err := tx.Exec(ctx, `INSERT INTO usage_events (
+ request_id,tenant_id,project_id,key_id,public_model,provider_id,upstream_model,protocol,stream,
+ status_code,success,error_type,attempts,started_at,duration_ms,input_tokens,output_tokens,total_tokens,
+ cache_creation_input_tokens,cache_read_input_tokens,cost_micros,charged_micros,uncollected_micros,
+ usage_reported,metering_status)
+ VALUES ($1,$2,$3,$4,$5,NULLIF($6,''),NULLIF($7,''),$8,$9,$10,$11,'usage_not_reported',$12,$13,$14,
+ $15,$16,$17,$18,$19,0,0,0,false,'missing')
+ ON CONFLICT (request_id) DO UPDATE SET error_type='usage_not_reported',usage_reported=false,metering_status='missing'`,
+ event.RequestID, tenantID, projectID, keyID, modelID, event.ProviderID, event.UpstreamModel,
+ string(event.Protocol), event.Stream, event.StatusCode, event.Success, event.Attempts, event.StartedAt,
+ event.DurationMS, event.Usage.InputTokens, event.Usage.OutputTokens, event.Usage.TotalTokens,
+ event.Usage.CacheCreationInputTokens, event.Usage.CacheReadInputTokens); err != nil {
+ return fmt.Errorf("persist unmetered usage event: %w", err)
+ }
+ if !usageAlreadyRecorded {
+ period := time.Date(event.StartedAt.UTC().Year(), event.StartedAt.UTC().Month(), 1, 0, 0, 0, 0, time.UTC)
+ if _, err := tx.Exec(ctx, `INSERT INTO usage_monthly_rollups
+ (period_start,tenant_id,project_id,request_count,successful_requests,input_tokens,output_tokens,total_tokens)
+ VALUES ($1,$2,$3,1,1,$4,$5,$6)
+ ON CONFLICT (project_id,period_start) DO UPDATE SET
+ request_count=usage_monthly_rollups.request_count+1,
+ successful_requests=usage_monthly_rollups.successful_requests+1,
+ input_tokens=usage_monthly_rollups.input_tokens+EXCLUDED.input_tokens,
+ output_tokens=usage_monthly_rollups.output_tokens+EXCLUDED.output_tokens,
+ total_tokens=usage_monthly_rollups.total_tokens+EXCLUDED.total_tokens,updated_at=now()`,
+ period, tenantID, projectID, event.Usage.InputTokens, event.Usage.OutputTokens, event.Usage.TotalTokens); err != nil {
+ return fmt.Errorf("roll up unmetered usage event: %w", err)
+ }
+ }
+ return tx.Commit(ctx)
+ }
+ if billableSuccess {
actualCost, err = usageCost(event.Usage, inputPrice, outputPrice, cacheReadPrice, cacheWritePrice)
if err != nil {
return err
@@ -195,8 +285,9 @@ func (s *Service) Settle(ctx context.Context, event domain.UsageEvent) error {
return fmt.Errorf("lock existing usage event: %w", err)
}
if usageAlreadyRecorded {
- if _, err := tx.Exec(ctx, `UPDATE usage_events SET cost_micros=$2, charged_micros=$3, uncollected_micros=$4 WHERE request_id=$1`,
- event.RequestID, actualCost, charged, uncollected); err != nil {
+ if _, err := tx.Exec(ctx, `UPDATE usage_events SET cost_micros=$2, charged_micros=$3, uncollected_micros=$4,
+ usage_reported=$5,metering_status=$6 WHERE request_id=$1`,
+ event.RequestID, actualCost, charged, uncollected, event.UsageReported, meteringStatus(event)); err != nil {
return fmt.Errorf("apply usage charge: %w", err)
}
} else if _, err := tx.Exec(ctx, `
@@ -204,14 +295,14 @@ func (s *Service) Settle(ctx context.Context, event domain.UsageEvent) error {
request_id, tenant_id, project_id, key_id, public_model, provider_id, upstream_model,
protocol, stream, status_code, success, error_type, attempts, started_at, duration_ms,
input_tokens, output_tokens, total_tokens, cache_creation_input_tokens, cache_read_input_tokens,
- cost_micros, charged_micros, uncollected_micros)
- VALUES ($1,$2,$3,$4,$5,NULLIF($6,''),NULLIF($7,''),$8,$9,$10,$11,$12,$13,$14,$15,$16,$17,$18,$19,$20,$21,$22,$23)
+ cost_micros, charged_micros, uncollected_micros, usage_reported, metering_status)
+ VALUES ($1,$2,$3,$4,$5,NULLIF($6,''),NULLIF($7,''),$8,$9,$10,$11,$12,$13,$14,$15,$16,$17,$18,$19,$20,$21,$22,$23,$24,$25)
ON CONFLICT (request_id) DO NOTHING`,
event.RequestID, tenantID, projectID, keyID, modelID, event.ProviderID, event.UpstreamModel,
string(event.Protocol), event.Stream, event.StatusCode, event.Success, event.ErrorType, event.Attempts,
event.StartedAt, event.DurationMS, event.Usage.InputTokens, event.Usage.OutputTokens,
event.Usage.TotalTokens, event.Usage.CacheCreationInputTokens, event.Usage.CacheReadInputTokens,
- actualCost, charged, uncollected); err != nil {
+ actualCost, charged, uncollected, event.UsageReported, meteringStatus(event)); err != nil {
return fmt.Errorf("persist usage event: %w", err)
}
if charged > 0 {
@@ -251,6 +342,199 @@ func (s *Service) Settle(ctx context.Context, event domain.UsageEvent) error {
return nil
}
+// RunSettlementWorker processes jobs using PostgreSQL row locks so any gateway
+// instance can resume work left by another instance after a crash.
+func (s *Service) RunSettlementWorker(ctx context.Context) {
+ ticker := time.NewTicker(time.Second)
+ defer ticker.Stop()
+ for {
+ s.drainSettlementSpool(ctx)
+ s.recoverStaleSettlements(ctx)
+ for i := 0; i < 32; i++ {
+ processed, err := s.processSettlementJob(ctx)
+ if err != nil || !processed {
+ break
+ }
+ }
+ select {
+ case <-ctx.Done():
+ return
+ case <-ticker.C:
+ }
+ }
+}
+
+func (s *Service) recoverStaleSettlements(ctx context.Context) {
+ // A process may die after authorization but before it can attach a response
+ // event. Releasing after an hour prevents permanent holds; such synthetic
+ // events stay visible in Usage for reconciliation.
+ rows, err := s.db.Query(ctx, `SELECT request_id,tenant_id::text,project_id::text,key_id::text,public_model,created_at
+ FROM billing_reservations WHERE status='pending' AND created_at<now()-interval '1 hour'
+ AND EXISTS(SELECT 1 FROM billing_settlement_jobs j WHERE j.request_id=billing_reservations.request_id AND j.status='awaiting_event') LIMIT 100`)
+ if err != nil {
+ return
+ }
+ defer rows.Close()
+ for rows.Next() {
+ var event domain.UsageEvent
+ if rows.Scan(&event.RequestID, &event.TenantID, &event.ProjectID, &event.KeyID, &event.PublicModel, &event.StartedAt) != nil {
+ continue
+ }
+ event.StatusCode = 500
+ event.Success = false
+ event.ErrorType = "gateway_interrupted_before_settlement"
+ event.DurationMS = time.Since(event.StartedAt).Milliseconds()
+ _ = s.EnqueueSettlement(ctx, event)
+ }
+ _, _ = s.db.Exec(ctx, `UPDATE billing_settlement_jobs SET status='retry',locked_at=NULL,available_at=now(),updated_at=now(),last_error='recovered stale processing lease'
+ WHERE status='processing' AND locked_at<now()-interval '5 minutes'`)
+}
+
+func (s *Service) processSettlementJob(ctx context.Context) (bool, error) {
+ tx, err := s.db.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.ReadCommitted})
+ if err != nil {
+ return false, err
+ }
+ defer tx.Rollback(ctx)
+ var requestID string
+ var payload []byte
+ err = tx.QueryRow(ctx, `
+ WITH selected AS (
+ SELECT request_id FROM billing_settlement_jobs
+ WHERE status IN ('pending','retry') AND available_at <= now()
+ ORDER BY available_at, created_at FOR UPDATE SKIP LOCKED LIMIT 1
+ )
+ UPDATE billing_settlement_jobs j SET status='processing', attempts=attempts+1,
+ locked_at=now(), updated_at=now()
+ FROM selected WHERE j.request_id=selected.request_id
+ RETURNING j.request_id, j.event`,
+ ).Scan(&requestID, &payload)
+ if errors.Is(err, pgx.ErrNoRows) {
+ return false, tx.Commit(ctx)
+ }
+ if err != nil {
+ return false, err
+ }
+ if err := tx.Commit(ctx); err != nil {
+ return false, err
+ }
+ var event domain.UsageEvent
+ if err := json.Unmarshal(payload, &event); err != nil {
+ s.retrySettlement(ctx, requestID, fmt.Errorf("decode settlement event: %w", err))
+ return true, err
+ }
+ jobCtx, cancel := context.WithTimeout(ctx, 15*time.Second)
+ err = s.settle(jobCtx, event)
+ cancel()
+ if err != nil {
+ s.retrySettlement(ctx, requestID, err)
+ return true, err
+ }
+ _, err = s.db.Exec(ctx, `UPDATE billing_settlement_jobs SET status='done', locked_at=NULL,
+ last_error='', updated_at=now(), completed_at=now() WHERE request_id=$1`, requestID)
+ return true, err
+}
+
+func (s *Service) retrySettlement(ctx context.Context, requestID string, cause error) {
+ message := cause.Error()
+ if len(message) > 1000 {
+ message = message[:1000]
+ }
+ _, _ = s.db.Exec(ctx, `UPDATE billing_settlement_jobs SET status='retry', locked_at=NULL,
+ last_error=$2, available_at=now() + make_interval(secs => LEAST(300, power(2, LEAST(attempts, 8))::int)),
+ updated_at=now() WHERE request_id=$1`, requestID, message)
+}
+
+func (s *Service) appendSettlementSpool(payload []byte) error {
+ if s.settlementSpoolPath == "" {
+ return errors.New("settlement spool path is not configured")
+ }
+ s.spoolMu.Lock()
+ defer s.spoolMu.Unlock()
+ if err := os.MkdirAll(filepath.Dir(s.settlementSpoolPath), 0o700); err != nil {
+ return err
+ }
+ file, err := os.OpenFile(s.settlementSpoolPath, os.O_CREATE|os.O_APPEND|os.O_WRONLY, 0o600)
+ if err != nil {
+ return err
+ }
+ defer file.Close()
+ if _, err := file.Write(append(payload, '\n')); err != nil {
+ return err
+ }
+ return file.Sync()
+}
+
+func (s *Service) drainSettlementSpool(ctx context.Context) {
+ if s.settlementSpoolPath == "" {
+ return
+ }
+ s.spoolMu.Lock()
+ defer s.spoolMu.Unlock()
+ file, err := os.Open(s.settlementSpoolPath)
+ if errors.Is(err, os.ErrNotExist) {
+ return
+ }
+ if err != nil {
+ return
+ }
+ var pending [][]byte
+ scanner := bufio.NewScanner(file)
+ scanner.Buffer(make([]byte, 64*1024), 2<<20)
+ for scanner.Scan() {
+ line := append([]byte(nil), scanner.Bytes()...)
+ var event domain.UsageEvent
+ if json.Unmarshal(line, &event) != nil || event.RequestID == "" {
+ pending = append(pending, line)
+ continue
+ }
+ if _, err := s.db.Exec(ctx, `UPDATE billing_settlement_jobs SET event=$2, status=CASE WHEN status='done' THEN 'done' ELSE 'pending' END,
+ available_at=now(), locked_at=NULL, updated_at=now() WHERE request_id=$1`, event.RequestID, line); err != nil {
+ pending = append(pending, line)
+ }
+ }
+ _ = file.Close()
+ if scanner.Err() != nil {
+ return
+ }
+ temporary := s.settlementSpoolPath + ".tmp"
+ out, err := os.OpenFile(temporary, os.O_CREATE|os.O_TRUNC|os.O_WRONLY, 0o600)
+ if err != nil {
+ return
+ }
+ for _, line := range pending {
+ _, _ = out.Write(append(line, '\n'))
+ }
+ _ = out.Sync()
+ _ = out.Close()
+ _ = os.Rename(temporary, s.settlementSpoolPath)
+}
+
+func (s *Service) SettlementQueueStatus(ctx context.Context) (SettlementQueueStatus, error) {
+ var result SettlementQueueStatus
+ err := s.db.QueryRow(ctx, `SELECT
+ count(*) FILTER (WHERE status='awaiting_event'), count(*) FILTER (WHERE status='pending'),
+ count(*) FILTER (WHERE status='processing'), count(*) FILTER (WHERE status='retry'),
+ min(created_at) FILTER (WHERE status IN ('awaiting_event','pending','processing','retry'))
+ FROM billing_settlement_jobs`).Scan(&result.AwaitingEvent, &result.Pending, &result.Processing, &result.Retrying, &result.OldestPending)
+ if err != nil {
+ return result, err
+ }
+ if s.settlementSpoolPath != "" {
+ s.spoolMu.Lock()
+ file, openErr := os.Open(s.settlementSpoolPath)
+ if openErr == nil {
+ scanner := bufio.NewScanner(file)
+ for scanner.Scan() {
+ result.SpoolRecords++
+ }
+ _ = file.Close()
+ }
+ s.spoolMu.Unlock()
+ }
+ return result, nil
+}
+
func boolToInt(value bool) int {
if value {
return 1
@@ -258,8 +542,21 @@ func boolToInt(value bool) int {
return 0
}
+func meteringStatus(event domain.UsageEvent) string {
+ if event.StatusCode < 200 || event.StatusCode >= 300 || !event.Success {
+ return "upstream_failed"
+ }
+ if event.UsageReported {
+ return "reported"
+ }
+ return "missing"
+}
+
func reservationCost(model domain.Model, body []byte, defaultMaxOutput int64) (int64, error) {
maxOutput := defaultMaxOutput
+ if model.MaxOutputTokens > 0 && (maxOutput == 0 || model.MaxOutputTokens < maxOutput) {
+ maxOutput = model.MaxOutputTokens
+ }
var limits struct {
MaxTokens int64 `json:"max_tokens"`
MaxCompletionTokens int64 `json:"max_completion_tokens"`
diff --git a/internal/billing/service_test.go b/internal/billing/service_test.go
index f675da7..6a1bb49 100644
--- a/internal/billing/service_test.go
+++ b/internal/billing/service_test.go
@@ -2,6 +2,7 @@ package billing
import (
"context"
+ "crypto/sha256"
"encoding/json"
"fmt"
"net/http"
@@ -97,6 +98,26 @@ func TestCheckoutReturnURLPreservesCallbackAndSessionPlaceholder(t *testing.T) {
}
}
+func TestSettlementSpoolWritesOneDurableRecord(t *testing.T) {
+ path := t.TempDir() + "/settlements.jsonl"
+ service := &Service{settlementSpoolPath: path}
+ event := domain.UsageEvent{RequestID: "req_spool_test", TenantID: "tenant", StartedAt: time.Now().UTC()}
+ payload, err := json.Marshal(event)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if err := service.appendSettlementSpool(payload); err != nil {
+ t.Fatal(err)
+ }
+ contents, err := os.ReadFile(path)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if strings.Count(string(contents), "req_spool_test") != 1 || !strings.HasSuffix(string(contents), "\n") {
+ t.Fatalf("unexpected spool contents %q", contents)
+ }
+}
+
func TestWebhookRejectsInvalidSignatureBeforeProcessing(t *testing.T) {
service := &Service{stripeWebhookSecret: "whsec_test"}
request := httptest.NewRequest(http.MethodPost, "/billing/stripe/webhook", strings.NewReader(`{"id":"evt_fake"}`))
@@ -133,6 +154,7 @@ func TestStripeWebhookCreditsPaidOrderExactlyOncePostgres(t *testing.T) {
t.Fatal(err)
}
eventID := fmt.Sprintf("evt_aigw_%d", time.Now().UnixNano())
+ followupEventID := eventID + "_async"
t.Cleanup(func() {
cleanupCtx := context.Background()
for _, statement := range []struct {
@@ -140,6 +162,7 @@ func TestStripeWebhookCreditsPaidOrderExactlyOncePostgres(t *testing.T) {
arg string
}{
{`DELETE FROM stripe_webhook_events WHERE event_id=$1`, eventID},
+ {`DELETE FROM stripe_webhook_events WHERE event_id=$1`, followupEventID},
{`DELETE FROM billing_ledger WHERE tenant_id=$1`, tenantID},
{`DELETE FROM tenant_wallets WHERE tenant_id=$1`, tenantID},
{`DELETE FROM topup_orders WHERE tenant_id=$1`, tenantID},
@@ -177,6 +200,28 @@ func TestStripeWebhookCreditsPaidOrderExactlyOncePostgres(t *testing.T) {
t.Fatalf("delivery %d status = %d, body = %s", delivery+1, response.Code, response.Body.String())
}
}
+ if _, err := service.db.Exec(ctx, `UPDATE topup_orders SET status='refunded',refunded_micros=amount_micros WHERE id=$1`, orderID); err != nil {
+ t.Fatal(err)
+ }
+ followupPayload, err := json.Marshal(map[string]any{
+ "id": followupEventID, "object": "event", "api_version": stripe.APIVersion,
+ "type": string(stripe.EventTypeCheckoutSessionAsyncPaymentSucceeded),
+ "data": map[string]any{"object": map[string]any{
+ "id": sessionID, "object": "checkout.session", "client_reference_id": orderID,
+ "amount_total": amountMinor, "currency": "usd", "payment_status": "paid",
+ }},
+ })
+ if err != nil {
+ t.Fatal(err)
+ }
+ followupSigned := webhook.GenerateTestSignedPayload(&webhook.UnsignedPayload{Payload: followupPayload, Secret: service.stripeWebhookSecret})
+ followupRequest := httptest.NewRequest(http.MethodPost, "/billing/stripe/webhook", strings.NewReader(string(followupPayload)))
+ followupRequest.Header.Set("Stripe-Signature", followupSigned.Header)
+ followupResponse := httptest.NewRecorder()
+ service.WebhookHandler().ServeHTTP(followupResponse, followupRequest)
+ if followupResponse.Code != http.StatusOK {
+ t.Fatalf("follow-up status = %d, body = %s", followupResponse.Code, followupResponse.Body.String())
+ }
var balance int64
if err := service.db.QueryRow(ctx, `SELECT balance_micros FROM tenant_wallets WHERE tenant_id=$1`, tenantID).Scan(&balance); err != nil {
@@ -190,13 +235,136 @@ func TestStripeWebhookCreditsPaidOrderExactlyOncePostgres(t *testing.T) {
if err := service.db.QueryRow(ctx, `SELECT count(*) FROM billing_ledger WHERE source_type='stripe_checkout' AND source_id=$1`, sessionID).Scan(&ledgerCount); err != nil {
t.Fatal(err)
}
- if err := service.db.QueryRow(ctx, `SELECT count(*) FROM stripe_webhook_events WHERE event_id=$1`, eventID).Scan(&webhookCount); err != nil {
+ if err := service.db.QueryRow(ctx, `SELECT count(*) FROM stripe_webhook_events WHERE event_id IN ($1,$2)`, eventID, followupEventID).Scan(&webhookCount); err != nil {
t.Fatal(err)
}
if err := service.db.QueryRow(ctx, `SELECT status FROM topup_orders WHERE id=$1`, orderID).Scan(&orderStatus); err != nil {
t.Fatal(err)
}
- if ledgerCount != 1 || webhookCount != 1 || orderStatus != "paid" {
- t.Fatalf("ledger=%d webhook=%d order=%s, want 1/1/paid", ledgerCount, webhookCount, orderStatus)
+ if ledgerCount != 1 || webhookCount != 2 || orderStatus != "refunded" {
+ t.Fatalf("ledger=%d webhook=%d order=%s, want 1/2/refunded", ledgerCount, webhookCount, orderStatus)
+ }
+}
+
+func TestSettlementWorkerPersistsUsageAndReleasesReservationPostgres(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", SettlementSpoolPath: t.TempDir() + "/settlements.jsonl"})
+ if err != nil {
+ t.Fatal(err)
+ }
+ t.Cleanup(service.Close)
+ slug := fmt.Sprintf("settle-%d", time.Now().UnixNano())
+ var tenantID, projectID, keyID string
+ if err := service.db.QueryRow(ctx, `INSERT INTO tenants (slug,name) VALUES ($1,'Settlement integration') 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','Settlement') RETURNING id::text`, tenantID).Scan(&projectID); err != nil {
+ t.Fatal(err)
+ }
+ keyHash := sha256.Sum256([]byte(slug))
+ if err := service.db.QueryRow(ctx, `INSERT INTO api_keys (tenant_id,project_id,name,key_prefix,key_hash) VALUES ($1,$2,'integration','sk-test',$3) RETURNING id::text`, tenantID, projectID, keyHash[:]).Scan(&keyID); err != nil {
+ t.Fatal(err)
+ }
+ requestID := fmt.Sprintf("req_settle_%d", time.Now().UnixNano())
+ missingRequestID := requestID + "_missing_usage"
+ t.Cleanup(func() {
+ for _, statement := range []struct {
+ query string
+ arg string
+ }{
+ {`DELETE FROM billing_ledger WHERE tenant_id=$1`, tenantID},
+ {`DELETE FROM usage_monthly_rollups WHERE tenant_id=$1`, tenantID},
+ {`DELETE FROM usage_events WHERE tenant_id=$1`, tenantID},
+ {`DELETE FROM billing_settlement_jobs WHERE request_id=$1`, requestID},
+ {`DELETE FROM billing_settlement_jobs WHERE request_id=$1`, missingRequestID},
+ {`DELETE FROM billing_reservations WHERE request_id=$1`, requestID},
+ {`DELETE FROM billing_reservations WHERE request_id=$1`, missingRequestID},
+ {`DELETE FROM tenant_wallets WHERE tenant_id=$1`, tenantID},
+ {`DELETE FROM api_keys WHERE id=$1`, keyID},
+ {`DELETE FROM projects WHERE id=$1`, projectID},
+ {`DELETE FROM tenants WHERE id=$1`, tenantID},
+ } {
+ if _, cleanupErr := service.db.Exec(context.Background(), statement.query, statement.arg); cleanupErr != nil {
+ t.Errorf("cleanup settlement integration data: %v", cleanupErr)
+ }
+ }
+ })
+ if _, err := service.db.Exec(ctx, `INSERT INTO tenant_wallets (tenant_id,currency,balance_micros,reserved_micros) VALUES ($1,'usd',1000,20)`, tenantID); err != nil {
+ t.Fatal(err)
+ }
+ if _, err := service.db.Exec(ctx, `INSERT INTO billing_reservations (request_id,tenant_id,project_id,key_id,public_model,currency,reserved_micros,input_price_micros_per_million,output_price_micros_per_million,cache_read_price_micros_per_million,cache_write_price_micros_per_million) VALUES ($1,$2,$3,$4,'demo/model','usd',20,1000000,1000000,0,0)`, requestID, tenantID, projectID, keyID); err != nil {
+ t.Fatal(err)
+ }
+ if _, err := service.db.Exec(ctx, `INSERT INTO billing_settlement_jobs (request_id) VALUES ($1)`, requestID); err != nil {
+ t.Fatal(err)
+ }
+ if err := service.EnqueueSettlement(ctx, domain.UsageEvent{RequestID: requestID, TenantID: tenantID, ProjectID: projectID, KeyID: keyID, PublicModel: "demo/model", Protocol: domain.ProtocolOpenAI, StatusCode: 200, Success: true, UsageReported: true, StartedAt: time.Now().UTC(), Usage: domain.Usage{InputTokens: 10}}); err != nil {
+ t.Fatal(err)
+ }
+ processed, err := service.processSettlementJob(ctx)
+ if err != nil || !processed {
+ t.Fatalf("process settlement: processed=%v err=%v", processed, err)
+ }
+ var reservationStatus, jobStatus string
+ var balance, reserved, charged, usageCount int64
+ if err := service.db.QueryRow(ctx, `SELECT status,charged_micros FROM billing_reservations WHERE request_id=$1`, requestID).Scan(&reservationStatus, &charged); err != nil {
+ t.Fatal(err)
+ }
+ if err := service.db.QueryRow(ctx, `SELECT status FROM billing_settlement_jobs WHERE request_id=$1`, requestID).Scan(&jobStatus); err != nil {
+ t.Fatal(err)
+ }
+ if err := service.db.QueryRow(ctx, `SELECT balance_micros,reserved_micros FROM tenant_wallets WHERE tenant_id=$1`, tenantID).Scan(&balance, &reserved); err != nil {
+ t.Fatal(err)
+ }
+ if err := service.db.QueryRow(ctx, `SELECT count(*) FROM usage_events WHERE request_id=$1`, requestID).Scan(&usageCount); err != nil {
+ t.Fatal(err)
+ }
+ if reservationStatus != "settled" || jobStatus != "done" || balance != 990 || reserved != 0 || charged != 10 || usageCount != 1 {
+ t.Fatalf("reservation=%s job=%s balance=%d reserved=%d charged=%d usage=%d", reservationStatus, jobStatus, balance, reserved, charged, usageCount)
+ }
+
+ if _, err := service.db.Exec(ctx, `UPDATE tenant_wallets SET reserved_micros=reserved_micros+30 WHERE tenant_id=$1`, tenantID); err != nil {
+ t.Fatal(err)
+ }
+ if _, err := service.db.Exec(ctx, `INSERT INTO billing_reservations
+ (request_id,tenant_id,project_id,key_id,public_model,currency,reserved_micros,input_price_micros_per_million,output_price_micros_per_million,cache_read_price_micros_per_million,cache_write_price_micros_per_million)
+ VALUES ($1,$2,$3,$4,'demo/model','usd',30,1000000,1000000,0,0)`, missingRequestID, tenantID, projectID, keyID); err != nil {
+ t.Fatal(err)
+ }
+ if _, err := service.db.Exec(ctx, `INSERT INTO billing_settlement_jobs (request_id) VALUES ($1)`, missingRequestID); err != nil {
+ t.Fatal(err)
+ }
+ if err := service.EnqueueSettlement(ctx, domain.UsageEvent{RequestID: missingRequestID, TenantID: tenantID, ProjectID: projectID,
+ KeyID: keyID, PublicModel: "demo/model", Protocol: domain.ProtocolOpenAI, StatusCode: 200, Success: true,
+ UsageReported: false, StartedAt: time.Now().UTC()}); err != nil {
+ t.Fatal(err)
+ }
+ processed, err = service.processSettlementJob(ctx)
+ if err != nil || !processed {
+ t.Fatalf("process missing-usage settlement: processed=%v err=%v", processed, err)
+ }
+ var missingReservationStatus, missingJobStatus, meteringStatus string
+ if err := service.db.QueryRow(ctx, `SELECT status FROM billing_reservations WHERE request_id=$1`, missingRequestID).Scan(&missingReservationStatus); err != nil {
+ t.Fatal(err)
+ }
+ if err := service.db.QueryRow(ctx, `SELECT status FROM billing_settlement_jobs WHERE request_id=$1`, missingRequestID).Scan(&missingJobStatus); err != nil {
+ t.Fatal(err)
+ }
+ if err := service.db.QueryRow(ctx, `SELECT balance_micros,reserved_micros FROM tenant_wallets WHERE tenant_id=$1`, tenantID).Scan(&balance, &reserved); err != nil {
+ t.Fatal(err)
+ }
+ if err := service.db.QueryRow(ctx, `SELECT metering_status FROM usage_events WHERE request_id=$1`, missingRequestID).Scan(&meteringStatus); err != nil {
+ t.Fatal(err)
+ }
+ if missingReservationStatus != "metering_failed" || missingJobStatus != "done" || balance != 990 || reserved != 30 || meteringStatus != "missing" {
+ t.Fatalf("missing usage reservation=%s job=%s balance=%d reserved=%d metering=%s", missingReservationStatus, missingJobStatus, balance, reserved, meteringStatus)
}
}
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)
+}
diff --git a/internal/billing/types.go b/internal/billing/types.go
index 3d0461e..fe5df1e 100644
--- a/internal/billing/types.go
+++ b/internal/billing/types.go
@@ -14,11 +14,13 @@ var (
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")
)
type Meter interface {
Authorize(context.Context, Authorization) error
- Settle(context.Context, domain.UsageEvent) error
+ EnqueueSettlement(context.Context, domain.UsageEvent) error
}
type Authorization struct {
@@ -40,6 +42,54 @@ type Options struct {
StripeWebhookSecret string
StripeSuccessURL string
StripeCancelURL string
+ StripePortalReturnURL string
+ StripeAutomaticTax bool
+ StripeProductTaxCode string
+ SettlementSpoolPath string
+ Metrics OperationalMetrics
+}
+
+type OperationalMetrics interface {
+ SetStripeOperations(OperationalStatus)
+}
+
+type OperationalStatus struct {
+ StripeEnabled bool `json:"stripe_enabled"`
+ ReconciliationStatus string `json:"reconciliation_status"`
+ ReconciliationMismatches int64 `json:"reconciliation_mismatches"`
+ ReconciliationCompletedAt *time.Time `json:"reconciliation_completed_at,omitempty"`
+ ReconciliationError string `json:"reconciliation_error,omitempty"`
+ UnprocessedWebhooks int64 `json:"unprocessed_webhooks"`
+ OldestUnprocessedWebhook *time.Time `json:"oldest_unprocessed_webhook,omitempty"`
+ RefundBacklog int64 `json:"refund_backlog"`
+ OldestRefund *time.Time `json:"oldest_refund,omitempty"`
+ UncollectedMicros int64 `json:"uncollected_micros"`
+ UnmeteredSuccesses int64 `json:"unmetered_successes"`
+}
+
+func (s OperationalStatus) Ready(now time.Time) bool {
+ if !s.StripeEnabled {
+ return s.UnmeteredSuccesses == 0
+ }
+ if s.ReconciliationCompletedAt == nil || now.Sub(*s.ReconciliationCompletedAt) > 2*time.Hour {
+ return false
+ }
+ webhooksStuck := s.UnprocessedWebhooks > 0 &&
+ (s.OldestUnprocessedWebhook == nil || now.Sub(*s.OldestUnprocessedWebhook) > 5*time.Minute)
+ refundsStuck := s.RefundBacklog > 0 &&
+ (s.OldestRefund == nil || now.Sub(*s.OldestRefund) > 15*time.Minute)
+ return s.ReconciliationStatus == "clean" && s.ReconciliationMismatches == 0 &&
+ s.ReconciliationError == "" && !webhooksStuck && !refundsStuck &&
+ s.UncollectedMicros == 0 && s.UnmeteredSuccesses == 0
+}
+
+type SettlementQueueStatus struct {
+ AwaitingEvent int64 `json:"awaiting_event"`
+ Pending int64 `json:"pending"`
+ Processing int64 `json:"processing"`
+ Retrying int64 `json:"retrying"`
+ OldestPending *time.Time `json:"oldest_pending,omitempty"`
+ SpoolRecords int `json:"spool_records"`
}
type Account struct {
@@ -73,8 +123,9 @@ type AdjustmentInput struct {
}
type CheckoutInput struct {
- TenantID string `json:"tenant_id"`
- AmountMinor int64 `json:"amount_minor"`
+ TenantID string `json:"tenant_id"`
+ AmountMinor int64 `json:"amount_minor"`
+ CustomerEmail string `json:"-"`
}
type CheckoutResult struct {
@@ -84,14 +135,99 @@ type CheckoutResult struct {
}
type TopUpOrder struct {
- ID string `json:"id"`
- TenantID string `json:"tenant_id"`
- AmountMinor int64 `json:"amount_minor"`
- AmountMicros int64 `json:"amount_micros"`
- Currency string `json:"currency"`
- Status string `json:"status"`
- StripeSessionID string `json:"stripe_session_id,omitempty"`
- CheckoutURL string `json:"checkout_url,omitempty"`
- CreatedAt time.Time `json:"created_at"`
- PaidAt *time.Time `json:"paid_at,omitempty"`
+ ID string `json:"id"`
+ TenantID string `json:"tenant_id"`
+ AmountMinor int64 `json:"amount_minor"`
+ AmountMicros int64 `json:"amount_micros"`
+ Currency string `json:"currency"`
+ Status string `json:"status"`
+ StripeSessionID string `json:"stripe_session_id,omitempty"`
+ CheckoutURL string `json:"checkout_url,omitempty"`
+ CreatedAt time.Time `json:"created_at"`
+ PaidAt *time.Time `json:"paid_at,omitempty"`
+ StripeCustomerID string `json:"stripe_customer_id,omitempty"`
+ StripePaymentIntentID string `json:"stripe_payment_intent_id,omitempty"`
+ StripeChargeID string `json:"stripe_charge_id,omitempty"`
+ StripeInvoiceID string `json:"stripe_invoice_id,omitempty"`
+ InvoiceURL string `json:"invoice_url,omitempty"`
+ InvoicePDFURL string `json:"invoice_pdf_url,omitempty"`
+ ReceiptURL string `json:"receipt_url,omitempty"`
+ RefundedMicros int64 `json:"refunded_micros"`
+ DisputedMicros int64 `json:"disputed_micros"`
+ ReconciliationStatus string `json:"reconciliation_status"`
+ ReconciledAt *time.Time `json:"reconciled_at,omitempty"`
+ ReconciliationError string `json:"reconciliation_error,omitempty"`
+}
+
+type ResolveMissingTopUpInput struct {
+ Reason string `json:"reason"`
+}
+
+type ResolutionActor struct {
+ ID string
+ Type string
+}
+
+type RefundInput struct {
+ AmountMinor int64 `json:"amount_minor"`
+ Reason string `json:"reason"`
+}
+
+type Refund struct {
+ ID string `json:"id"`
+ TenantID string `json:"tenant_id"`
+ TopUpOrderID string `json:"topup_order_id"`
+ StripeRefundID string `json:"stripe_refund_id,omitempty"`
+ AmountMinor int64 `json:"amount_minor"`
+ AmountMicros int64 `json:"amount_micros"`
+ Currency string `json:"currency"`
+ Reason string `json:"reason"`
+ Status string `json:"status"`
+ LastError string `json:"last_error,omitempty"`
+ CreatedAt time.Time `json:"created_at"`
+ CompletedAt *time.Time `json:"completed_at,omitempty"`
+}
+
+type PortalResult struct {
+ URL string `json:"url"`
+}
+
+type ReconciliationResult struct {
+ ID string `json:"id"`
+ Status string `json:"status"`
+ CheckedOrders int64 `json:"checked_orders"`
+ MismatchCount int64 `json:"mismatch_count"`
+ Mismatches []map[string]any `json:"mismatches"`
+ Repairs []map[string]any `json:"repairs,omitempty"`
+}
+
+type PaymentDispute struct {
+ ID string `json:"id"`
+ TenantID string `json:"tenant_id,omitempty"`
+ TopUpOrderID string `json:"topup_order_id,omitempty"`
+ AmountMinor int64 `json:"amount_minor"`
+ AmountMicros int64 `json:"amount_micros"`
+ Currency string `json:"currency"`
+ Status string `json:"status"`
+ Reason string `json:"reason"`
+ DebitedMicros int64 `json:"debited_micros"`
+ UncollectedMicros int64 `json:"uncollected_micros"`
+ DueBy *time.Time `json:"due_by,omitempty"`
+ UpdatedAt time.Time `json:"updated_at"`
+}
+
+type Invoice struct {
+ ID string `json:"id"`
+ TenantID string `json:"tenant_id,omitempty"`
+ TopUpOrderID string `json:"topup_order_id,omitempty"`
+ Status string `json:"status"`
+ Currency string `json:"currency"`
+ AmountDueMinor int64 `json:"amount_due_minor"`
+ AmountPaidMinor int64 `json:"amount_paid_minor"`
+ AttemptCount int `json:"attempt_count"`
+ NextPaymentAttempt *time.Time `json:"next_payment_attempt,omitempty"`
+ HostedInvoiceURL string `json:"hosted_invoice_url,omitempty"`
+ InvoicePDFURL string `json:"invoice_pdf_url,omitempty"`
+ LastFailure string `json:"last_failure,omitempty"`
+ UpdatedAt time.Time `json:"updated_at"`
}