Updated subscriptions
This commit is contained in:
@@ -61,7 +61,7 @@ func main() {
|
||||
notifConfigHandler := handlers.NewNotificationConfigHandler(notifConfigModel, userModel, cfg)
|
||||
emailDigestHandler := handlers.NewEmailDigestHandler(emailDigestSvc, notifConfigModel)
|
||||
statusHandler := handlers.NewStatusHandler(checkpointModel)
|
||||
stripeHandler := handlers.NewStripeHandler(userModel, cfg)
|
||||
stripeHandler := handlers.NewStripeHandler(userModel, alertRuleModel, cfg)
|
||||
accountHandler := handlers.NewAccountHandler(userModel, addressModel, cfg)
|
||||
|
||||
authenticate := middleware.Authenticate(userModel)
|
||||
|
||||
68
backend/cmd/subscription-sweep/main.go
Normal file
68
backend/cmd/subscription-sweep/main.go
Normal file
@@ -0,0 +1,68 @@
|
||||
// Package main runs a daily (cron-invoked) job: for paid tiers without an
|
||||
// active or trialling Stripe subscription, disable Firebase login and turn
|
||||
// off all alert rules until billing is restored via Stripe webhook / checkout.
|
||||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log"
|
||||
|
||||
"github.com/joho/godotenv"
|
||||
|
||||
"github.com/kjannette/koin-ping/backend/internal/config"
|
||||
"github.com/kjannette/koin-ping/backend/internal/database"
|
||||
"github.com/kjannette/koin-ping/backend/internal/firebase"
|
||||
"github.com/kjannette/koin-ping/backend/internal/models"
|
||||
)
|
||||
|
||||
func main() {
|
||||
_ = godotenv.Load()
|
||||
|
||||
cfg, err := config.Load()
|
||||
if err != nil {
|
||||
log.Fatalf("Failed to load config: %v", err)
|
||||
}
|
||||
|
||||
if err := firebase.Init(cfg.FirebaseProjectID); err != nil {
|
||||
log.Fatalf("Failed to initialize Firebase: %v", err)
|
||||
}
|
||||
|
||||
pool, connErr := database.Connect(cfg.DSN())
|
||||
if connErr != nil {
|
||||
log.Fatalf("Failed to connect to database: %v", connErr)
|
||||
}
|
||||
|
||||
defer database.Close()
|
||||
|
||||
ctx := context.Background()
|
||||
userModel := models.NewUserModel(pool)
|
||||
alertModel := models.NewAlertRuleModel(pool)
|
||||
|
||||
users, listErr := userModel.ListPaidUsersWithoutActiveSubscription(ctx)
|
||||
if listErr != nil {
|
||||
log.Fatalf("Failed to list lapsed subscriptions: %v", listErr)
|
||||
}
|
||||
|
||||
if len(users) == 0 {
|
||||
log.Println("Subscription sweep: no lapsed paid users")
|
||||
return
|
||||
}
|
||||
|
||||
for _, u := range users {
|
||||
if disableErr := firebase.SetUserDisabled(ctx, u.FirebaseUID, true); disableErr != nil {
|
||||
log.Printf("Subscription sweep: Firebase disable failed user %s: %v", u.ID, disableErr)
|
||||
continue
|
||||
}
|
||||
|
||||
n, rulesErr := alertModel.DisableAllForUser(ctx, u.ID)
|
||||
if rulesErr != nil {
|
||||
log.Printf("Subscription sweep: disabling alerts failed for user %s: %v", u.ID, rulesErr)
|
||||
continue
|
||||
}
|
||||
|
||||
log.Printf(
|
||||
"Subscription sweep: suspended user %s (%s tier, status=%s), disabled %d alert rules",
|
||||
u.ID, u.SubscriptionTier, u.SubscriptionStatus, n,
|
||||
)
|
||||
}
|
||||
}
|
||||
@@ -32,12 +32,14 @@ type Config struct {
|
||||
ResendAPIKey string
|
||||
EmailFrom string
|
||||
DigestIntervalHours int
|
||||
StripeSecretKey string
|
||||
StripeWebhookSecret string
|
||||
StripePriceIDPremium string
|
||||
StripePriceIDPro string
|
||||
StripePublishableKey string
|
||||
FrontendURL string
|
||||
StripeSecretKey string
|
||||
StripeWebhookSecret string
|
||||
StripePriceIDPremium string
|
||||
StripePriceIDPro string
|
||||
StripePriceIDPremiumAnnual string
|
||||
StripePriceIDProAnnual string
|
||||
StripePublishableKey string
|
||||
FrontendURL string
|
||||
}
|
||||
|
||||
// Load reads configuration from environment variables and returns a Config.
|
||||
@@ -58,12 +60,14 @@ func Load() (*Config, error) {
|
||||
ResendAPIKey: os.Getenv("RESEND_API_KEY"),
|
||||
EmailFrom: getEnv("EMAIL_FROM", "Koin Ping <alerts@koinping.com>"),
|
||||
DigestIntervalHours: getEnvInt("DIGEST_INTERVAL_HOURS", defaultDigestIntervalHours),
|
||||
StripeSecretKey: os.Getenv("STRIPE_SECRET_KEY"),
|
||||
StripeWebhookSecret: os.Getenv("STRIPE_WEBHOOK_SECRET"),
|
||||
StripePriceIDPremium: os.Getenv("STRIPE_PRICE_ID_PREMIUM"),
|
||||
StripePriceIDPro: os.Getenv("STRIPE_PRICE_ID_PRO"),
|
||||
StripePublishableKey: os.Getenv("STRIPE_PUBLISHABLE_KEY"),
|
||||
FrontendURL: getEnv("FRONTEND_URL", "http://localhost:3000"),
|
||||
StripeSecretKey: os.Getenv("STRIPE_SECRET_KEY"),
|
||||
StripeWebhookSecret: os.Getenv("STRIPE_WEBHOOK_SECRET"),
|
||||
StripePriceIDPremium: os.Getenv("STRIPE_PRICE_ID_PREMIUM"),
|
||||
StripePriceIDPro: os.Getenv("STRIPE_PRICE_ID_PRO"),
|
||||
StripePriceIDPremiumAnnual: os.Getenv("STRIPE_PRICE_ID_PREMIUM_ANNUAL"),
|
||||
StripePriceIDProAnnual: os.Getenv("STRIPE_PRICE_ID_PRO_ANNUAL"),
|
||||
StripePublishableKey: os.Getenv("STRIPE_PUBLISHABLE_KEY"),
|
||||
FrontendURL: getEnv("FRONTEND_URL", "http://localhost:3000"),
|
||||
}
|
||||
|
||||
if cfg.PollIntervalMS < minPollIntervalMS {
|
||||
@@ -97,10 +101,13 @@ func (c *Config) DSN() string {
|
||||
// TierForPriceID maps a Stripe price ID back to the corresponding
|
||||
// subscription tier. Returns empty string if the price is unrecognised.
|
||||
func (c *Config) TierForPriceID(priceID string) string {
|
||||
if priceID == "" {
|
||||
return ""
|
||||
}
|
||||
switch priceID {
|
||||
case c.StripePriceIDPremium:
|
||||
case c.StripePriceIDPremium, c.StripePriceIDPremiumAnnual:
|
||||
return "premium"
|
||||
case c.StripePriceIDPro:
|
||||
case c.StripePriceIDPro, c.StripePriceIDProAnnual:
|
||||
return "pro"
|
||||
default:
|
||||
return ""
|
||||
|
||||
@@ -6,8 +6,10 @@ func TestTierForPriceID(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
cfg := &Config{
|
||||
StripePriceIDPremium: "price_premium_123",
|
||||
StripePriceIDPro: "price_pro_456",
|
||||
StripePriceIDPremium: "price_premium_123",
|
||||
StripePriceIDPro: "price_pro_456",
|
||||
StripePriceIDPremiumAnnual: "price_premium_yr",
|
||||
StripePriceIDProAnnual: "price_pro_yr",
|
||||
}
|
||||
|
||||
tests := []struct {
|
||||
@@ -16,6 +18,8 @@ func TestTierForPriceID(t *testing.T) {
|
||||
}{
|
||||
{"price_premium_123", "premium"},
|
||||
{"price_pro_456", "pro"},
|
||||
{"price_premium_yr", "premium"},
|
||||
{"price_pro_yr", "pro"},
|
||||
{"price_unknown", ""},
|
||||
{"", ""},
|
||||
}
|
||||
|
||||
22
backend/internal/firebase/users.go
Normal file
22
backend/internal/firebase/users.go
Normal file
@@ -0,0 +1,22 @@
|
||||
package firebase
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
"firebase.google.com/go/v4/auth"
|
||||
)
|
||||
|
||||
// SetUserDisabled updates the Firebase user record's disabled flag.
|
||||
func SetUserDisabled(ctx context.Context, uid string, disabled bool) error {
|
||||
if authClient == nil {
|
||||
return fmt.Errorf("firebase auth not initialized") //nolint:err113
|
||||
}
|
||||
if uid == "" {
|
||||
return nil
|
||||
}
|
||||
|
||||
params := (&auth.UserToUpdate{}).Disabled(disabled)
|
||||
_, err := authClient.UpdateUser(ctx, uid, params)
|
||||
return err
|
||||
}
|
||||
@@ -38,8 +38,8 @@ type accountResponse struct {
|
||||
|
||||
var tierPlanLabels = map[domain.SubscriptionTier]string{ //nolint:gochecknoglobals
|
||||
domain.TierFree: "Free Trial",
|
||||
domain.TierPremium: "Premium / $1.99 mo",
|
||||
domain.TierPro: "Pro / $11.99 mo",
|
||||
domain.TierPremium: "Premium / $8.78 mo",
|
||||
domain.TierPro: "Pro / $16.78 mo",
|
||||
}
|
||||
|
||||
func (h *AccountHandler) GetAccount(w http.ResponseWriter, r *http.Request) {
|
||||
|
||||
@@ -1,36 +1,86 @@
|
||||
package handlers
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"log"
|
||||
"net/http"
|
||||
|
||||
"github.com/stripe/stripe-go/v82"
|
||||
portalsession "github.com/stripe/stripe-go/v82/billingportal/session"
|
||||
checkoutsession "github.com/stripe/stripe-go/v82/checkout/session"
|
||||
"github.com/stripe/stripe-go/v82/webhook"
|
||||
|
||||
kpfirebase "github.com/kjannette/koin-ping/backend/internal/firebase"
|
||||
"github.com/kjannette/koin-ping/backend/internal/config"
|
||||
"github.com/kjannette/koin-ping/backend/internal/domain"
|
||||
"github.com/kjannette/koin-ping/backend/internal/middleware"
|
||||
"github.com/kjannette/koin-ping/backend/internal/models"
|
||||
"github.com/stripe/stripe-go/v82"
|
||||
portalsession "github.com/stripe/stripe-go/v82/billingportal/session"
|
||||
checkoutsession "github.com/stripe/stripe-go/v82/checkout/session"
|
||||
"github.com/stripe/stripe-go/v82/webhook"
|
||||
)
|
||||
|
||||
const webhookMaxBodyBytes = 65536
|
||||
|
||||
type StripeHandler struct {
|
||||
users *models.UserModel
|
||||
cfg *config.Config
|
||||
users *models.UserModel
|
||||
alerts *models.AlertRuleModel
|
||||
cfg *config.Config
|
||||
}
|
||||
|
||||
func NewStripeHandler(users *models.UserModel, cfg *config.Config) *StripeHandler {
|
||||
func NewStripeHandler(users *models.UserModel, alerts *models.AlertRuleModel, cfg *config.Config) *StripeHandler {
|
||||
stripe.Key = cfg.StripeSecretKey
|
||||
return &StripeHandler{users: users, cfg: cfg}
|
||||
return &StripeHandler{users: users, alerts: alerts, cfg: cfg}
|
||||
}
|
||||
|
||||
func (h *StripeHandler) priceIDForTier(tier domain.SubscriptionTier) (string, error) {
|
||||
func (h *StripeHandler) ensureUserFirebaseAndAlertsEnabled(ctx context.Context, localUserID string) {
|
||||
user, err := h.users.GetByID(ctx, localUserID)
|
||||
if err != nil || user == nil || user.FirebaseUID == "" {
|
||||
return
|
||||
}
|
||||
if firebaseErr := kpfirebase.SetUserDisabled(ctx, user.FirebaseUID, false); firebaseErr != nil {
|
||||
log.Printf("Billing access restore: firebase enable failed for user %s: %v", localUserID, firebaseErr)
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
n, alertsErr := h.alerts.EnableAllForUser(ctx, localUserID)
|
||||
if alertsErr != nil {
|
||||
log.Printf("Billing access restore: enable alerts failed for user %s: %v", localUserID, alertsErr)
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
log.Printf("Billing access restored: user %s, %d alert rules enabled", localUserID, n)
|
||||
}
|
||||
|
||||
func (h *StripeHandler) restorePaidSubscriptionAccess(ctx context.Context, stripeCustomerID, status string) {
|
||||
if stripeCustomerID == "" || (status != "active" && status != "trialing") {
|
||||
return
|
||||
}
|
||||
|
||||
u, err := h.users.GetByStripeCustomerID(ctx, stripeCustomerID)
|
||||
if err != nil || u == nil {
|
||||
if err != nil {
|
||||
log.Printf("restorePaidSubscriptionAccess: lookup %s: %v", stripeCustomerID, err)
|
||||
}
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
h.ensureUserFirebaseAndAlertsEnabled(ctx, u.ID)
|
||||
}
|
||||
|
||||
func (h *StripeHandler) priceIDForTier(tier domain.SubscriptionTier, interval string) (string, error) {
|
||||
if interval == "annual" {
|
||||
switch tier {
|
||||
case domain.TierPremium:
|
||||
return h.cfg.StripePriceIDPremiumAnnual, nil
|
||||
case domain.TierPro:
|
||||
return h.cfg.StripePriceIDProAnnual, nil
|
||||
default:
|
||||
return "", fmt.Errorf("no Stripe price for tier %q", tier) //nolint:err113
|
||||
}
|
||||
}
|
||||
switch tier {
|
||||
case domain.TierPremium:
|
||||
return h.cfg.StripePriceIDPremium, nil
|
||||
@@ -46,7 +96,8 @@ func (h *StripeHandler) CreateCheckoutSession(w http.ResponseWriter, r *http.Req
|
||||
userID := middleware.GetUserID(r.Context())
|
||||
|
||||
var body struct {
|
||||
Tier string `json:"tier"`
|
||||
Tier string `json:"tier"`
|
||||
Interval string `json:"interval"`
|
||||
}
|
||||
if err := json.NewDecoder(r.Body).Decode(&body); err != nil {
|
||||
writeError(w, http.StatusBadRequest, "BAD_REQUEST", "Invalid request body")
|
||||
@@ -56,6 +107,9 @@ func (h *StripeHandler) CreateCheckoutSession(w http.ResponseWriter, r *http.Req
|
||||
if body.Tier == "" {
|
||||
body.Tier = "premium"
|
||||
}
|
||||
if body.Interval == "" {
|
||||
body.Interval = "annual"
|
||||
}
|
||||
|
||||
tier := domain.SubscriptionTier(body.Tier)
|
||||
if tier != domain.TierPremium && tier != domain.TierPro {
|
||||
@@ -63,7 +117,7 @@ func (h *StripeHandler) CreateCheckoutSession(w http.ResponseWriter, r *http.Req
|
||||
return
|
||||
}
|
||||
|
||||
priceID, err := h.priceIDForTier(tier)
|
||||
priceID, err := h.priceIDForTier(tier, body.Interval)
|
||||
if err != nil {
|
||||
writeError(w, http.StatusBadRequest, "VALIDATION_ERROR", err.Error())
|
||||
return
|
||||
@@ -177,6 +231,8 @@ func (h *StripeHandler) VerifyCheckoutSession(w http.ResponseWriter, r *http.Req
|
||||
if subscriptionID != "" && customerID != "" {
|
||||
if err := h.users.ActivateSubscription(r.Context(), customerID, subscriptionID, "active", tier); err != nil {
|
||||
log.Printf("VerifyCheckout: failed to activate subscription: %v", err)
|
||||
} else {
|
||||
h.restorePaidSubscriptionAccess(r.Context(), customerID, "active")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -197,6 +253,8 @@ func (h *StripeHandler) ActivateFreeTier(w http.ResponseWriter, r *http.Request)
|
||||
return
|
||||
}
|
||||
|
||||
h.ensureUserFirebaseAndAlertsEnabled(r.Context(), userID)
|
||||
|
||||
log.Printf("Free tier activated for user %s", userID)
|
||||
writeJSON(w, http.StatusOK, map[string]string{
|
||||
"subscription_status": "active",
|
||||
@@ -209,8 +267,9 @@ func (h *StripeHandler) ActivateFreeTier(w http.ResponseWriter, r *http.Request)
|
||||
// The Firebase account is created on the frontend only after payment succeeds.
|
||||
func (h *StripeHandler) CreateOnboardingCheckout(w http.ResponseWriter, r *http.Request) {
|
||||
var body struct {
|
||||
Email string `json:"email"`
|
||||
Tier string `json:"tier"`
|
||||
Email string `json:"email"`
|
||||
Tier string `json:"tier"`
|
||||
Interval string `json:"interval"`
|
||||
}
|
||||
if err := json.NewDecoder(r.Body).Decode(&body); err != nil {
|
||||
writeError(w, http.StatusBadRequest, "BAD_REQUEST", "Invalid request body")
|
||||
@@ -225,6 +284,9 @@ func (h *StripeHandler) CreateOnboardingCheckout(w http.ResponseWriter, r *http.
|
||||
if body.Tier == "" {
|
||||
body.Tier = "premium"
|
||||
}
|
||||
if body.Interval == "" {
|
||||
body.Interval = "annual"
|
||||
}
|
||||
|
||||
tier := domain.SubscriptionTier(body.Tier)
|
||||
if tier != domain.TierPremium && tier != domain.TierPro {
|
||||
@@ -232,7 +294,7 @@ func (h *StripeHandler) CreateOnboardingCheckout(w http.ResponseWriter, r *http.
|
||||
return
|
||||
}
|
||||
|
||||
priceID, err := h.priceIDForTier(tier)
|
||||
priceID, err := h.priceIDForTier(tier, body.Interval)
|
||||
if err != nil {
|
||||
writeError(w, http.StatusBadRequest, "VALIDATION_ERROR", err.Error())
|
||||
return
|
||||
@@ -362,6 +424,8 @@ func (h *StripeHandler) handleCheckoutCompleted(r *http.Request, event stripe.Ev
|
||||
if subscriptionID != "" && customerID != "" {
|
||||
if err := h.users.ActivateSubscription(r.Context(), customerID, subscriptionID, "active", tier); err != nil {
|
||||
log.Printf("Failed to activate subscription: %v", err)
|
||||
} else {
|
||||
h.restorePaidSubscriptionAccess(r.Context(), customerID, "active")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -401,8 +465,12 @@ func (h *StripeHandler) handleSubscriptionUpdated(r *http.Request, event stripe.
|
||||
|
||||
if err := h.users.ActivateSubscription(r.Context(), customerID, sub.ID, status, tier); err != nil {
|
||||
log.Printf("Failed to update subscription: %v", err)
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
h.restorePaidSubscriptionAccess(r.Context(), customerID, status)
|
||||
|
||||
log.Printf("Subscription %s updated to %s (tier %s) for customer %s", sub.ID, status, tier, customerID)
|
||||
}
|
||||
|
||||
|
||||
@@ -140,3 +140,33 @@ func (m *AlertRuleModel) Remove(ctx context.Context, id int) (bool, error) {
|
||||
}
|
||||
return tag.RowsAffected() > 0, nil
|
||||
}
|
||||
|
||||
// DisableAllForUser sets enabled = false on every alert rule owned by addresses of userID.
|
||||
func (m *AlertRuleModel) DisableAllForUser(ctx context.Context, userID string) (int64, error) {
|
||||
tag, err := m.pool.Exec(ctx,
|
||||
`UPDATE alert_rules ar
|
||||
SET enabled = FALSE
|
||||
FROM addresses a
|
||||
WHERE ar.address_id = a.id AND a.user_id = $1`,
|
||||
userID,
|
||||
)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return tag.RowsAffected(), nil
|
||||
}
|
||||
|
||||
// EnableAllForUser sets enabled = true on every alert rule owned by addresses of userID.
|
||||
func (m *AlertRuleModel) EnableAllForUser(ctx context.Context, userID string) (int64, error) {
|
||||
tag, err := m.pool.Exec(ctx,
|
||||
`UPDATE alert_rules ar
|
||||
SET enabled = TRUE
|
||||
FROM addresses a
|
||||
WHERE ar.address_id = a.id AND a.user_id = $1`,
|
||||
userID,
|
||||
)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return tag.RowsAffected(), nil
|
||||
}
|
||||
|
||||
@@ -130,6 +130,46 @@ func (m *UserModel) GetByID(ctx context.Context, id string) (*domain.User, error
|
||||
return scanUser(row)
|
||||
}
|
||||
|
||||
// GetByStripeCustomerID loads a user by their Stripe Customer ID if set.
|
||||
func (m *UserModel) GetByStripeCustomerID(ctx context.Context, stripeCustomerID string) (*domain.User, error) {
|
||||
row := m.pool.QueryRow(ctx,
|
||||
`SELECT `+userColumns+` FROM users WHERE stripe_customer_id = $1`,
|
||||
stripeCustomerID,
|
||||
)
|
||||
return scanUser(row)
|
||||
}
|
||||
|
||||
// ListPaidUsersWithoutActiveSubscription finds paid-tier rows whose Stripe
|
||||
// subscription is not active or trialing. Used by the periodic billing sweep.
|
||||
func (m *UserModel) ListPaidUsersWithoutActiveSubscription(ctx context.Context) ([]domain.User, error) {
|
||||
rows, err := m.pool.Query(ctx,
|
||||
`SELECT `+userColumns+` FROM users
|
||||
WHERE subscription_tier IN ('premium', 'pro')
|
||||
AND subscription_status NOT IN ('active', 'trialing')
|
||||
AND COALESCE(trim(firebase_uid), '') <> ''`,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
out := []domain.User{}
|
||||
for rows.Next() {
|
||||
var u domain.User
|
||||
rowErr := rows.Scan(
|
||||
&u.ID, &u.FirebaseUID, &u.Email, &u.DisplayName,
|
||||
&u.StripeCustomerID, &u.StripeSubscriptionID, &u.SubscriptionStatus,
|
||||
&u.SubscriptionTier, &u.SubscriptionCreatedAt, &u.CreatedAt, &u.UpdatedAt,
|
||||
)
|
||||
if rowErr != nil {
|
||||
return nil, rowErr
|
||||
}
|
||||
out = append(out, u)
|
||||
}
|
||||
|
||||
return out, rows.Err()
|
||||
}
|
||||
|
||||
func (m *UserModel) UpdateStripeCustomer(ctx context.Context, userID, stripeCustomerID string) error {
|
||||
_, err := m.pool.Exec(ctx,
|
||||
`UPDATE users SET stripe_customer_id = $2, updated_at = NOW() WHERE id = $1`,
|
||||
|
||||
Reference in New Issue
Block a user