summaryrefslogtreecommitdiff
path: root/internal/billing/profile.go
blob: 3b08a3c04451f4a2b136f6f6d914f61b29302416 (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
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
}