summaryrefslogtreecommitdiff
path: root/internal/billing/profile.go
diff options
context:
space:
mode:
Diffstat (limited to 'internal/billing/profile.go')
-rw-r--r--internal/billing/profile.go197
1 files changed, 197 insertions, 0 deletions
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
+}