diff options
Diffstat (limited to 'internal/billing/operations.go')
| -rw-r--r-- | internal/billing/operations.go | 8 |
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 |
