implement core foundation
This commit is contained in:
@@ -0,0 +1,237 @@
|
||||
package billing
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
|
||||
"github.com/jackc/pgx/v5/pgxpool"
|
||||
"github.com/stripe/stripe-go/v81"
|
||||
billingportalsession "github.com/stripe/stripe-go/v81/billingportal/session"
|
||||
checkoutsession "github.com/stripe/stripe-go/v81/checkout/session"
|
||||
stripecustomer "github.com/stripe/stripe-go/v81/customer"
|
||||
stripeprice "github.com/stripe/stripe-go/v81/price"
|
||||
stripesub "github.com/stripe/stripe-go/v81/subscription"
|
||||
)
|
||||
|
||||
type Service interface {
|
||||
GetStatus(ctx context.Context, networkID string) (*Status, error)
|
||||
CreateCheckoutSession(ctx context.Context, p CheckoutParams) (url string, err error)
|
||||
CreatePortalSession(ctx context.Context, networkID string) (url string, err error)
|
||||
// SyncSeats updates the Stripe subscription quantity with proration.
|
||||
// No-op if the network has no active subscription.
|
||||
SyncSeats(ctx context.Context, networkID string, seats int) error
|
||||
HandleWebhook(ctx context.Context, payload []byte, signature string) error
|
||||
}
|
||||
|
||||
var (
|
||||
ErrNoActiveSubscription = errors.New("network has no stripe customer yet")
|
||||
ErrInvalidCadence = errors.New("invalid billing cadence")
|
||||
)
|
||||
|
||||
type serviceImpl struct {
|
||||
cfg Config
|
||||
repo repository
|
||||
priceMonthlyCents int64
|
||||
priceAnnualCents int64
|
||||
}
|
||||
|
||||
func NewService(ctx context.Context, pool *pgxpool.Pool, cfg Config) (Service, error) {
|
||||
if err := cfg.Validate(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
stripe.Key = cfg.SecretKey
|
||||
|
||||
monthly, err := stripeprice.Get(cfg.PriceMonthlyID, &stripe.PriceParams{
|
||||
Params: stripe.Params{Context: ctx},
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("fetch stripe monthly price: %w", err)
|
||||
}
|
||||
annual, err := stripeprice.Get(cfg.PriceAnnualID, &stripe.PriceParams{
|
||||
Params: stripe.Params{Context: ctx},
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("fetch stripe annual price: %w", err)
|
||||
}
|
||||
|
||||
return &serviceImpl{
|
||||
cfg: cfg,
|
||||
repo: newRepository(pool),
|
||||
priceMonthlyCents: monthly.UnitAmount,
|
||||
priceAnnualCents: annual.UnitAmount,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s *serviceImpl) GetStatus(ctx context.Context, networkID string) (*Status, error) {
|
||||
sub, err := s.repo.getSubscriptionByNetworkID(ctx, networkID)
|
||||
if err != nil && !errors.Is(err, errNotFound) {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
status := &Status{
|
||||
Plan: PlanFree,
|
||||
PlanStatus: "active",
|
||||
PriceMonthlyCents: s.priceMonthlyCents,
|
||||
PriceAnnualCents: s.priceAnnualCents,
|
||||
}
|
||||
|
||||
if sub != nil {
|
||||
cadence := sub.Cadence
|
||||
periodEnd := sub.CurrentPeriodEnd
|
||||
// past_due keeps access: Stripe still considers the subscription live
|
||||
// during the dunning window.
|
||||
switch stripe.SubscriptionStatus(sub.Status) {
|
||||
case stripe.SubscriptionStatusActive,
|
||||
stripe.SubscriptionStatusTrialing,
|
||||
stripe.SubscriptionStatusPastDue:
|
||||
status.Plan = PlanPro
|
||||
}
|
||||
status.PlanStatus = sub.Status
|
||||
status.Cadence = &cadence
|
||||
status.Seats = sub.Quantity
|
||||
status.CurrentPeriodEnd = &periodEnd
|
||||
status.CancelAtPeriodEnd = sub.CancelAtPeriodEnd
|
||||
}
|
||||
|
||||
return status, nil
|
||||
}
|
||||
|
||||
func (s *serviceImpl) CreateCheckoutSession(ctx context.Context, p CheckoutParams) (string, error) {
|
||||
if !p.Cadence.IsValid() {
|
||||
return "", ErrInvalidCadence
|
||||
}
|
||||
if p.Seats < 1 {
|
||||
return "", fmt.Errorf("seats must be >= 1")
|
||||
}
|
||||
|
||||
customerID, err := s.ensureStripeCustomer(ctx, p)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("ensure stripe customer: %w", err)
|
||||
}
|
||||
|
||||
priceID := s.cfg.PriceAnnualID
|
||||
switch p.Cadence {
|
||||
case CadenceAnnual:
|
||||
priceID = s.cfg.PriceAnnualID
|
||||
break
|
||||
case CadenceMonthly:
|
||||
priceID = s.cfg.PriceMonthlyID
|
||||
break
|
||||
}
|
||||
|
||||
params := &stripe.CheckoutSessionParams{
|
||||
Mode: stripe.String(string(stripe.CheckoutSessionModeSubscription)),
|
||||
Customer: stripe.String(customerID),
|
||||
ClientReferenceID: stripe.String(p.NetworkID),
|
||||
SuccessURL: stripe.String(s.cfg.SuccessURL),
|
||||
CancelURL: stripe.String(s.cfg.CancelURL),
|
||||
LineItems: []*stripe.CheckoutSessionLineItemParams{{
|
||||
Price: stripe.String(priceID),
|
||||
Quantity: stripe.Int64(int64(p.Seats)),
|
||||
}},
|
||||
SubscriptionData: &stripe.CheckoutSessionSubscriptionDataParams{
|
||||
Metadata: map[string]string{
|
||||
"network_id": p.NetworkID,
|
||||
"admin_human_id": p.AdminHumanID,
|
||||
"cadence": string(p.Cadence),
|
||||
},
|
||||
},
|
||||
}
|
||||
params.Context = ctx
|
||||
|
||||
sess, err := checkoutsession.New(params)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("stripe checkout: %w", err)
|
||||
}
|
||||
return sess.URL, nil
|
||||
}
|
||||
|
||||
func (s *serviceImpl) CreatePortalSession(ctx context.Context, networkID string) (string, error) {
|
||||
sub, err := s.repo.getSubscriptionByNetworkID(ctx, networkID)
|
||||
if errors.Is(err, errNotFound) {
|
||||
return "", ErrNoActiveSubscription
|
||||
}
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
params := &stripe.BillingPortalSessionParams{
|
||||
Customer: stripe.String(sub.StripeCustomerID),
|
||||
ReturnURL: stripe.String(s.cfg.SuccessURL),
|
||||
}
|
||||
params.Context = ctx
|
||||
|
||||
sess, err := billingportalsession.New(params)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("stripe portal: %w", err)
|
||||
}
|
||||
return sess.URL, nil
|
||||
}
|
||||
|
||||
func (s *serviceImpl) SyncSeats(ctx context.Context, networkID string, seats int) error {
|
||||
sub, err := s.repo.getSubscriptionByNetworkID(ctx, networkID)
|
||||
if errors.Is(err, errNotFound) {
|
||||
return nil
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if sub.Quantity == seats {
|
||||
return nil
|
||||
}
|
||||
|
||||
liveSub, err := stripesub.Get(sub.ID, &stripe.SubscriptionParams{
|
||||
Params: stripe.Params{Context: ctx},
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("fetch stripe subscription: %w", err)
|
||||
}
|
||||
if len(liveSub.Items.Data) == 0 {
|
||||
return fmt.Errorf("stripe subscription %s has no items", sub.ID)
|
||||
}
|
||||
|
||||
params := &stripe.SubscriptionParams{
|
||||
ProrationBehavior: stripe.String(string(stripe.SubscriptionSchedulePhaseProrationBehaviorCreateProrations)),
|
||||
Items: []*stripe.SubscriptionItemsParams{{
|
||||
ID: stripe.String(liveSub.Items.Data[0].ID),
|
||||
Quantity: stripe.Int64(int64(seats)),
|
||||
}},
|
||||
}
|
||||
params.Context = ctx
|
||||
|
||||
if _, err := stripesub.Update(sub.ID, params); err != nil {
|
||||
return fmt.Errorf("update stripe subscription quantity: %w", err)
|
||||
}
|
||||
// customer.subscription.updated webhook arrives within seconds and
|
||||
// reconciles quantity in our DB.
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *serviceImpl) ensureStripeCustomer(ctx context.Context, p CheckoutParams) (string, error) {
|
||||
existing, err := s.repo.getStripeCustomerID(ctx, p.NetworkID)
|
||||
if err == nil {
|
||||
return existing, nil
|
||||
}
|
||||
if !errors.Is(err, errNotFound) {
|
||||
return "", err
|
||||
}
|
||||
|
||||
params := &stripe.CustomerParams{
|
||||
Email: stripe.String(p.AdminEmail),
|
||||
Metadata: map[string]string{
|
||||
"network_id": p.NetworkID,
|
||||
"admin_human_id": p.AdminHumanID,
|
||||
},
|
||||
}
|
||||
params.Context = ctx
|
||||
|
||||
cust, err := stripecustomer.New(params)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("create stripe customer: %w", err)
|
||||
}
|
||||
if err := s.repo.setStripeCustomerID(ctx, p.NetworkID, cust.ID); err != nil {
|
||||
return "", fmt.Errorf("persist stripe customer id: %w", err)
|
||||
}
|
||||
return cust.ID, nil
|
||||
}
|
||||
Reference in New Issue
Block a user