summaryrefslogtreecommitdiff
path: root/internal/billing/operations.go
diff options
context:
space:
mode:
Diffstat (limited to '')
-rw-r--r--internal/billing/operations.go8
1 files changed, 5 insertions, 3 deletions
diff --git a/internal/billing/operations.go b/internal/billing/operations.go
index 461c59a..508ebd5 100644
--- a/internal/billing/operations.go
+++ b/internal/billing/operations.go
@@ -15,8 +15,10 @@ import (
"github.com/stripe/stripe-go/v86"
)
+type stripePortalSessionCreator func(context.Context, *stripe.BillingPortalSessionCreateParams) (*stripe.BillingPortalSession, error)
+
func (s *Service) CreatePortalSession(ctx context.Context, tenantID string) (PortalResult, error) {
- if !s.stripeEnabled || s.stripeClient == nil {
+ if !s.stripeEnabled || s.createStripePortalSession == nil {
return PortalResult{}, ErrStripeDisabled
}
customerID, err := s.ensureStripeCustomer(ctx, tenantID)
@@ -26,13 +28,13 @@ func (s *Service) CreatePortalSession(ctx context.Context, tenantID string) (Por
if customerID == "" {
return PortalResult{}, errors.New("no Stripe customer exists for this account")
}
- session, err := s.stripeClient.V1BillingPortalSessions.Create(ctx, &stripe.BillingPortalSessionCreateParams{
+ session, err := s.createStripePortalSession(ctx, &stripe.BillingPortalSessionCreateParams{
Customer: stripe.String(customerID), ReturnURL: stripe.String(s.stripePortalReturnURL),
})
if err != nil {
return PortalResult{}, fmt.Errorf("create Stripe customer portal session: %w", err)
}
- if session.URL == "" {
+ if session == nil || session.URL == "" {
return PortalResult{}, errors.New("Stripe returned an incomplete portal session")
}
return PortalResult{URL: session.URL}, nil