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 }