From 41e322c53d7b4b796eb377d0df9c29ecd10ba431 Mon Sep 17 00:00:00 2001 From: Chia Date: Thu, 6 Aug 2026 09:29:41 +1200 Subject: feat: complete commercial control plane, billing, auth, and model catalog - add PostgreSQL control-plane persistence with Redis-degraded hot reload - implement prepaid balance, usage ledger, Stripe top-up and reconciliation - add registration, email verification, password reset, invitations and RBAC - support TOTP, Passkey MFA, device sessions, quotas and rate limits - add tenant billing profiles, audit logs and operational readiness checks - build authenticated admin console, Quickstart, Playground and usage analytics - add public model catalog with pricing, filtering and cost estimation - support OpenAI Responses providers and provider health failover - validate real upstream usage reporting and balance settlement --- internal/billing/profile.go | 197 ++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 197 insertions(+) create mode 100644 internal/billing/profile.go (limited to 'internal/billing/profile.go') diff --git a/internal/billing/profile.go b/internal/billing/profile.go new file mode 100644 index 0000000..3b08a3c --- /dev/null +++ b/internal/billing/profile.go @@ -0,0 +1,197 @@ +package billing + +import ( + "context" + "errors" + "fmt" + "net/mail" + "strings" + "unicode" + + "github.com/jackc/pgx/v5" + "github.com/stripe/stripe-go/v86" +) + +type stripeCustomerCreator func(context.Context, *stripe.CustomerCreateParams) (*stripe.Customer, error) +type stripeCustomerUpdater func(context.Context, string, *stripe.CustomerUpdateParams) (*stripe.Customer, error) + +func (s *Service) GetBillingProfile(ctx context.Context, tenantID string) (BillingProfile, error) { + tenantID = strings.TrimSpace(tenantID) + if tenantID == "" { + return BillingProfile{}, ErrBillingAccountNotFound + } + var result BillingProfile + err := s.db.QueryRow(ctx, `SELECT t.id::text, + COALESCE(p.legal_name,t.name),COALESCE(p.billing_email,sc.email,''), + COALESCE(p.address_line1,''),COALESCE(p.address_line2,''),COALESCE(p.city,''), + COALESCE(p.region,''),COALESCE(p.postal_code,''),COALESCE(p.country,''), + p.tenant_id IS NOT NULL,sc.stripe_customer_id IS NOT NULL, + COALESCE(p.stripe_sync_status,CASE WHEN sc.stripe_customer_id IS NOT NULL THEN 'checkout_managed' ELSE 'not_configured' END), + p.stripe_synced_at,COALESCE(p.stripe_sync_error,''),p.updated_at + FROM tenants t + LEFT JOIN tenant_billing_profiles p ON p.tenant_id=t.id + LEFT JOIN stripe_customers sc ON sc.tenant_id=t.id + WHERE t.id=$1`, tenantID).Scan(&result.TenantID, &result.LegalName, &result.BillingEmail, + &result.AddressLine1, &result.AddressLine2, &result.City, &result.Region, &result.PostalCode, + &result.Country, &result.Configured, &result.StripeCustomerConfigured, &result.StripeSyncStatus, + &result.StripeSyncedAt, &result.StripeSyncError, &result.UpdatedAt) + if errors.Is(err, pgx.ErrNoRows) { + return BillingProfile{}, ErrBillingAccountNotFound + } + if err != nil { + return BillingProfile{}, fmt.Errorf("get billing profile: %w", err) + } + return result, nil +} + +func (s *Service) UpdateBillingProfile(ctx context.Context, input UpdateBillingProfileInput) (BillingProfile, error) { + normalized, err := normalizeBillingProfile(input) + if err != nil { + return BillingProfile{}, err + } + status := "disabled" + if s.stripeEnabled { + status = "pending" + } + _, err = s.db.Exec(ctx, `INSERT INTO tenant_billing_profiles + (tenant_id,legal_name,billing_email,address_line1,address_line2,city,region,postal_code,country,stripe_sync_status,stripe_synced_at,stripe_sync_error) + VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,NULL,'') + ON CONFLICT (tenant_id) DO UPDATE SET legal_name=EXCLUDED.legal_name, + billing_email=EXCLUDED.billing_email,address_line1=EXCLUDED.address_line1,address_line2=EXCLUDED.address_line2, + city=EXCLUDED.city,region=EXCLUDED.region,postal_code=EXCLUDED.postal_code,country=EXCLUDED.country, + stripe_sync_status=EXCLUDED.stripe_sync_status,stripe_synced_at=NULL,stripe_sync_error='',updated_at=now()`, + normalized.TenantID, normalized.LegalName, normalized.BillingEmail, normalized.AddressLine1, + normalized.AddressLine2, normalized.City, normalized.Region, normalized.PostalCode, normalized.Country, status) + if err != nil { + return BillingProfile{}, fmt.Errorf("save billing profile: %w", err) + } + if !s.stripeEnabled { + return s.GetBillingProfile(ctx, normalized.TenantID) + } + if _, err := s.ensureStripeCustomer(ctx, normalized.TenantID); err != nil { + return BillingProfile{}, err + } + return s.GetBillingProfile(ctx, normalized.TenantID) +} + +func (s *Service) ensureStripeCustomer(ctx context.Context, tenantID string) (string, error) { + var profile BillingProfile + err := s.db.QueryRow(ctx, `SELECT tenant_id::text,legal_name,billing_email,address_line1,address_line2, + city,region,postal_code,country,TRUE FROM tenant_billing_profiles WHERE tenant_id=$1`, tenantID). + Scan(&profile.TenantID, &profile.LegalName, &profile.BillingEmail, &profile.AddressLine1, + &profile.AddressLine2, &profile.City, &profile.Region, &profile.PostalCode, &profile.Country, &profile.Configured) + if errors.Is(err, pgx.ErrNoRows) { + var customerID string + err = s.db.QueryRow(ctx, `SELECT stripe_customer_id FROM stripe_customers WHERE tenant_id=$1`, tenantID).Scan(&customerID) + if errors.Is(err, pgx.ErrNoRows) { + return "", nil + } + return customerID, err + } + if err != nil { + return "", fmt.Errorf("load billing profile for Stripe: %w", err) + } + if !s.stripeEnabled || s.createStripeCustomer == nil || s.updateStripeCustomer == nil { + return "", ErrStripeDisabled + } + var customerID string + err = s.db.QueryRow(ctx, `SELECT stripe_customer_id FROM stripe_customers WHERE tenant_id=$1`, tenantID).Scan(&customerID) + if err != nil && !errors.Is(err, pgx.ErrNoRows) { + return "", fmt.Errorf("load Stripe customer: %w", err) + } + if customerID == "" { + params := billingProfileCustomerCreateParams(profile) + params.SetIdempotencyKey("aigw_customer_" + tenantID) + customer, createErr := s.createStripeCustomer(ctx, params) + if createErr != nil || customer == nil || strings.TrimSpace(customer.ID) == "" { + if createErr == nil { + createErr = errors.New("Stripe returned an incomplete Customer") + } + return "", s.failBillingProfileSync(ctx, tenantID, createErr) + } + customerID = customer.ID + } else { + customer, updateErr := s.updateStripeCustomer(ctx, customerID, billingProfileCustomerUpdateParams(profile)) + if updateErr != nil || customer == nil || strings.TrimSpace(customer.ID) == "" { + if updateErr == nil { + updateErr = errors.New("Stripe returned an incomplete Customer") + } + return "", s.failBillingProfileSync(ctx, tenantID, updateErr) + } + } + if _, err := s.db.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=EXCLUDED.email,updated_at=now()`, + tenantID, customerID, profile.BillingEmail); err != nil { + return "", fmt.Errorf("persist Stripe customer: %w", err) + } + if _, err := s.db.Exec(ctx, `UPDATE tenant_billing_profiles SET stripe_sync_status='synced',stripe_synced_at=now(), + stripe_sync_error='',updated_at=now() WHERE tenant_id=$1`, tenantID); err != nil { + return "", fmt.Errorf("record billing profile Stripe synchronization: %w", err) + } + return customerID, nil +} + +func (s *Service) failBillingProfileSync(ctx context.Context, tenantID string, cause error) error { + _, _ = s.db.Exec(ctx, `UPDATE tenant_billing_profiles SET stripe_sync_status='failed', + stripe_sync_error='Stripe customer synchronization failed',stripe_synced_at=NULL,updated_at=now() WHERE tenant_id=$1`, tenantID) + return fmt.Errorf("%w: %v", ErrBillingProfileSync, cause) +} + +func billingProfileCustomerCreateParams(profile BillingProfile) *stripe.CustomerCreateParams { + return &stripe.CustomerCreateParams{ + Name: stripe.String(profile.LegalName), BusinessName: stripe.String(profile.LegalName), + Email: stripe.String(profile.BillingEmail), Address: billingProfileAddress(profile), + Metadata: map[string]string{"aigw_tenant_id": profile.TenantID}, + } +} + +func billingProfileCustomerUpdateParams(profile BillingProfile) *stripe.CustomerUpdateParams { + return &stripe.CustomerUpdateParams{ + Name: stripe.String(profile.LegalName), BusinessName: stripe.String(profile.LegalName), + Email: stripe.String(profile.BillingEmail), Address: billingProfileAddress(profile), + Metadata: map[string]string{"aigw_tenant_id": profile.TenantID}, + } +} + +func billingProfileAddress(profile BillingProfile) *stripe.AddressParams { + return &stripe.AddressParams{Line1: stripe.String(profile.AddressLine1), Line2: stripe.String(profile.AddressLine2), + City: stripe.String(profile.City), State: stripe.String(profile.Region), PostalCode: stripe.String(profile.PostalCode), + Country: stripe.String(profile.Country)} +} + +func normalizeBillingProfile(input UpdateBillingProfileInput) (UpdateBillingProfileInput, error) { + input.TenantID = strings.TrimSpace(input.TenantID) + if input.TenantID == "" { + return UpdateBillingProfileInput{}, fmt.Errorf("%w: tenant is required", ErrInvalidBillingProfile) + } + var err error + for _, field := range []struct { + value *string + name string + max int + required bool + }{ + {&input.LegalName, "legal name", 150, true}, {&input.BillingEmail, "billing email", 254, true}, + {&input.AddressLine1, "address line 1", 200, true}, {&input.AddressLine2, "address line 2", 200, false}, + {&input.City, "city", 100, true}, {&input.Region, "state or region", 100, false}, + {&input.PostalCode, "postal code", 32, true}, + } { + *field.value = strings.TrimSpace(*field.value) + if field.required && *field.value == "" { + return UpdateBillingProfileInput{}, fmt.Errorf("%w: %s is required", ErrInvalidBillingProfile, field.name) + } + if len([]rune(*field.value)) > field.max || strings.IndexFunc(*field.value, unicode.IsControl) >= 0 { + return UpdateBillingProfileInput{}, fmt.Errorf("%w: %s is invalid", ErrInvalidBillingProfile, field.name) + } + } + parsedEmail, err := mail.ParseAddress(input.BillingEmail) + if err != nil || !strings.EqualFold(parsedEmail.Address, input.BillingEmail) { + return UpdateBillingProfileInput{}, fmt.Errorf("%w: billing email is invalid", ErrInvalidBillingProfile) + } + input.BillingEmail = strings.ToLower(parsedEmail.Address) + input.Country = strings.ToUpper(strings.TrimSpace(input.Country)) + if len(input.Country) != 2 || input.Country[0] < 'A' || input.Country[0] > 'Z' || input.Country[1] < 'A' || input.Country[1] > 'Z' { + return UpdateBillingProfileInput{}, fmt.Errorf("%w: country must be a two-letter code", ErrInvalidBillingProfile) + } + return input, nil +} -- cgit v1.2.3