From 279a3e52c2c7f6c5d8be34dd59f1fdee1d9a7a06 Mon Sep 17 00:00:00 2001 From: talksik Date: Tue, 14 Apr 2026 11:00:46 -0700 Subject: [PATCH] implement core foundation --- go/cmd/emailnotifierjob/main.go | 3 +- go/cmd/orion/main.go | 26 +- go/cmd/pusherservice/main.go | 5 +- go/go.mod | 1 + go/go.sum | 5 + go/internal/billing/config.go | 37 ++ go/internal/billing/models.go | 57 +++ go/internal/billing/noop.go | 26 ++ go/internal/billing/repository.go | 112 ++++++ go/internal/billing/service.go | 237 ++++++++++++ go/internal/billing/webhooks.go | 115 ++++++ go/internal/handler/billing.go | 172 +++++++++ go/internal/handler/handler.go | 15 +- go/internal/network/models.go | 10 +- go/internal/network/repository.go | 79 ++-- go/internal/network/service.go | 97 +++-- go/k8s/dev/orion.yaml | 18 + go/k8s/prod/orion.yaml | 24 ++ go/migrations/000013_network_billing.down.sql | 2 + go/migrations/000013_network_billing.up.sql | 24 ++ js/src/App.tsx | 2 + js/src/api/client.ts | 31 ++ js/src/api/types.ts | 42 +++ js/src/features/network-billing.tsx | 350 ++++++++++++++++++ js/src/hooks/use-billing.ts | 27 ++ 25 files changed, 1455 insertions(+), 62 deletions(-) create mode 100644 go/internal/billing/config.go create mode 100644 go/internal/billing/models.go create mode 100644 go/internal/billing/noop.go create mode 100644 go/internal/billing/repository.go create mode 100644 go/internal/billing/service.go create mode 100644 go/internal/billing/webhooks.go create mode 100644 go/internal/handler/billing.go create mode 100644 go/migrations/000013_network_billing.down.sql create mode 100644 go/migrations/000013_network_billing.up.sql create mode 100644 js/src/features/network-billing.tsx create mode 100644 js/src/hooks/use-billing.ts diff --git a/go/cmd/emailnotifierjob/main.go b/go/cmd/emailnotifierjob/main.go index e536279..84871d9 100644 --- a/go/cmd/emailnotifierjob/main.go +++ b/go/cmd/emailnotifierjob/main.go @@ -11,6 +11,7 @@ import ( "cloud.google.com/go/firestore" pbaero "github.com/flowy-live/llink/genproto/aero" pbpusher "github.com/flowy-live/llink/genproto/llink/pusher" + "github.com/flowy-live/llink/internal/billing" "github.com/flowy-live/llink/internal/db" "github.com/flowy-live/llink/internal/human" "github.com/flowy-live/llink/internal/network" @@ -67,7 +68,7 @@ func main() { // Initialize services humanSvc := human.NewService(db.Pool()) - networkSvc := network.NewService(db.Pool(), aeroSvc) + networkSvc := network.NewService(db.Pool(), aeroSvc, billing.Noop()) slog.Info("starting email notification cycle") if err := runNotificationCycle(ctx, firestoreClient, aeroSvc, pusherSvc, humanSvc, networkSvc); err != nil { diff --git a/go/cmd/orion/main.go b/go/cmd/orion/main.go index 9880ad8..e27f2da 100644 --- a/go/cmd/orion/main.go +++ b/go/cmd/orion/main.go @@ -12,6 +12,7 @@ import ( pbaero "github.com/flowy-live/llink/genproto/aero" "github.com/flowy-live/llink/internal" "github.com/flowy-live/llink/internal/auth" + "github.com/flowy-live/llink/internal/billing" "github.com/flowy-live/llink/internal/db" "github.com/flowy-live/llink/internal/depot" "github.com/flowy-live/llink/internal/handler" @@ -64,10 +65,25 @@ func main() { defer aeroServer.Close() aeroSvc := pbaero.NewPrimaryClient(aeroServer) - // Initialize services + // Initialize services. Billing is built before network because network + // reports seat-count changes to billing on member writes. authSvc := auth.NewAuthService(redisClient, aeroSvc) humanSvc := human.NewService(db.Pool()) - networkSvc := network.NewService(db.Pool(), aeroSvc) + + billingSvc, err := billing.NewService(ctx, db.Pool(), billing.Config{ + SecretKey: utils.MustGetEnv("STRIPE_SECRET_KEY"), + WebhookSecret: utils.MustGetEnv("STRIPE_WEBHOOK_SECRET"), + PriceMonthlyID: utils.MustGetEnv("STRIPE_PRICE_PRO_MONTHLY"), + PriceAnnualID: utils.MustGetEnv("STRIPE_PRICE_PRO_ANNUAL"), + SuccessURL: utils.MustGetEnv("BILLING_SUCCESS_URL"), + CancelURL: utils.MustGetEnv("BILLING_CANCEL_URL"), + }) + if err != nil { + slog.Error("failed to initialize billing service", "error", err) + os.Exit(1) + } + + networkSvc := network.NewService(db.Pool(), aeroSvc, billingSvc) particleSvc := particle.NewService(db.Pool(), networkSvc) depotSvc := depot.NewService(db.Pool(), storageClient, depot.Config{ GoogleServiceAccountEmail: utils.MustGetEnv("GOOGLE_SERVICE_ACCOUNT_EMAIL"), @@ -108,6 +124,7 @@ func main() { mux.HandleFunc("POST /auth/sign-in", h.SignIn) mux.HandleFunc("POST /waitlist", h.AddToWaitlist) mux.HandleFunc("POST /livekit/webhook", h.HandleLivekitWebhook) + mux.HandleFunc("POST /webhooks/stripe", h.HandleStripeWebhook) // ========================================================================== // Protected routes (auth required) @@ -127,6 +144,11 @@ func main() { mux.Handle("POST /networks/{id}/members", withAuth(h.AddMembersToNetwork)) // mux.Handle("DELETE /networks/{id}/members/{humanId}", withAuth(h.RemoveMemberFromNetwork)) + // Billing (network admin only; admin check happens inside each handler) + mux.Handle("GET /networks/{id}/billing", withAuth(h.GetNetworkBilling)) + mux.Handle("POST /networks/{id}/billing/checkout-session", withAuth(h.CreateCheckoutSession)) + mux.Handle("POST /networks/{id}/billing/portal-session", withAuth(h.CreatePortalSession)) + // Network Invitations mux.Handle("GET /networks/{id}/invitations", withAuth(h.ListInvitationsForNetwork)) mux.Handle("DELETE /networks/{id}/invitations", withAuth(h.RevokeInvitation)) diff --git a/go/cmd/pusherservice/main.go b/go/cmd/pusherservice/main.go index 2e36865..0dab5cb 100644 --- a/go/cmd/pusherservice/main.go +++ b/go/cmd/pusherservice/main.go @@ -12,6 +12,7 @@ import ( "github.com/flowy-live/llink/internal" "github.com/flowy-live/llink/internal/auth" + "github.com/flowy-live/llink/internal/billing" "github.com/flowy-live/llink/internal/db" "github.com/flowy-live/llink/internal/network" "github.com/flowy-live/llink/internal/pusher" @@ -36,8 +37,8 @@ func main() { pusherRedis := internal.ConnectAndTestRedis(db.RedisDBPusher) // Services - authSvc := auth.NewAuthService(authRedis, nil) // nil aeroSvc — pusher only calls GetSession - networkSvc := network.NewService(db.Pool(), nil) // nil aeroSvc — pusher never calls InviteByEmail + authSvc := auth.NewAuthService(authRedis, nil) // nil aeroSvc — pusher only calls GetSession + networkSvc := network.NewService(db.Pool(), nil, billing.Noop()) // nil aeroSvc / noop billing — pusher never mutates membership // Pod identity (use hostname in k8s, which is the pod name) podID, err := os.Hostname() diff --git a/go/go.mod b/go/go.mod index 2ff2e45..ef5952a 100644 --- a/go/go.mod +++ b/go/go.mod @@ -142,6 +142,7 @@ require ( github.com/puzpuzpuz/xsync/v3 v3.5.1 // indirect github.com/shirou/gopsutil/v4 v4.25.6 // indirect github.com/spiffe/go-spiffe/v2 v2.6.0 // indirect + github.com/stripe/stripe-go/v81 v81.4.0 // indirect github.com/tklauser/go-sysconf v0.3.12 // indirect github.com/tklauser/numcpus v0.6.1 // indirect github.com/twitchtv/twirp v8.1.3+incompatible // indirect diff --git a/go/go.sum b/go/go.sum index c849499..8fedaf0 100644 --- a/go/go.sum +++ b/go/go.sum @@ -335,6 +335,8 @@ github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UV github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U= github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U= +github.com/stripe/stripe-go/v81 v81.4.0 h1:AuD9XzdAvl193qUCSaLocf8H+nRopOouXhxqJUzCLbw= +github.com/stripe/stripe-go/v81 v81.4.0/go.mod h1:C/F4jlmnGNacvYtBp/LUHCvVUJEZffFQCobkzwY1WOo= github.com/testcontainers/testcontainers-go v0.40.0 h1:pSdJYLOVgLE8YdUY2FHQ1Fxu+aMnb6JfVz1mxk7OeMU= github.com/testcontainers/testcontainers-go v0.40.0/go.mod h1:FSXV5KQtX2HAMlm7U3APNyLkkap35zNLxukw9oBi/MY= github.com/testcontainers/testcontainers-go/modules/postgres v0.40.0 h1:s2bIayFXlbDFexo96y+htn7FzuhpXLYJNnIuglNKqOk= @@ -415,6 +417,7 @@ golang.org/x/mod v0.33.0 h1:tHFzIWbBifEmbwtGz65eaWyGiGZatSrT9prnU8DbVL8= golang.org/x/mod v0.33.0/go.mod h1:swjeQEj+6r7fODbD2cqrnje9PnziFuw4bmLbBZFrQ5w= golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s= golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg= +golang.org/x/net v0.0.0-20210520170846-37e1c6afe023/go.mod h1:9nx3DQGgdP8bBQD5qxJ1jj9UTztislL4KSBs9R2vV5Y= golang.org/x/net v0.0.0-20220722155237-a158d28d115b/go.mod h1:XRhObCWvk6IyKnWLug+ECip1KBveYUHfp+8e9klMJ9c= golang.org/x/net v0.6.0/go.mod h1:2Tu9+aMcznHK/AK1HMvgo6xiTLG5rD5rZLDS+rp2Bjs= golang.org/x/net v0.10.0/go.mod h1:0qNGK6F8kojg2nk9dLZ2mShWaEBan6FAoqfSigmmuDg= @@ -431,6 +434,7 @@ golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5h golang.org/x/sys v0.0.0-20190916202348-b4ddaad3f8a3/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/sys v0.0.0-20201204225414-ed752295db88/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= +golang.org/x/sys v0.0.0-20210423082822-04245dca01da/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/sys v0.0.0-20210615035016-665e8c7367d1/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.0.0-20210616094352-59db8d763f22/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.0.0-20220520151302-bc2c85ada10a/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= @@ -453,6 +457,7 @@ golang.org/x/term v0.40.0 h1:36e4zGLqU4yhjlmxEaagx2KuYbJq3EwY8K943ZsHcvg= golang.org/x/term v0.40.0/go.mod h1:w2P8uVp06p2iyKKuvXIm7N/y0UCRt3UfJTfZ7oOpglM= golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ= golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= +golang.org/x/text v0.3.6/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= golang.org/x/text v0.3.7/go.mod h1:u+2+/6zg+i71rQMx5EYifcz6MCKuco9NR6JIITiCfzQ= golang.org/x/text v0.7.0/go.mod h1:mrYo+phRRbMaCq/xk9113O4dZlRixOauAjOtrjsXDZ8= golang.org/x/text v0.9.0/go.mod h1:e1OnstbJyHTd6l/uOt8jFFHp6TRDWZR/bV3emEE/zU8= diff --git a/go/internal/billing/config.go b/go/internal/billing/config.go new file mode 100644 index 0000000..c64c321 --- /dev/null +++ b/go/internal/billing/config.go @@ -0,0 +1,37 @@ +package billing + +import ( + "fmt" + "strings" +) + +type Config struct { + SecretKey string + WebhookSecret string + PriceMonthlyID string + PriceAnnualID string + + SuccessURL string + CancelURL string +} + +func (c Config) Validate() error { + required := map[string]string{ + "SecretKey": c.SecretKey, + "WebhookSecret": c.WebhookSecret, + "PriceMonthlyID": c.PriceMonthlyID, + "PriceAnnualID": c.PriceAnnualID, + "SuccessURL": c.SuccessURL, + "CancelURL": c.CancelURL, + } + var missing []string + for name, value := range required { + if value == "" { + missing = append(missing, name) + } + } + if len(missing) > 0 { + return fmt.Errorf("billing config missing: %s", strings.Join(missing, ", ")) + } + return nil +} diff --git a/go/internal/billing/models.go b/go/internal/billing/models.go new file mode 100644 index 0000000..9a0a4b6 --- /dev/null +++ b/go/internal/billing/models.go @@ -0,0 +1,57 @@ +package billing + +import "time" + +type Cadence string + +const ( + CadenceMonthly Cadence = "monthly" + CadenceAnnual Cadence = "annual" +) + +func (c Cadence) IsValid() bool { + return c == CadenceMonthly || c == CadenceAnnual +} + +type Plan string + +const ( + PlanFree Plan = "free" + PlanPro Plan = "pro" +) + +// Subscription is the persisted projection of a Stripe subscription, +// reconciled on every webhook event. +type Subscription struct { + ID string + NetworkID string + StripeCustomerID string + Status string + PriceID string + Cadence Cadence + Quantity int + CancelAtPeriodEnd bool + CurrentPeriodStart time.Time + CurrentPeriodEnd time.Time + CreatedAt time.Time + UpdatedAt time.Time +} + +type Status struct { + Plan Plan `json:"plan"` + PlanStatus string `json:"plan_status"` + Cadence *Cadence `json:"cadence"` + Seats int `json:"seats"` // 0 on free plan; sub.Quantity on pro + CurrentPeriodEnd *time.Time `json:"current_period_end"` + CancelAtPeriodEnd bool `json:"cancel_at_period_end"` + PriceMonthlyCents int64 `json:"price_monthly_cents"` + PriceAnnualCents int64 `json:"price_annual_cents"` +} + +type CheckoutParams struct { + NetworkID string + AdminHumanID string + AdminEmail string + Cadence Cadence + Seats int +} diff --git a/go/internal/billing/noop.go b/go/internal/billing/noop.go new file mode 100644 index 0000000..764a2bd --- /dev/null +++ b/go/internal/billing/noop.go @@ -0,0 +1,26 @@ +package billing + +import ( + "context" + "errors" +) + +// Noop returns a billing service for binaries that depend on network.Service +// but never mutate membership (jobs, pusher). Stripe isn't configured. +func Noop() Service { return noopService{} } + +type noopService struct{} + +var errNoopBilling = errors.New("billing: not configured in this process") + +func (noopService) GetStatus(context.Context, string) (*Status, error) { + return nil, errNoopBilling +} +func (noopService) CreateCheckoutSession(context.Context, CheckoutParams) (string, error) { + return "", errNoopBilling +} +func (noopService) CreatePortalSession(context.Context, string) (string, error) { + return "", errNoopBilling +} +func (noopService) SyncSeats(context.Context, string, int) error { return nil } +func (noopService) HandleWebhook(context.Context, []byte, string) error { return errNoopBilling } diff --git a/go/internal/billing/repository.go b/go/internal/billing/repository.go new file mode 100644 index 0000000..06ccf76 --- /dev/null +++ b/go/internal/billing/repository.go @@ -0,0 +1,112 @@ +package billing + +import ( + "context" + "errors" + + "github.com/jackc/pgx/v5" + "github.com/jackc/pgx/v5/pgxpool" +) + +var errNotFound = errors.New("not found") + +type repository interface { + getSubscriptionByNetworkID(ctx context.Context, networkID string) (*Subscription, error) + upsertSubscription(ctx context.Context, sub *Subscription) error + deleteSubscriptionByID(ctx context.Context, subscriptionID string) error + + getStripeCustomerID(ctx context.Context, networkID string) (string, error) + setStripeCustomerID(ctx context.Context, networkID, customerID string) error +} + +type repositoryImpl struct { + pool *pgxpool.Pool +} + +func newRepository(pool *pgxpool.Pool) repository { + return &repositoryImpl{pool: pool} +} + +const subscriptionColumns = `id, network_id, stripe_customer_id, status, price_id, cadence, quantity, cancel_at_period_end, current_period_start, current_period_end, created_at, updated_at` + +func scanSubscription(row pgx.Row, s *Subscription) error { + return row.Scan( + &s.ID, &s.NetworkID, &s.StripeCustomerID, &s.Status, &s.PriceID, + &s.Cadence, &s.Quantity, &s.CancelAtPeriodEnd, + &s.CurrentPeriodStart, &s.CurrentPeriodEnd, + &s.CreatedAt, &s.UpdatedAt, + ) +} + +func (r *repositoryImpl) getSubscriptionByNetworkID(ctx context.Context, networkID string) (*Subscription, error) { + var s Subscription + err := scanSubscription( + r.pool.QueryRow(ctx, + `SELECT `+subscriptionColumns+` FROM network_subscriptions WHERE network_id = $1`, + networkID, + ), + &s, + ) + if errors.Is(err, pgx.ErrNoRows) { + return nil, errNotFound + } + if err != nil { + return nil, err + } + return &s, nil +} + +func (r *repositoryImpl) upsertSubscription(ctx context.Context, sub *Subscription) error { + _, err := r.pool.Exec(ctx, ` + INSERT INTO network_subscriptions ( + id, network_id, stripe_customer_id, status, price_id, cadence, + quantity, cancel_at_period_end, current_period_start, current_period_end + ) VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10) + ON CONFLICT (id) DO UPDATE SET + stripe_customer_id = EXCLUDED.stripe_customer_id, + status = EXCLUDED.status, + price_id = EXCLUDED.price_id, + cadence = EXCLUDED.cadence, + quantity = EXCLUDED.quantity, + cancel_at_period_end = EXCLUDED.cancel_at_period_end, + current_period_start = EXCLUDED.current_period_start, + current_period_end = EXCLUDED.current_period_end, + updated_at = NOW() + `, + sub.ID, sub.NetworkID, sub.StripeCustomerID, sub.Status, sub.PriceID, sub.Cadence, + sub.Quantity, sub.CancelAtPeriodEnd, sub.CurrentPeriodStart, sub.CurrentPeriodEnd, + ) + return err +} + +func (r *repositoryImpl) deleteSubscriptionByID(ctx context.Context, subscriptionID string) error { + _, err := r.pool.Exec(ctx, + `DELETE FROM network_subscriptions WHERE id = $1`, + subscriptionID, + ) + return err +} + +func (r *repositoryImpl) getStripeCustomerID(ctx context.Context, networkID string) (string, error) { + var customerID string + err := r.pool.QueryRow(ctx, + `SELECT stripe_customer_id FROM network_stripe_customers WHERE network_id = $1`, + networkID, + ).Scan(&customerID) + if errors.Is(err, pgx.ErrNoRows) { + return "", errNotFound + } + if err != nil { + return "", err + } + return customerID, nil +} + +func (r *repositoryImpl) setStripeCustomerID(ctx context.Context, networkID, customerID string) error { + _, err := r.pool.Exec(ctx, ` + INSERT INTO network_stripe_customers (network_id, stripe_customer_id) + VALUES ($1, $2) + ON CONFLICT (network_id) DO NOTHING + `, networkID, customerID) + return err +} diff --git a/go/internal/billing/service.go b/go/internal/billing/service.go new file mode 100644 index 0000000..a97217b --- /dev/null +++ b/go/internal/billing/service.go @@ -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 +} diff --git a/go/internal/billing/webhooks.go b/go/internal/billing/webhooks.go new file mode 100644 index 0000000..a0e50ab --- /dev/null +++ b/go/internal/billing/webhooks.go @@ -0,0 +1,115 @@ +package billing + +import ( + "context" + "encoding/json" + "fmt" + "log/slog" + "time" + + "github.com/stripe/stripe-go/v81" + "github.com/stripe/stripe-go/v81/webhook" +) + +func (s *serviceImpl) HandleWebhook(ctx context.Context, payload []byte, signature string) error { + event, err := webhook.ConstructEvent(payload, signature, s.cfg.WebhookSecret) + if err != nil { + return fmt.Errorf("verify stripe signature: %w", err) + } + + slog.Info("stripe webhook", "type", event.Type, "id", event.ID) + + switch event.Type { + case "checkout.session.completed": + // subscription.created fires right after with full detail; we handle + // the subscription there. + return nil + case "customer.subscription.created", "customer.subscription.updated": + return s.handleSubscriptionUpsert(ctx, event) + case "customer.subscription.deleted": + return s.handleSubscriptionDeleted(ctx, event) + default: + return nil + } +} + +func (s *serviceImpl) handleSubscriptionUpsert(ctx context.Context, event stripe.Event) error { + var sub stripe.Subscription + if err := json.Unmarshal(event.Data.Raw, &sub); err != nil { + return fmt.Errorf("decode subscription: %w", err) + } + + networkID := sub.Metadata["network_id"] + if networkID == "" { + return fmt.Errorf("subscription %s missing network_id metadata", sub.ID) + } + + local, err := subscriptionFromStripe(&sub, networkID) + if err != nil { + return err + } + if err := s.repo.upsertSubscription(ctx, local); err != nil { + return fmt.Errorf("upsert subscription: %w", err) + } + return nil +} + +func (s *serviceImpl) handleSubscriptionDeleted(ctx context.Context, event stripe.Event) error { + var sub stripe.Subscription + if err := json.Unmarshal(event.Data.Raw, &sub); err != nil { + return fmt.Errorf("decode subscription: %w", err) + } + + networkID := sub.Metadata["network_id"] + if networkID == "" { + return fmt.Errorf("subscription %s missing network_id metadata", sub.ID) + } + + if err := s.repo.deleteSubscriptionByID(ctx, sub.ID); err != nil { + return fmt.Errorf("delete subscription: %w", err) + } + return nil +} + +func subscriptionFromStripe(sub *stripe.Subscription, networkID string) (*Subscription, error) { + if len(sub.Items.Data) == 0 { + return nil, fmt.Errorf("subscription %s has no items", sub.ID) + } + item := sub.Items.Data[0] + cadence, err := cadenceFromInterval(item.Price) + if err != nil { + return nil, fmt.Errorf("subscription %s: %w", sub.ID, err) + } + + var customerID string + if sub.Customer != nil { + customerID = sub.Customer.ID + } + + return &Subscription{ + ID: sub.ID, + NetworkID: networkID, + StripeCustomerID: customerID, + Status: string(sub.Status), + PriceID: item.Price.ID, + Cadence: cadence, + Quantity: int(item.Quantity), + CancelAtPeriodEnd: sub.CancelAtPeriodEnd, + CurrentPeriodStart: time.Unix(sub.CurrentPeriodStart, 0).UTC(), + CurrentPeriodEnd: time.Unix(sub.CurrentPeriodEnd, 0).UTC(), + }, nil +} + +func cadenceFromInterval(price *stripe.Price) (Cadence, error) { + if price == nil || price.Recurring == nil { + return "", fmt.Errorf("price is not recurring") + } + switch price.Recurring.Interval { + case stripe.PriceRecurringIntervalMonth: + return CadenceMonthly, nil + case stripe.PriceRecurringIntervalYear: + return CadenceAnnual, nil + default: + return "", fmt.Errorf("unsupported price interval %q", price.Recurring.Interval) + } +} diff --git a/go/internal/handler/billing.go b/go/internal/handler/billing.go new file mode 100644 index 0000000..6ef2890 --- /dev/null +++ b/go/internal/handler/billing.go @@ -0,0 +1,172 @@ +package handler + +import ( + "encoding/json" + "errors" + "fmt" + "io" + "log/slog" + "net/http" + + "github.com/flowy-live/llink/internal/billing" + "github.com/flowy-live/llink/internal/middleware" + "github.com/flowy-live/llink/internal/network" +) + +type CreateCheckoutSessionRequest struct { + Cadence string `json:"cadence"` +} + +type CheckoutSessionResponse struct { + URL string `json:"url"` +} + +type PortalSessionResponse struct { + URL string `json:"url"` +} + +func (h *Handler) GetNetworkBilling(w http.ResponseWriter, r *http.Request) { + net, _, ok := h.loadNetworkForAdmin(w, r) + if !ok { + return + } + + status, err := h.billingSvc.GetStatus(r.Context(), net.ID) + if err != nil { + slog.Error("failed to get billing status", "error", err, "network_id", net.ID) + http.Error(w, "internal server error", http.StatusInternalServerError) + return + } + + writeJSON(w, status) +} + +func (h *Handler) CreateCheckoutSession(w http.ResponseWriter, r *http.Request) { + net, adminHumanId, ok := h.loadNetworkForAdmin(w, r) + if !ok { + return + } + + var req CreateCheckoutSessionRequest + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + http.Error(w, "invalid request body", http.StatusBadRequest) + return + } + + cadence := billing.Cadence(req.Cadence) + if !cadence.IsValid() { + http.Error(w, "invalid cadence", http.StatusBadRequest) + return + } + + adminHuman, err := h.humanSvc.GetByID(r.Context(), adminHumanId) + if err != nil { + slog.Error("failed to load admin human", "error", err, "human_id", adminHumanId) + http.Error(w, "internal server error", http.StatusInternalServerError) + return + } + + seats, err := h.networkSvc.CountSeats(r.Context(), net.ID) + if err != nil { + slog.Error("failed to count seats", "error", err, "network_id", net.ID) + http.Error(w, "internal server error", http.StatusInternalServerError) + return + } + + url, err := h.billingSvc.CreateCheckoutSession(r.Context(), billing.CheckoutParams{ + NetworkID: net.ID, + AdminHumanID: adminHumanId, + AdminEmail: adminHuman.Email, + Cadence: cadence, + Seats: seats, + }) + if err != nil { + slog.Error("failed to create checkout session", "error", err, "network_id", net.ID) + http.Error(w, "internal server error", http.StatusInternalServerError) + return + } + + writeJSON(w, CheckoutSessionResponse{URL: url}) +} + +func (h *Handler) CreatePortalSession(w http.ResponseWriter, r *http.Request) { + net, _, ok := h.loadNetworkForAdmin(w, r) + if !ok { + return + } + + url, err := h.billingSvc.CreatePortalSession(r.Context(), net.ID) + if errors.Is(err, billing.ErrNoActiveSubscription) { + http.Error(w, "no active subscription", http.StatusConflict) + return + } + if err != nil { + slog.Error("failed to create portal session", "error", err, "network_id", net.ID) + http.Error(w, "internal server error", http.StatusInternalServerError) + return + } + + writeJSON(w, PortalSessionResponse{URL: url}) +} + +const maxStripeWebhookBytes = 1 << 20 // 1 MiB + +func (h *Handler) HandleStripeWebhook(w http.ResponseWriter, r *http.Request) { + payload, err := io.ReadAll(io.LimitReader(r.Body, maxStripeWebhookBytes)) + if err != nil { + slog.Warn("stripe webhook: failed to read body", "error", err) + http.Error(w, "bad request", http.StatusBadRequest) + return + } + signature := r.Header.Get("Stripe-Signature") + + if err := h.billingSvc.HandleWebhook(r.Context(), payload, signature); err != nil { + slog.Error("stripe webhook failed", "error", err) + formattedErr := fmt.Errorf("webhook processing failed: %w", err) + http.Error(w, formattedErr.Error(), http.StatusBadRequest) + return + } + + w.WriteHeader(http.StatusOK) +} + +// loadNetworkForAdmin resolves the {id} path param and verifies the caller +// is the network's admin. On failure it writes the HTTP error and returns ok=false. +func (h *Handler) loadNetworkForAdmin(w http.ResponseWriter, r *http.Request) (*network.Network, string, bool) { + humanId, ok := middleware.HumanIdFromContext(r.Context()) + if !ok { + http.Error(w, "unauthorized", http.StatusUnauthorized) + return nil, "", false + } + + networkID := r.PathValue("id") + if networkID == "" { + http.Error(w, "network id is required", http.StatusBadRequest) + return nil, "", false + } + + net, err := h.networkSvc.GetByID(r.Context(), networkID) + if errors.Is(err, network.ErrNotFound) { + http.Error(w, "network not found", http.StatusNotFound) + return nil, "", false + } + if err != nil { + slog.Error("failed to load network", "error", err, "network_id", networkID) + http.Error(w, "internal server error", http.StatusInternalServerError) + return nil, "", false + } + + if net.AdminHumanId != humanId { + http.Error(w, "forbidden", http.StatusForbidden) + return nil, "", false + } + + return net, humanId, true +} + +func writeJSON(w http.ResponseWriter, body any) { + w.Header().Set("Content-Type", "application/json") + if err := json.NewEncoder(w).Encode(body); err != nil { + slog.Error("failed to write json", "error", err) + } +} diff --git a/go/internal/handler/handler.go b/go/internal/handler/handler.go index 3db4cee..d6d90e8 100644 --- a/go/internal/handler/handler.go +++ b/go/internal/handler/handler.go @@ -12,6 +12,7 @@ import ( "cloud.google.com/go/firestore" "github.com/flowy-live/llink/internal/auth" + "github.com/flowy-live/llink/internal/billing" "github.com/flowy-live/llink/internal/depot" "github.com/flowy-live/llink/internal/human" "github.com/flowy-live/llink/internal/livekit" @@ -30,11 +31,22 @@ type Handler struct { particleSvc particle.Service depotSvc depot.Service waitlistSvc waitlist.Service + billingSvc billing.Service livekitClient livekit.Client firestoreClient *firestore.Client } -func NewHandler(authSvc auth.AuthService, humanSvc human.Service, networkSvc network.Service, particleSvc particle.Service, depotSvc depot.Service, waitlistSvc waitlist.Service, livekitClient livekit.Client, firestoreClient *firestore.Client) *Handler { +func NewHandler( + authSvc auth.AuthService, + humanSvc human.Service, + networkSvc network.Service, + particleSvc particle.Service, + depotSvc depot.Service, + waitlistSvc waitlist.Service, + billingSvc billing.Service, + livekitClient livekit.Client, + firestoreClient *firestore.Client, +) *Handler { return &Handler{ authSvc: authSvc, humanSvc: humanSvc, @@ -42,6 +54,7 @@ func NewHandler(authSvc auth.AuthService, humanSvc human.Service, networkSvc net particleSvc: particleSvc, depotSvc: depotSvc, waitlistSvc: waitlistSvc, + billingSvc: billingSvc, livekitClient: livekitClient, firestoreClient: firestoreClient, } diff --git a/go/internal/network/models.go b/go/internal/network/models.go index a0906b1..9a3cc76 100644 --- a/go/internal/network/models.go +++ b/go/internal/network/models.go @@ -3,11 +3,11 @@ package network import "time" type Network struct { - ID string - Name string - AdminHumanId string - MemberHumanIds []string - CreatedAt time.Time + ID string + Name string + AdminHumanId string + MemberHumanIds []string + CreatedAt time.Time } type Invitation struct { diff --git a/go/internal/network/repository.go b/go/internal/network/repository.go index b7db2f7..523a89b 100644 --- a/go/internal/network/repository.go +++ b/go/internal/network/repository.go @@ -5,10 +5,20 @@ import ( "errors" "github.com/jackc/pgx/v5" + "github.com/jackc/pgx/v5/pgconn" "github.com/jackc/pgx/v5/pgxpool" "go.jetify.com/typeid" ) +// dbtx is the subset of pgx's query API shared by *pgxpool.Pool and pgx.Tx. +// Used by repository helpers that the service layer may run either standalone +// (against the pool) or inside a transaction. +type dbtx interface { + Exec(ctx context.Context, sql string, args ...any) (pgconn.CommandTag, error) + Query(ctx context.Context, sql string, args ...any) (pgx.Rows, error) + QueryRow(ctx context.Context, sql string, args ...any) pgx.Row +} + var errNotFound = errors.New("not found") type networkIDPrefix struct{} @@ -30,9 +40,10 @@ type repository interface { getByID(ctx context.Context, id string) (*Network, error) updateName(ctx context.Context, id, name string) error delete(ctx context.Context, id string) error - addMember(ctx context.Context, networkID, humanId string) error - removeMember(ctx context.Context, networkID, humanId string) error + addMember(ctx context.Context, db dbtx, networkID, humanId string) error + removeMember(ctx context.Context, db dbtx, networkID, humanId string) error getMemberHumanIds(ctx context.Context, networkID string) ([]string, error) + countSeats(ctx context.Context, db dbtx, networkID string) (int, error) getNetworksForHuman(ctx context.Context, humanId string) ([]*Network, error) isMember(ctx context.Context, networkID, humanId string) (bool, error) listAll(ctx context.Context) ([]*Network, error) @@ -41,7 +52,15 @@ type repository interface { createInvitation(ctx context.Context, networkID, email string) error getInvitationsByEmail(ctx context.Context, email string) ([]*Invitation, error) getInvitationsByNetwork(ctx context.Context, networkID string) ([]*Invitation, error) - deleteInvitation(ctx context.Context, networkID, email string) error + deleteInvitation(ctx context.Context, db dbtx, networkID, email string) error +} + +// networkColumns lists every column selected when hydrating a Network. +// Centralized to keep SELECTs and Scan() calls in sync. +const networkColumns = `id, name, admin_human_id, created_at` + +func scanNetwork(row pgx.Row, n *Network) error { + return row.Scan(&n.ID, &n.Name, &n.AdminHumanId, &n.CreatedAt) } type repositoryImpl struct { @@ -59,12 +78,12 @@ func (r *repositoryImpl) create(ctx context.Context, name, adminHumanId string) } var n Network - err = r.pool.QueryRow(ctx, + row := r.pool.QueryRow(ctx, `INSERT INTO networks (id, name, admin_human_id) VALUES ($1, $2, $3) - RETURNING id, name, admin_human_id, created_at`, + RETURNING `+networkColumns, id.String(), name, adminHumanId, - ).Scan(&n.ID, &n.Name, &n.AdminHumanId, &n.CreatedAt) - if err != nil { + ) + if err := scanNetwork(row, &n); err != nil { return nil, err } @@ -74,21 +93,22 @@ func (r *repositoryImpl) create(ctx context.Context, name, adminHumanId string) func (r *repositoryImpl) getByID(ctx context.Context, id string) (*Network, error) { var n Network - err := r.pool.QueryRow(ctx, - `SELECT id, name, admin_human_id, created_at FROM networks WHERE id = $1`, + row := r.pool.QueryRow(ctx, + `SELECT `+networkColumns+` FROM networks WHERE id = $1`, id, - ).Scan(&n.ID, &n.Name, &n.AdminHumanId, &n.CreatedAt) - if err != nil { + ) + if err := scanNetwork(row, &n); err != nil { if errors.Is(err, pgx.ErrNoRows) { return nil, errNotFound } return nil, err } - n.MemberHumanIds, err = r.getMemberHumanIds(ctx, id) + memberIds, err := r.getMemberHumanIds(ctx, id) if err != nil { return nil, err } + n.MemberHumanIds = memberIds return &n, nil } @@ -118,8 +138,8 @@ func (r *repositoryImpl) delete(ctx context.Context, id string) error { return nil } -func (r *repositoryImpl) addMember(ctx context.Context, networkID, humanId string) error { - _, err := r.pool.Exec(ctx, +func (r *repositoryImpl) addMember(ctx context.Context, db dbtx, networkID, humanId string) error { + _, err := db.Exec(ctx, `INSERT INTO network_members (network_id, human_id) VALUES ($1, $2) ON CONFLICT (network_id, human_id) DO NOTHING`, networkID, humanId, @@ -127,14 +147,23 @@ func (r *repositoryImpl) addMember(ctx context.Context, networkID, humanId strin return err } -func (r *repositoryImpl) removeMember(ctx context.Context, networkID, humanId string) error { - _, err := r.pool.Exec(ctx, +func (r *repositoryImpl) removeMember(ctx context.Context, db dbtx, networkID, humanId string) error { + _, err := db.Exec(ctx, `DELETE FROM network_members WHERE network_id = $1 AND human_id = $2`, networkID, humanId, ) return err } +func (r *repositoryImpl) countSeats(ctx context.Context, db dbtx, networkID string) (int, error) { + var count int + err := db.QueryRow(ctx, + `SELECT COUNT(*) FROM network_members WHERE network_id = $1`, + networkID, + ).Scan(&count) + return count, err +} + func (r *repositoryImpl) getMemberHumanIds(ctx context.Context, networkID string) ([]string, error) { rows, err := r.pool.Query(ctx, `SELECT human_id FROM network_members WHERE network_id = $1`, @@ -158,10 +187,10 @@ func (r *repositoryImpl) getMemberHumanIds(ctx context.Context, networkID string func (r *repositoryImpl) getNetworksForHuman(ctx context.Context, humanId string) ([]*Network, error) { rows, err := r.pool.Query(ctx, - `SELECT n.id, n.name, n.admin_human_id, n.created_at - FROM networks n - WHERE n.admin_human_id = $1 - OR EXISTS (SELECT 1 FROM network_members nm WHERE nm.network_id = n.id AND nm.human_id = $1)`, + `SELECT `+networkColumns+` + FROM networks + WHERE admin_human_id = $1 + OR EXISTS (SELECT 1 FROM network_members nm WHERE nm.network_id = id AND nm.human_id = $1)`, humanId, ) if err != nil { @@ -172,7 +201,7 @@ func (r *repositoryImpl) getNetworksForHuman(ctx context.Context, humanId string var networks []*Network for rows.Next() { var n Network - if err := rows.Scan(&n.ID, &n.Name, &n.AdminHumanId, &n.CreatedAt); err != nil { + if err := scanNetwork(rows, &n); err != nil { return nil, err } networks = append(networks, &n) @@ -205,7 +234,7 @@ func (r *repositoryImpl) isMember(ctx context.Context, networkID, humanId string func (r *repositoryImpl) listAll(ctx context.Context) ([]*Network, error) { rows, err := r.pool.Query(ctx, - `SELECT id, name, admin_human_id, created_at FROM networks`, + `SELECT `+networkColumns+` FROM networks`, ) if err != nil { return nil, err @@ -215,7 +244,7 @@ func (r *repositoryImpl) listAll(ctx context.Context) ([]*Network, error) { var networks []*Network for rows.Next() { var n Network - if err := rows.Scan(&n.ID, &n.Name, &n.AdminHumanId, &n.CreatedAt); err != nil { + if err := scanNetwork(rows, &n); err != nil { return nil, err } networks = append(networks, &n) @@ -293,8 +322,8 @@ func (r *repositoryImpl) getInvitationsByNetwork(ctx context.Context, networkID return invitations, rows.Err() } -func (r *repositoryImpl) deleteInvitation(ctx context.Context, networkID, email string) error { - _, err := r.pool.Exec(ctx, +func (r *repositoryImpl) deleteInvitation(ctx context.Context, db dbtx, networkID, email string) error { + _, err := db.Exec(ctx, `DELETE FROM network_invitations WHERE network_id = $1 AND email = $2`, networkID, email, ) diff --git a/go/internal/network/service.go b/go/internal/network/service.go index 738ed21..a5b9f01 100644 --- a/go/internal/network/service.go +++ b/go/internal/network/service.go @@ -9,8 +9,11 @@ import ( "strings" pbaero "github.com/flowy-live/llink/genproto/aero" + "github.com/flowy-live/llink/internal/billing" "github.com/flowy-live/llink/internal/utils" + "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/pgxpool" + "slices" ) var ErrNotFound = errors.New("network not found") @@ -19,17 +22,16 @@ var ErrCapacityExceeded = errors.New("active stream capacity exceeded") var ErrInvalidRetentionHours = errors.New("message retention hours must be between 24 and 336") type Service interface { - // Create creates a network and adds adminHumanId as the first member. Returns ErrInvalidName if name is empty. Create(ctx context.Context, name, adminHumanId string) (*Network, error) - // GetByID returns ErrNotFound if network doesn't exist. GetByID(ctx context.Context, id string) (*Network, error) - // SetName returns ErrNotFound or ErrInvalidName. SetName(ctx context.Context, id, name string) error + // AddMembers inserts members and syncs the new seat count to billing + // atomically; a Stripe failure rolls the insert back. AddMembers(ctx context.Context, networkID string, humanIds []string) error RemoveMember(ctx context.Context, networkID, humanId string) error + CountSeats(ctx context.Context, networkID string) (int, error) ListForHuman(ctx context.Context, humanId string) ([]*Network, error) IsMember(ctx context.Context, networkID, humanId string) (bool, error) - // ListAll returns all networks with their members ListAll(ctx context.Context) ([]*Network, error) // Invitations (email-based, for users who haven't registered yet) @@ -41,12 +43,19 @@ type Service interface { } type serviceImpl struct { - repo repository - aeroSvc pbaero.PrimaryClient + pool *pgxpool.Pool + repo repository + aeroSvc pbaero.PrimaryClient + billingSvc billing.Service } -func NewService(pool *pgxpool.Pool, aeroSvc pbaero.PrimaryClient) Service { - return &serviceImpl{repo: newRepository(pool), aeroSvc: aeroSvc} +func NewService(pool *pgxpool.Pool, aeroSvc pbaero.PrimaryClient, billingSvc billing.Service) Service { + return &serviceImpl{ + pool: pool, + repo: newRepository(pool), + aeroSvc: aeroSvc, + billingSvc: billingSvc, + } } func (s *serviceImpl) Create(ctx context.Context, name, adminHumanId string) (*Network, error) { @@ -60,8 +69,7 @@ func (s *serviceImpl) Create(ctx context.Context, name, adminHumanId string) (*N return nil, err } - err = s.AddMembers(ctx, network.ID, []string{adminHumanId}) - if err != nil { + if err := s.AddMembers(ctx, network.ID, []string{adminHumanId}); err != nil { return nil, err } @@ -90,22 +98,58 @@ func (s *serviceImpl) SetName(ctx context.Context, id, name string) error { } func (s *serviceImpl) AddMembers(ctx context.Context, networkID string, humanIds []string) error { - for _, humanId := range humanIds { - if humanId == "" { - return fmt.Errorf("invalid humanId") - } - if err := s.repo.addMember(ctx, networkID, humanId); err != nil { - return err - } + if slices.Contains(humanIds, "") { + return fmt.Errorf("invalid humanId") } - return nil + return s.mutateMembers(ctx, networkID, func(tx pgx.Tx) error { + for _, humanId := range humanIds { + if err := s.repo.addMember(ctx, tx, networkID, humanId); err != nil { + return err + } + } + return nil + }) } func (s *serviceImpl) RemoveMember(ctx context.Context, networkID, humanId string) error { if humanId == "" { return fmt.Errorf("invalid humanId") } - return s.repo.removeMember(ctx, networkID, humanId) + return s.mutateMembers(ctx, networkID, func(tx pgx.Tx) error { + return s.repo.removeMember(ctx, tx, networkID, humanId) + }) +} + +// mutateMembers runs fn in a tx, recounts seats, calls billing.SyncSeats, +// and commits. Any error rolls the membership change back. +func (s *serviceImpl) mutateMembers(ctx context.Context, networkID string, fn func(pgx.Tx) error) error { + tx, err := s.pool.Begin(ctx) + if err != nil { + return fmt.Errorf("begin tx: %w", err) + } + defer tx.Rollback(ctx) + + if err := fn(tx); err != nil { + return err + } + + seats, err := s.repo.countSeats(ctx, tx, networkID) + if err != nil { + return fmt.Errorf("count seats: %w", err) + } + + if err := s.billingSvc.SyncSeats(ctx, networkID, seats); err != nil { + return fmt.Errorf("sync billing seats: %w", err) + } + + if err := tx.Commit(ctx); err != nil { + return fmt.Errorf("commit tx: %w", err) + } + return nil +} + +func (s *serviceImpl) CountSeats(ctx context.Context, networkID string) (int, error) { + return s.repo.countSeats(ctx, s.pool, networkID) } func (s *serviceImpl) ListForHuman(ctx context.Context, humanId string) ([]*Network, error) { @@ -126,8 +170,6 @@ func (s *serviceImpl) ListAll(ctx context.Context) ([]*Network, error) { return s.repo.listAll(ctx) } -// Invitation methods - func (s *serviceImpl) InviteByEmail(ctx context.Context, networkID string, emails []string) error { network, err := s.repo.getByID(ctx, networkID) if err != nil { @@ -180,10 +222,13 @@ func (s *serviceImpl) AcceptInvitation(ctx context.Context, networkID, email, hu return fmt.Errorf("invalid humanId") } - if err := s.repo.deleteInvitation(ctx, networkID, normalized); err != nil { - return err - } - return s.repo.addMember(ctx, networkID, humanId) + return s.mutateMembers(ctx, networkID, func(tx pgx.Tx) error { + err := s.repo.deleteInvitation(ctx, tx, networkID, normalized) + if err != nil { + return err + } + return s.repo.addMember(ctx, tx, networkID, humanId) + }) } func (s *serviceImpl) RevokeInvitation(ctx context.Context, networkID, email string) error { @@ -191,7 +236,7 @@ func (s *serviceImpl) RevokeInvitation(ctx context.Context, networkID, email str if err != nil { return fmt.Errorf("invalid email: %w", err) } - return s.repo.deleteInvitation(ctx, networkID, normalized) + return s.repo.deleteInvitation(ctx, s.pool, networkID, normalized) } func buildInvitationHTML(networkName string) string { diff --git a/go/k8s/dev/orion.yaml b/go/k8s/dev/orion.yaml index fe5a099..0596ed6 100644 --- a/go/k8s/dev/orion.yaml +++ b/go/k8s/dev/orion.yaml @@ -61,6 +61,24 @@ spec: secretKeyRef: name: shared-secrets key: LIVEKIT_URL + - name: "STRIPE_SECRET_KEY" + valueFrom: + secretKeyRef: + name: shared-secrets + key: STRIPE_SECRET_KEY + - name: "STRIPE_WEBHOOK_SECRET" + valueFrom: + secretKeyRef: + name: shared-secrets + key: STRIPE_WEBHOOK_SECRET + - name: "STRIPE_PRICE_PRO_MONTHLY" + value: "price_1TM9abF9Z3lE6HEQJQmVsWES" + - name: "STRIPE_PRICE_PRO_ANNUAL" + value: "price_1TM9abF9Z3lE6HEQsToXflbq" + - name: "BILLING_SUCCESS_URL" + value: "llink://billing/success" + - name: "BILLING_CANCEL_URL" + value: "llink://billing/cancel" --- diff --git a/go/k8s/prod/orion.yaml b/go/k8s/prod/orion.yaml index 50b4e7f..9da0ff5 100644 --- a/go/k8s/prod/orion.yaml +++ b/go/k8s/prod/orion.yaml @@ -58,6 +58,30 @@ spec: secretKeyRef: name: shared-secrets key: LIVEKIT_URL + - name: "STRIPE_SECRET_KEY" + valueFrom: + secretKeyRef: + name: shared-secrets + key: STRIPE_SECRET_KEY + - name: "STRIPE_WEBHOOK_SECRET" + valueFrom: + secretKeyRef: + name: shared-secrets + key: STRIPE_WEBHOOK_SECRET + - name: "STRIPE_PRICE_PRO_MONTHLY" + valueFrom: + secretKeyRef: + name: shared-secrets + key: STRIPE_PRICE_PRO_MONTHLY + - name: "STRIPE_PRICE_PRO_ANNUAL" + valueFrom: + secretKeyRef: + name: shared-secrets + key: STRIPE_PRICE_PRO_ANNUAL + - name: "BILLING_SUCCESS_URL" + value: "llink://billing/success" + - name: "BILLING_CANCEL_URL" + value: "llink://billing/cancel" --- diff --git a/go/migrations/000013_network_billing.down.sql b/go/migrations/000013_network_billing.down.sql new file mode 100644 index 0000000..073d83b --- /dev/null +++ b/go/migrations/000013_network_billing.down.sql @@ -0,0 +1,2 @@ +DROP TABLE IF EXISTS network_subscriptions; +DROP TABLE IF EXISTS network_stripe_customers; diff --git a/go/migrations/000013_network_billing.up.sql b/go/migrations/000013_network_billing.up.sql new file mode 100644 index 0000000..034e143 --- /dev/null +++ b/go/migrations/000013_network_billing.up.sql @@ -0,0 +1,24 @@ +CREATE TABLE network_stripe_customers ( + network_id TEXT PRIMARY KEY, + stripe_customer_id TEXT NOT NULL UNIQUE, + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW() +); + +CREATE TABLE network_subscriptions ( + id TEXT PRIMARY KEY, + network_id TEXT NOT NULL REFERENCES network_stripe_customers(id) ON DELETE CASCADE, + stripe_customer_id TEXT NOT NULL, + status TEXT NOT NULL, + price_id TEXT NOT NULL, + cadence TEXT NOT NULL, + quantity INT NOT NULL, + cancel_at_period_end BOOLEAN NOT NULL DEFAULT FALSE, + current_period_start TIMESTAMPTZ NOT NULL, + current_period_end TIMESTAMPTZ NOT NULL, + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), + CONSTRAINT network_subscriptions_cadence_check CHECK (cadence IN ('monthly', 'annual')) +); + +CREATE UNIQUE INDEX idx_network_subscriptions_network_id + ON network_subscriptions (network_id); diff --git a/js/src/App.tsx b/js/src/App.tsx index 09999e1..c0ce124 100644 --- a/js/src/App.tsx +++ b/js/src/App.tsx @@ -14,6 +14,7 @@ import NetworkRoot from "@/features/network-root"; import ParticleViewResolver from "@/features/particles/particle-view-resolver"; import Layout from "@/features/layout"; import NetworkSettingsPage from "@/features/network-settings"; +import NetworkBillingPage from "@/features/network-billing"; import { Toaster } from "@/components/ui/sonner"; import { PusherProvider } from "@/lib/pusher-provider"; @@ -71,6 +72,7 @@ function AuthenticatedApp() { } /> } /> + } /> } /> diff --git a/js/src/api/client.ts b/js/src/api/client.ts index 24139ba..2a3c31f 100644 --- a/js/src/api/client.ts +++ b/js/src/api/client.ts @@ -2,18 +2,22 @@ import { appConfig } from "@/config/env"; import { useSessionStore } from "@/stores/session-store"; import type { z } from "zod"; import { + BillingStatusSchema, + CheckoutSessionResponseSchema, DepotObjectSchema, GetLivekitTokenResponseSchema, HumanSchema, ListInvitationsResponseSchema, ListNetworksResponseSchema, NetworkSchema, + PortalSessionResponseSchema, PrepareUploadResponseSchema, SignInResponseSchema, } from "./types"; import type { AcceptInvitationRequest, AddMembersRequest, + BillingCadence, CreateNetworkRequest, PrepareUploadRequest, RequestCodeRequest, @@ -216,6 +220,33 @@ class ApiClient { async getLivekitToken(networkId: string, streamId: string) { return this.request(GetLivekitTokenResponseSchema, "POST", "/livekit/token", { network_id: networkId, stream_id: streamId }); } + + // --- Billing (network admin only) --- + + async getNetworkBilling(networkId: string) { + return this.request( + BillingStatusSchema, + "GET", + `/networks/${networkId}/billing`, + ); + } + + async createCheckoutSession(networkId: string, cadence: BillingCadence) { + return this.request( + CheckoutSessionResponseSchema, + "POST", + `/networks/${networkId}/billing/checkout-session`, + { cadence }, + ); + } + + async createPortalSession(networkId: string) { + return this.request( + PortalSessionResponseSchema, + "POST", + `/networks/${networkId}/billing/portal-session`, + ); + } } export const apiClient = new ApiClient({ diff --git a/js/src/api/types.ts b/js/src/api/types.ts index 5e1ae03..ace151f 100644 --- a/js/src/api/types.ts +++ b/js/src/api/types.ts @@ -262,3 +262,45 @@ export const SignInResponseSchema = z.object({ token: z.string(), }); export type SignInResponse = z.infer; + +// --- Billing types --- + +export const BillingCadenceSchema = z.enum(["monthly", "annual"]); +export type BillingCadence = z.infer; + +export const NetworkPlanSchema = z.enum(["free", "pro"]); +export type NetworkPlan = z.infer; + +// Mirrors Stripe subscription.status plus "active" as the default free-tier value. +export const BillingPlanStatusSchema = z.enum([ + "active", + "trialing", + "past_due", + "canceled", + "incomplete", + "incomplete_expired", + "unpaid", +]); +export type BillingPlanStatus = z.infer; + +export const BillingStatusSchema = z.object({ + plan: NetworkPlanSchema, + plan_status: BillingPlanStatusSchema, + cadence: BillingCadenceSchema.nullable(), + seats: z.number().int(), + current_period_end: z.coerce.date().nullable(), + cancel_at_period_end: z.boolean(), + price_monthly_cents: z.number().int(), + price_annual_cents: z.number().int(), +}); +export type BillingStatus = z.infer; + +export const CheckoutSessionResponseSchema = z.object({ + url: z.string().url(), +}); +export type CheckoutSessionResponse = z.infer; + +export const PortalSessionResponseSchema = z.object({ + url: z.string().url(), +}); +export type PortalSessionResponse = z.infer; diff --git a/js/src/features/network-billing.tsx b/js/src/features/network-billing.tsx new file mode 100644 index 0000000..ac233e3 --- /dev/null +++ b/js/src/features/network-billing.tsx @@ -0,0 +1,350 @@ +import { useNavigate, useParams } from "react-router-dom"; +import { ArrowLeft, Check, ExternalLink } from "lucide-react"; +import { toast } from "sonner"; +import { Badge } from "@/components/ui/badge"; +import { Button } from "@/components/ui/button"; +import { + Card, + CardContent, + CardDescription, + CardFooter, + CardHeader, + CardTitle, +} from "@/components/ui/card"; +import { ScrollArea } from "@/components/ui/scroll-area"; +import { Separator } from "@/components/ui/separator"; +import { Muted } from "@/components/ui/typography"; +import { WindowControls } from "@/components/window-controls"; +import { useNetworks } from "@/hooks/use-networks"; +import { + useCreateCheckoutSession, + useCreatePortalSession, + useNetworkBilling, +} from "@/hooks/use-billing"; +import { useAuthStore } from "@/stores/auth-store"; +import type { BillingCadence, BillingStatus } from "@/api/types"; + +function formatCents(cents: number): string { + if (cents % 100 === 0) return `$${cents / 100}`; + return `$${(cents / 100).toFixed(2)}`; +} + +function formatDate(date: Date): string { + return date.toLocaleDateString(undefined, { + month: "long", + day: "numeric", + year: "numeric", + }); +} + +function PlanStatusBadge({ status }: { status: BillingStatus["plan_status"] }) { + if (status === "past_due") { + return Past due; + } + if (status === "canceled") { + return Canceled; + } + if (status === "trialing") { + return Trialing; + } + return null; +} + +function PricingCard({ + cadence, + pricePerSeatCents, + seats, + saveBadge, + billedNote, + onUpgrade, + isLoading, +}: { + cadence: BillingCadence; + pricePerSeatCents: number; + seats: number; + saveBadge?: string; + billedNote: string; + onUpgrade: () => void; + isLoading: boolean; +}) { + const label = cadence === "monthly" ? "Monthly" : "Annual"; + const perSeat = formatCents(pricePerSeatCents); + const total = formatCents(pricePerSeatCents * seats); + + return ( + + + + {label} + {saveBadge && ( + + {saveBadge} + + )} + + {billedNote} + + +
+ {perSeat} + / seat / month +
+ + {total} / month for {seats} {seats === 1 ? "seat" : "seats"} + +
+ + + +
+ ); +} + +function FreePlanView({ + networkId, + billing, +}: { + networkId: string; + billing: BillingStatus; +}) { + const createCheckout = useCreateCheckoutSession(networkId); + + const handleUpgrade = (cadence: BillingCadence) => { + createCheckout.mutate(cadence, { + onSuccess: ({ url }) => { + window.electronLink.openExternal(url); + }, + onError: (err) => { + toast.error(err.message || "Failed to start checkout"); + }, + }); + }; + + const features = [ + "Unlimited members", + "Priority support", + "All current and future features", + ]; + + const annualPerSeatMonthlyCents = Math.round(billing.price_annual_cents / 12); + const savingsPct = Math.round( + (1 - annualPerSeatMonthlyCents / billing.price_monthly_cents) * 100, + ); + + return ( +
+ + +
+
+ Llink Free + + Current plan · up to 50 particles per day + +
+ Free +
+
+
+ +
+

+ Upgrade to Pro +

+
    + {features.map((f) => ( +
  • + + {f} +
  • + ))} +
+ +
+ handleUpgrade("monthly")} + isLoading={createCheckout.isPending} + /> + 0 ? `Save ${savingsPct}%` : undefined} + billedNote="Billed annually" + onUpgrade={() => handleUpgrade("annual")} + isLoading={createCheckout.isPending} + /> +
+
+
+ ); +} + +function ProPlanView({ + networkId, + billing, +}: { + networkId: string; + billing: BillingStatus; +}) { + const createPortal = useCreatePortalSession(networkId); + + const handleManage = () => { + createPortal.mutate(undefined, { + onSuccess: ({ url }) => { + window.electronLink.openExternal(url); + }, + onError: (err) => { + toast.error(err.message || "Failed to open billing portal"); + }, + }); + }; + + const cadenceLabel = + billing.cadence === "annual" ? "Annual" : "Monthly"; + const perSeatCents = + billing.cadence === "annual" + ? Math.round(billing.price_annual_cents / 12) + : billing.price_monthly_cents; + const renewal = billing.current_period_end + ? formatDate(billing.current_period_end) + : null; + + return ( +
+ {billing.cancel_at_period_end && renewal && ( +
+ Your subscription is set to downgrade to Free on {renewal}. You can + reactivate from the billing portal before then. +
+ )} + + {billing.plan_status === "past_due" && ( +
+ Your last payment failed. Update your payment method in the billing + portal to keep Pro features active. +
+ )} + + + +
+
+ + Llink Pro + + + + {cadenceLabel} · {formatCents(perSeatCents)} per seat / month + +
+ Pro +
+
+ +
+
Seats
+
{billing.seats}
+
+ {billing.cancel_at_period_end ? "Ends" : "Renews"} +
+
{renewal ?? "—"}
+
+
+ + + +
+ + + Seats are synced automatically when you add or remove members. Changes + are prorated. + +
+ ); +} + +export default function NetworkBillingPage() { + const navigate = useNavigate(); + const { networkId } = useParams<{ networkId: string }>(); + const { data: networks } = useNetworks(); + const network = networks?.find((n) => n.id === networkId); + const currentUser = useAuthStore((s) => s.user); + const isAdmin = currentUser?.id === network?.admin_human.id; + + const { data: billing, isLoading, error } = useNetworkBilling( + isAdmin ? networkId : undefined, + ); + + const networkName = network?.name ?? "Network"; + + return ( +
+
+ + + {networkName} · Billing +
+
+ + + {!isAdmin ? ( +
+ + Only the network admin can manage billing. + +
+ ) : isLoading || !billing ? ( +
+ Loading billing... +
+ ) : error ? ( +
+ + Failed to load billing. Try again later. + +
+ ) : billing.plan === "pro" ? ( + + ) : ( + + )} + + +
+ + Payments are processed securely by Stripe. You can view invoices + and update payment methods from the subscription portal. + +
+
+
+ ); +} diff --git a/js/src/hooks/use-billing.ts b/js/src/hooks/use-billing.ts new file mode 100644 index 0000000..bdaf09f --- /dev/null +++ b/js/src/hooks/use-billing.ts @@ -0,0 +1,27 @@ +import { useQuery, useMutation } from "@tanstack/react-query"; +import { apiClient } from "@/api/client"; +import type { BillingCadence } from "@/api/types"; + +export function useNetworkBilling(networkId: string | undefined) { + return useQuery({ + queryKey: ["network-billing", networkId], + queryFn: () => apiClient.getNetworkBilling(networkId!), + enabled: !!networkId, + // Refetch on window focus so the UI catches up after the user returns + // from Stripe Checkout (webhook may land a second or two later). + refetchOnWindowFocus: true, + }); +} + +export function useCreateCheckoutSession(networkId: string) { + return useMutation({ + mutationFn: (cadence: BillingCadence) => + apiClient.createCheckoutSession(networkId, cadence), + }); +} + +export function useCreatePortalSession(networkId: string) { + return useMutation({ + mutationFn: () => apiClient.createPortalSession(networkId), + }); +}