//go:build billing package billing import ( "encoding/json" "errors" "fmt" "io" "log/slog" "net/http" "os" "sort" "strings" "sync" "time" "github.com/stripe/stripe-go/v84" portalsession "github.com/stripe/stripe-go/v84/billingportal/session" "github.com/stripe/stripe-go/v84/checkout/session" "github.com/stripe/stripe-go/v84/customer" "github.com/stripe/stripe-go/v84/price" "github.com/stripe/stripe-go/v84/subscription" "github.com/stripe/stripe-go/v84/webhook" "atcr.io/pkg/hold/quota" ) // Manager handles Stripe billing integration. type Manager struct { quotaMgr *quota.Manager billingCfg *BillingConfig holdPublicURL string stripeKey string webhookSecret string publishableKey string // In-memory cache for customer lookups (DID -> customer) customerCache map[string]*cachedCustomer customerCacheMu sync.RWMutex } type cachedCustomer struct { customer *stripe.Customer expiresAt time.Time } const customerCacheTTL = 10 * time.Minute // New creates a new billing manager with Stripe integration. // configPath is the path to the hold config YAML file (for billing config parsing). func New(quotaMgr *quota.Manager, holdPublicURL string, configPath string) *Manager { stripeKey := os.Getenv("STRIPE_SECRET_KEY") if stripeKey != "" { stripe.Key = stripeKey } billingCfg, err := LoadBillingConfig(configPath) if err != nil { slog.Warn("Failed to load billing config", "error", err) } // Validate billing tier names against quota tiers if billingCfg != nil && billingCfg.Enabled { for tierName := range billingCfg.Tiers { if quotaMgr.GetTierLimit(tierName) == nil && tierName != quotaMgr.GetDefaultTier() { slog.Warn("Billing tier has no matching quota tier", "tier", tierName) } } } return &Manager{ quotaMgr: quotaMgr, billingCfg: billingCfg, holdPublicURL: holdPublicURL, stripeKey: stripeKey, webhookSecret: os.Getenv("STRIPE_WEBHOOK_SECRET"), publishableKey: os.Getenv("STRIPE_PUBLISHABLE_KEY"), customerCache: make(map[string]*cachedCustomer), } } // Enabled returns true if billing is properly configured. func (m *Manager) Enabled() bool { return m.billingCfg != nil && m.billingCfg.Enabled && m.stripeKey != "" } // GetSubscriptionInfo returns subscription and quota information for a user. func (m *Manager) GetSubscriptionInfo(userDID string) (*SubscriptionInfo, error) { if !m.Enabled() { return nil, ErrBillingDisabled } info := &SubscriptionInfo{ UserDID: userDID, PaymentsEnabled: true, Tiers: m.buildTierList(userDID), } // Try to find existing customer cust, err := m.findCustomerByDID(userDID) if err != nil { slog.Debug("No Stripe customer found for user", "userDid", userDID) } else if cust != nil { info.CustomerID = cust.ID // Get active subscription if any (check all nil pointers) if cust.Subscriptions != nil && len(cust.Subscriptions.Data) > 0 { sub := cust.Subscriptions.Data[0] info.SubscriptionID = sub.ID // Safely access subscription items if sub.Items != nil && len(sub.Items.Data) > 0 && sub.Items.Data[0].Price != nil { info.CurrentTier = m.billingCfg.GetTierByPriceID(sub.Items.Data[0].Price.ID) if sub.Items.Data[0].Price.Recurring != nil { switch sub.Items.Data[0].Price.Recurring.Interval { case stripe.PriceRecurringIntervalMonth: info.BillingInterval = "monthly" case stripe.PriceRecurringIntervalYear: info.BillingInterval = "yearly" } } } } } // If no subscription, use default tier if info.CurrentTier == "" { info.CurrentTier = m.quotaMgr.GetDefaultTier() } // Get quota limit for current tier limit := m.quotaMgr.GetTierLimit(info.CurrentTier) info.CurrentLimit = limit // Mark current tier in tier list for i := range info.Tiers { if info.Tiers[i].ID == info.CurrentTier { info.Tiers[i].IsCurrent = true } } return info, nil } // buildTierList creates the list of available tiers by merging quota limits // from the quota manager with billing metadata from the billing config. func (m *Manager) buildTierList(userDID string) []TierInfo { quotaTiers := m.quotaMgr.ListTiers() if len(quotaTiers) == 0 { return nil } result := make([]TierInfo, 0, len(quotaTiers)) for _, qt := range quotaTiers { var quotaBytes int64 if qt.Limit != nil { quotaBytes = *qt.Limit } // Capitalize tier ID for display name (e.g., "swabbie" -> "Swabbie") name := strings.ToUpper(qt.Key[:1]) + qt.Key[1:] tier := TierInfo{ ID: qt.Key, Name: name, QuotaBytes: quotaBytes, QuotaFormatted: quota.FormatHumanBytes(quotaBytes), } // Merge billing metadata if available if bt := m.billingCfg.GetTierPricing(qt.Key); bt != nil { tier.Description = bt.Description // Fetch actual prices from Stripe if bt.StripePriceMonthly != "" { if p, err := price.Get(bt.StripePriceMonthly, nil); err == nil && p != nil { tier.PriceCentsMonthly = int(p.UnitAmount) } else { slog.Debug("Failed to fetch monthly price", "priceId", bt.StripePriceMonthly, "error", err) tier.PriceCentsMonthly = -1 } } if bt.StripePriceYearly != "" { if p, err := price.Get(bt.StripePriceYearly, nil); err == nil && p != nil { tier.PriceCentsYearly = int(p.UnitAmount) } else { slog.Debug("Failed to fetch yearly price", "priceId", bt.StripePriceYearly, "error", err) tier.PriceCentsYearly = -1 } } } result = append(result, tier) } // Sort tiers by quota size (ascending) sort.Slice(result, func(i, j int) bool { return result[i].QuotaBytes < result[j].QuotaBytes }) return result } // CreateCheckoutSession creates a Stripe checkout session for subscription. func (m *Manager) CreateCheckoutSession(r *http.Request, req *CheckoutSessionRequest) (*CheckoutSessionResponse, error) { if !m.Enabled() { return nil, ErrBillingDisabled } // Get user DID from request context (set by auth middleware) userDID := r.Header.Get("X-User-DID") if userDID == "" { return nil, errors.New("user not authenticated") } // Get tier config tierCfg := m.billingCfg.GetTierPricing(req.Tier) if tierCfg == nil { return nil, fmt.Errorf("tier not found: %s", req.Tier) } // Determine price ID - prefer requested interval, fall back to what's available var priceID string switch req.Interval { case "monthly": priceID = tierCfg.StripePriceMonthly case "yearly": priceID = tierCfg.StripePriceYearly default: // No interval specified - prefer monthly, fall back to yearly if tierCfg.StripePriceMonthly != "" { priceID = tierCfg.StripePriceMonthly } else { priceID = tierCfg.StripePriceYearly } } if priceID == "" { return nil, fmt.Errorf("tier %s has no Stripe price configured", req.Tier) } // Get or create customer cust, err := m.getOrCreateCustomer(userDID) if err != nil { return nil, fmt.Errorf("failed to get/create customer: %w", err) } // Build success/cancel URLs successURL := strings.ReplaceAll(m.billingCfg.SuccessURL, "{hold_url}", m.holdPublicURL) cancelURL := strings.ReplaceAll(m.billingCfg.CancelURL, "{hold_url}", m.holdPublicURL) if req.ReturnURL != "" { successURL = req.ReturnURL + "?success=true" cancelURL = req.ReturnURL + "?cancelled=true" } // Create checkout session params := &stripe.CheckoutSessionParams{ Customer: stripe.String(cust.ID), Mode: stripe.String(string(stripe.CheckoutSessionModeSubscription)), LineItems: []*stripe.CheckoutSessionLineItemParams{ { Price: stripe.String(priceID), Quantity: stripe.Int64(1), }, }, SuccessURL: stripe.String(successURL), CancelURL: stripe.String(cancelURL), } sess, err := session.New(params) if err != nil { return nil, fmt.Errorf("failed to create checkout session: %w", err) } return &CheckoutSessionResponse{ CheckoutURL: sess.URL, SessionID: sess.ID, }, nil } // GetBillingPortalURL returns a URL to the Stripe billing portal. func (m *Manager) GetBillingPortalURL(userDID string, returnURL string) (*BillingPortalResponse, error) { if !m.Enabled() { return nil, ErrBillingDisabled } // Find existing customer cust, err := m.findCustomerByDID(userDID) if err != nil || cust == nil { return nil, errors.New("no billing account found") } if returnURL == "" { returnURL = m.holdPublicURL } params := &stripe.BillingPortalSessionParams{ Customer: stripe.String(cust.ID), ReturnURL: stripe.String(returnURL), } sess, err := portalsession.New(params) if err != nil { return nil, fmt.Errorf("failed to create portal session: %w", err) } return &BillingPortalResponse{ PortalURL: sess.URL, }, nil } // HandleWebhook processes a Stripe webhook event. func (m *Manager) HandleWebhook(r *http.Request) (*WebhookEvent, error) { if !m.Enabled() { return nil, ErrBillingDisabled } body, err := io.ReadAll(r.Body) if err != nil { return nil, fmt.Errorf("failed to read request body: %w", err) } // Verify webhook signature event, err := webhook.ConstructEvent(body, r.Header.Get("Stripe-Signature"), m.webhookSecret) if err != nil { return nil, fmt.Errorf("failed to verify webhook signature: %w", err) } result := &WebhookEvent{ Type: string(event.Type), } switch event.Type { case "checkout.session.completed": var sess stripe.CheckoutSession if err := json.Unmarshal(event.Data.Raw, &sess); err != nil { return nil, fmt.Errorf("failed to parse checkout session: %w", err) } result.CustomerID = sess.Customer.ID result.SubscriptionID = sess.Subscription.ID result.Status = "active" // Fetch customer to get DID from metadata result.UserDID = m.getCustomerDID(sess.Customer.ID) // Get subscription to find the price/tier if sess.Subscription != nil && sess.Subscription.ID != "" { if sub, err := m.getSubscription(sess.Subscription.ID); err == nil && sub != nil { if len(sub.Items.Data) > 0 { result.PriceID = sub.Items.Data[0].Price.ID result.NewTier = m.billingCfg.GetTierByPriceID(result.PriceID) } } } if result.UserDID != "" && result.NewTier != "" { slog.Info("Checkout completed", "userDid", result.UserDID, "tier", result.NewTier, "subscriptionId", result.SubscriptionID, ) } case "customer.subscription.created", "customer.subscription.updated": var sub stripe.Subscription if err := json.Unmarshal(event.Data.Raw, &sub); err != nil { return nil, fmt.Errorf("failed to parse subscription: %w", err) } result.SubscriptionID = sub.ID result.CustomerID = sub.Customer.ID result.Status = string(sub.Status) if len(sub.Items.Data) > 0 { result.PriceID = sub.Items.Data[0].Price.ID result.NewTier = m.billingCfg.GetTierByPriceID(result.PriceID) } // Fetch customer to get DID from metadata (webhook doesn't include expanded customer) result.UserDID = m.getCustomerDID(sub.Customer.ID) // If we have user DID and new tier, this signals that crew tier should be updated if result.UserDID != "" && result.NewTier != "" && sub.Status == stripe.SubscriptionStatusActive { slog.Info("Subscription activated", "userDid", result.UserDID, "tier", result.NewTier, "subscriptionId", result.SubscriptionID, ) } case "customer.subscription.deleted", "customer.subscription.paused": var sub stripe.Subscription if err := json.Unmarshal(event.Data.Raw, &sub); err != nil { return nil, fmt.Errorf("failed to parse subscription: %w", err) } result.SubscriptionID = sub.ID result.CustomerID = sub.Customer.ID if event.Type == "customer.subscription.deleted" { result.Status = "cancelled" } else { result.Status = "paused" } // Fetch customer to get DID from metadata result.UserDID = m.getCustomerDID(sub.Customer.ID) // Set tier to default (downgrade on cancellation/pause) result.NewTier = m.quotaMgr.GetDefaultTier() if result.UserDID != "" { slog.Info("Subscription inactive, downgrading to default tier", "userDid", result.UserDID, "tier", result.NewTier, "status", result.Status, ) } case "customer.subscription.resumed": var sub stripe.Subscription if err := json.Unmarshal(event.Data.Raw, &sub); err != nil { return nil, fmt.Errorf("failed to parse subscription: %w", err) } result.SubscriptionID = sub.ID result.CustomerID = sub.Customer.ID result.Status = "active" if len(sub.Items.Data) > 0 { result.PriceID = sub.Items.Data[0].Price.ID result.NewTier = m.billingCfg.GetTierByPriceID(result.PriceID) } // Fetch customer to get DID from metadata result.UserDID = m.getCustomerDID(sub.Customer.ID) if result.UserDID != "" && result.NewTier != "" { slog.Info("Subscription resumed, restoring tier", "userDid", result.UserDID, "tier", result.NewTier, ) } } return result, nil } // getOrCreateCustomer finds or creates a Stripe customer for the given DID. func (m *Manager) getOrCreateCustomer(userDID string) (*stripe.Customer, error) { // Check cache first m.customerCacheMu.RLock() if cached, ok := m.customerCache[userDID]; ok && time.Now().Before(cached.expiresAt) { m.customerCacheMu.RUnlock() return cached.customer, nil } m.customerCacheMu.RUnlock() // Try to find existing customer cust, err := m.findCustomerByDID(userDID) if err == nil && cust != nil { m.cacheCustomer(userDID, cust) return cust, nil } // Create new customer params := &stripe.CustomerParams{ Metadata: map[string]string{ "user_did": userDID, "hold_did": m.holdPublicURL, // Not actually a DID but useful for tracking }, } cust, err = customer.New(params) if err != nil { return nil, fmt.Errorf("failed to create customer: %w", err) } m.cacheCustomer(userDID, cust) return cust, nil } // findCustomerByDID searches Stripe for a customer with the given DID in metadata. func (m *Manager) findCustomerByDID(userDID string) (*stripe.Customer, error) { // Check cache first m.customerCacheMu.RLock() if cached, ok := m.customerCache[userDID]; ok && time.Now().Before(cached.expiresAt) { m.customerCacheMu.RUnlock() return cached.customer, nil } m.customerCacheMu.RUnlock() // Search Stripe by metadata params := &stripe.CustomerSearchParams{ SearchParams: stripe.SearchParams{ Query: fmt.Sprintf("metadata['user_did']:'%s'", userDID), }, } params.AddExpand("data.subscriptions") iter := customer.Search(params) if iter.Next() { cust := iter.Customer() m.cacheCustomer(userDID, cust) return cust, nil } if err := iter.Err(); err != nil { return nil, err } return nil, nil // Not found } // cacheCustomer adds a customer to the in-memory cache. func (m *Manager) cacheCustomer(userDID string, cust *stripe.Customer) { m.customerCacheMu.Lock() defer m.customerCacheMu.Unlock() m.customerCache[userDID] = &cachedCustomer{ customer: cust, expiresAt: time.Now().Add(customerCacheTTL), } } // InvalidateCustomerCache removes a customer from the cache. func (m *Manager) InvalidateCustomerCache(userDID string) { m.customerCacheMu.Lock() defer m.customerCacheMu.Unlock() delete(m.customerCache, userDID) } // getCustomerDID fetches a customer by ID and returns the user_did from metadata. func (m *Manager) getCustomerDID(customerID string) string { if customerID == "" { return "" } cust, err := customer.Get(customerID, nil) if err != nil { slog.Debug("Failed to fetch customer", "customerId", customerID, "error", err) return "" } if cust.Metadata != nil { return cust.Metadata["user_did"] } return "" } // getSubscription fetches a subscription by ID. func (m *Manager) getSubscription(subscriptionID string) (*stripe.Subscription, error) { if subscriptionID == "" { return nil, nil } params := &stripe.SubscriptionParams{} params.AddExpand("items.data.price") return subscription.Get(subscriptionID, params) }