diff options
Diffstat (limited to '')
| -rw-r--r-- | internal/billing/ledger.go | 41 | ||||
| -rw-r--r-- | internal/billing/operations.go | 881 | ||||
| -rw-r--r-- | internal/billing/operations_test.go | 36 | ||||
| -rw-r--r-- | internal/billing/service.go | 319 | ||||
| -rw-r--r-- | internal/billing/service_test.go | 174 | ||||
| -rw-r--r-- | internal/billing/stripe.go | 207 | ||||
| -rw-r--r-- | internal/billing/types.go | 162 |
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, ¤cy, &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, ¤cy, &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, ¤tStatus) + 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, ¤cy); 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, ¤cy); 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"` } |
