From f9fa7def2b6b7a561bd761c6b8cf1743e99cc924 Mon Sep 17 00:00:00 2001 From: KS Jannette Date: Tue, 10 Mar 2026 15:22:48 -0400 Subject: [PATCH 1/5] add subscription tiers --- backend/cmd/api/main.go | 10 +- .../migrations/009_add_subscription_tier.sql | 1 + backend/infra/schema.sql | 1 + backend/internal/config/config.go | 20 +- backend/internal/domain/types.go | 80 ++++++- backend/internal/handlers/account.go | 56 +++-- backend/internal/handlers/address.go | 28 ++- backend/internal/handlers/alert_rule.go | 27 ++- .../internal/handlers/notification_config.go | 25 ++- backend/internal/handlers/stripe.go | 90 +++++++- backend/internal/middleware/auth.go | 9 + backend/internal/models/address.go | 9 + backend/internal/models/alert_rule.go | 9 + backend/internal/models/user.go | 30 ++- frontend/src/api/stripe.jsx | 18 +- frontend/src/components/TierPicker.css | 120 ++++++++++ frontend/src/components/TierPicker.jsx | 94 ++++++++ frontend/src/components/UpgradeBanner.css | 32 +++ frontend/src/components/UpgradeBanner.jsx | 18 ++ frontend/src/contexts/AuthContext.jsx | 54 +++-- frontend/src/pages/addresses/Addresses.jsx | 23 +- frontend/src/pages/alerts/Alerts.css | 14 ++ frontend/src/pages/alerts/Alerts.jsx | 62 +++++- frontend/src/pages/subscribe/Subscribe.css | 27 +++ frontend/src/pages/subscribe/Subscribe.jsx | 205 +++++++++++++----- frontend/src/pages/user_account/Account.jsx | 42 +++- 26 files changed, 944 insertions(+), 160 deletions(-) create mode 100644 backend/infra/migrations/009_add_subscription_tier.sql create mode 100644 frontend/src/components/TierPicker.css create mode 100644 frontend/src/components/TierPicker.jsx create mode 100644 frontend/src/components/UpgradeBanner.css create mode 100644 frontend/src/components/UpgradeBanner.jsx diff --git a/backend/cmd/api/main.go b/backend/cmd/api/main.go index 84c6d4d..a403a4b 100644 --- a/backend/cmd/api/main.go +++ b/backend/cmd/api/main.go @@ -55,14 +55,14 @@ func main() { cfg.ResendAPIKey, cfg.EmailFrom, alertEventModel, notifConfigModel, ) - addressHandler := handlers.NewAddressHandler(addressModel) - alertRuleHandler := handlers.NewAlertRuleHandler(alertRuleModel, addressModel) + addressHandler := handlers.NewAddressHandler(addressModel, userModel) + alertRuleHandler := handlers.NewAlertRuleHandler(alertRuleModel, addressModel, userModel) alertEventHandler := handlers.NewAlertEventHandler(alertEventModel) - notifConfigHandler := handlers.NewNotificationConfigHandler(notifConfigModel, cfg) + notifConfigHandler := handlers.NewNotificationConfigHandler(notifConfigModel, userModel, cfg) emailDigestHandler := handlers.NewEmailDigestHandler(emailDigestSvc, notifConfigModel) statusHandler := handlers.NewStatusHandler(checkpointModel) stripeHandler := handlers.NewStripeHandler(userModel, cfg) - accountHandler := handlers.NewAccountHandler(userModel, cfg) + accountHandler := handlers.NewAccountHandler(userModel, addressModel, cfg) authenticate := middleware.Authenticate(userModel) requireSub := middleware.RequireSubscription(userModel) @@ -91,6 +91,8 @@ func main() { authenticate(http.HandlerFunc(stripeHandler.VerifyCheckoutSession))) mux.Handle("POST "+b+"/stripe/create-portal-session", authenticate(http.HandlerFunc(stripeHandler.CreatePortalSession))) + mux.Handle("POST "+b+"/stripe/activate-free", + authenticate(http.HandlerFunc(stripeHandler.ActivateFreeTier))) // Account route (auth required, NO subscription required) mux.Handle("GET "+b+"/user/account", diff --git a/backend/infra/migrations/009_add_subscription_tier.sql b/backend/infra/migrations/009_add_subscription_tier.sql new file mode 100644 index 0000000..92ce8fd --- /dev/null +++ b/backend/infra/migrations/009_add_subscription_tier.sql @@ -0,0 +1 @@ +ALTER TABLE users ADD COLUMN IF NOT EXISTS subscription_tier VARCHAR(20) DEFAULT 'free'; diff --git a/backend/infra/schema.sql b/backend/infra/schema.sql index 9b02bec..31346c8 100644 --- a/backend/infra/schema.sql +++ b/backend/infra/schema.sql @@ -14,6 +14,7 @@ CREATE TABLE users ( stripe_customer_id VARCHAR(255), stripe_subscription_id VARCHAR(255), subscription_status VARCHAR(50) DEFAULT 'none', + subscription_tier VARCHAR(20) DEFAULT 'free', subscription_created_at TIMESTAMP, created_at TIMESTAMP DEFAULT NOW(), updated_at TIMESTAMP DEFAULT NOW() diff --git a/backend/internal/config/config.go b/backend/internal/config/config.go index 9bb4d22..ff4cd51 100644 --- a/backend/internal/config/config.go +++ b/backend/internal/config/config.go @@ -32,11 +32,12 @@ type Config struct { ResendAPIKey string EmailFrom string DigestIntervalHours int - StripeSecretKey string - StripeWebhookSecret string - StripePriceID string - StripePublishableKey string - FrontendURL string + StripeSecretKey string + StripeWebhookSecret string + StripePriceIDPremium string + StripePriceIDPro string + StripePublishableKey string + FrontendURL string } // Load reads configuration from environment variables and returns a Config. @@ -57,10 +58,11 @@ func Load() (*Config, error) { ResendAPIKey: os.Getenv("RESEND_API_KEY"), EmailFrom: getEnv("EMAIL_FROM", "Koin Ping "), DigestIntervalHours: getEnvInt("DIGEST_INTERVAL_HOURS", defaultDigestIntervalHours), - StripeSecretKey: os.Getenv("STRIPE_SECRET_KEY"), - StripeWebhookSecret: os.Getenv("STRIPE_WEBHOOK_SECRET"), - StripePriceID: os.Getenv("STRIPE_PRICE_ID"), - StripePublishableKey: os.Getenv("STRIPE_PUBLISHABLE_KEY"), + 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"), } diff --git a/backend/internal/domain/types.go b/backend/internal/domain/types.go index 5a18e87..09e0299 100644 --- a/backend/internal/domain/types.go +++ b/backend/internal/domain/types.go @@ -2,17 +2,77 @@ package domain import "time" +type SubscriptionTier string + +const ( + TierFree SubscriptionTier = "free" + TierPremium SubscriptionTier = "premium" + TierPro SubscriptionTier = "pro" +) + +func IsValidTier(t string) bool { + switch SubscriptionTier(t) { + case TierFree, TierPremium, TierPro: + return true + } + return false +} + +const unlimitedLimit = -1 + +type TierLimits struct { + MaxAddresses int `json:"max_addresses"` + MaxAlertTypes int `json:"max_alert_types"` + AllowedChannels []string `json:"allowed_channels"` +} + +func GetTierLimits(tier SubscriptionTier) TierLimits { + switch tier { + case TierPremium: + return TierLimits{ + MaxAddresses: 3, + MaxAlertTypes: 2, + AllowedChannels: []string{"email", "discord", "telegram"}, + } + case TierPro: + return TierLimits{ + MaxAddresses: unlimitedLimit, + MaxAlertTypes: unlimitedLimit, + AllowedChannels: []string{"email", "discord", "telegram", "slack"}, + } + default: + return TierLimits{ + MaxAddresses: 1, + MaxAlertTypes: 1, + AllowedChannels: []string{"email"}, + } + } +} + +func (l TierLimits) IsUnlimitedAddresses() bool { return l.MaxAddresses == unlimitedLimit } +func (l TierLimits) IsUnlimitedAlertTypes() bool { return l.MaxAlertTypes == unlimitedLimit } + +func (l TierLimits) ChannelAllowed(channel string) bool { + for _, c := range l.AllowedChannels { + if c == channel { + return true + } + } + return false +} + type User struct { - ID string `json:"id"` - FirebaseUID string `json:"-"` - Email string `json:"email"` - DisplayName *string `json:"display_name"` //nolint:tagliatelle - StripeCustomerID *string `json:"-"` - StripeSubscriptionID *string `json:"-"` - SubscriptionStatus string `json:"subscription_status"` //nolint:tagliatelle - SubscriptionCreatedAt *time.Time `json:"subscription_created_at,omitempty"` //nolint:tagliatelle - CreatedAt time.Time `json:"created_at"` //nolint:tagliatelle - UpdatedAt time.Time `json:"updated_at"` //nolint:tagliatelle + ID string `json:"id"` + FirebaseUID string `json:"-"` + Email string `json:"email"` + DisplayName *string `json:"display_name"` //nolint:tagliatelle + StripeCustomerID *string `json:"-"` + StripeSubscriptionID *string `json:"-"` + SubscriptionStatus string `json:"subscription_status"` //nolint:tagliatelle + SubscriptionTier SubscriptionTier `json:"subscription_tier"` //nolint:tagliatelle + SubscriptionCreatedAt *time.Time `json:"subscription_created_at,omitempty"` //nolint:tagliatelle + CreatedAt time.Time `json:"created_at"` //nolint:tagliatelle + UpdatedAt time.Time `json:"updated_at"` //nolint:tagliatelle } type Address struct { diff --git a/backend/internal/handlers/account.go b/backend/internal/handlers/account.go index 030a4a3..5bcdb42 100644 --- a/backend/internal/handlers/account.go +++ b/backend/internal/handlers/account.go @@ -6,29 +6,40 @@ import ( "time" "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" ) type AccountHandler struct { - users *models.UserModel - cfg *config.Config + users *models.UserModel + addresses *models.AddressModel + cfg *config.Config } -func NewAccountHandler(users *models.UserModel, cfg *config.Config) *AccountHandler { - return &AccountHandler{users: users, cfg: cfg} +func NewAccountHandler(users *models.UserModel, addresses *models.AddressModel, cfg *config.Config) *AccountHandler { + return &AccountHandler{users: users, addresses: addresses, cfg: cfg} } type accountResponse struct { - UserID string `json:"user_id"` - Email string `json:"email"` - UserName string `json:"user_name"` - SubscriptionStatus string `json:"subscription_status"` - SubscriptionPlan string `json:"subscription_plan"` - MemberSince *string `json:"member_since,omitempty"` - NextBillingDate *string `json:"next_billing_date,omitempty"` - CancelAtPeriodEnd bool `json:"cancel_at_period_end"` - PeriodEndDate *string `json:"period_end_date,omitempty"` + UserID string `json:"user_id"` + Email string `json:"email"` + UserName string `json:"user_name"` + SubscriptionStatus string `json:"subscription_status"` + SubscriptionTier string `json:"subscription_tier"` + SubscriptionPlan string `json:"subscription_plan"` + TierLimits domain.TierLimits `json:"tier_limits"` + AddressCount int `json:"address_count"` + MemberSince *string `json:"member_since,omitempty"` + NextBillingDate *string `json:"next_billing_date,omitempty"` + CancelAtPeriodEnd bool `json:"cancel_at_period_end"` + PeriodEndDate *string `json:"period_end_date,omitempty"` +} + +var tierPlanLabels = map[domain.SubscriptionTier]string{ //nolint:gochecknoglobals + domain.TierFree: "Free Trial", + domain.TierPremium: "Premium / $1.99 mo", + domain.TierPro: "Pro / $11.99 mo", } func (h *AccountHandler) GetAccount(w http.ResponseWriter, r *http.Request) { @@ -42,12 +53,26 @@ func (h *AccountHandler) GetAccount(w http.ResponseWriter, r *http.Request) { return } + addrCount, err := h.addresses.CountByUser(r.Context(), userID) + if err != nil { + log.Printf("Account: failed to count addresses for %s: %v", userID, err) + addrCount = 0 + } + + planLabel := tierPlanLabels[user.SubscriptionTier] + if planLabel == "" { + planLabel = "Free Trial" + } + resp := accountResponse{ UserID: user.ID, Email: email, UserName: email, SubscriptionStatus: user.SubscriptionStatus, - SubscriptionPlan: "monthly/$1.99", + SubscriptionTier: string(user.SubscriptionTier), + SubscriptionPlan: planLabel, + TierLimits: domain.GetTierLimits(user.SubscriptionTier), + AddressCount: addrCount, } if user.SubscriptionCreatedAt != nil { @@ -55,9 +80,6 @@ func (h *AccountHandler) GetAccount(w http.ResponseWriter, r *http.Request) { resp.MemberSince = &t } - // MOCKED: NextBillingDate, CancelAtPeriodEnd, PeriodEndDate. - // TODO: stripe-go v82 removed Subscription.CurrentPeriodEnd / CancelAtPeriodEnd. - // Research SubscriptionItem.CurrentPeriodEnd or use Stripe REST API directly. resp.NextBillingDate = nil resp.CancelAtPeriodEnd = false resp.PeriodEndDate = nil diff --git a/backend/internal/handlers/address.go b/backend/internal/handlers/address.go index 63e1b6b..28e6a10 100644 --- a/backend/internal/handlers/address.go +++ b/backend/internal/handlers/address.go @@ -3,6 +3,7 @@ package handlers import ( "encoding/json" + "fmt" "log" "net/http" "regexp" @@ -17,10 +18,11 @@ var ethAddressRe = regexp.MustCompile(`^0x[a-fA-F0-9]{40}$`) type AddressHandler struct { addresses *models.AddressModel + users *models.UserModel } -func NewAddressHandler(addresses *models.AddressModel) *AddressHandler { - return &AddressHandler{addresses: addresses} +func NewAddressHandler(addresses *models.AddressModel, users *models.UserModel) *AddressHandler { + return &AddressHandler{addresses: addresses, users: users} } func (h *AddressHandler) Create(w http.ResponseWriter, r *http.Request) { @@ -49,6 +51,28 @@ func (h *AddressHandler) Create(w http.ResponseWriter, r *http.Request) { return } + user, err := h.users.GetByID(r.Context(), userID) + if err != nil || user == nil { + log.Printf("Failed to get user %s for tier check: %v", userID, err) + writeError(w, http.StatusInternalServerError, "INTERNAL_ERROR", "Failed to verify account") + return + } + + limits := domain.GetTierLimits(user.SubscriptionTier) + if !limits.IsUnlimitedAddresses() { + count, err := h.addresses.CountByUser(r.Context(), userID) + if err != nil { + log.Printf("Failed to count addresses for user %s: %v", userID, err) + writeError(w, http.StatusInternalServerError, "INTERNAL_ERROR", "Failed to create address") + return + } + if count >= limits.MaxAddresses { + writeError(w, http.StatusForbidden, "TIER_LIMIT_REACHED", + fmt.Sprintf("Your %s plan allows %d address(es). Upgrade to track more.", user.SubscriptionTier, limits.MaxAddresses)) + return + } + } + log.Printf("User %s creating address: %s", userID, body.Address) addr, err := h.addresses.Create(r.Context(), userID, body.Address, body.Label) diff --git a/backend/internal/handlers/alert_rule.go b/backend/internal/handlers/alert_rule.go index 0a88b0e..7f758c7 100644 --- a/backend/internal/handlers/alert_rule.go +++ b/backend/internal/handlers/alert_rule.go @@ -19,10 +19,11 @@ var errThresholdFormat = errors.New("unsupported threshold format") type AlertRuleHandler struct { alertRules *models.AlertRuleModel addresses *models.AddressModel + users *models.UserModel } -func NewAlertRuleHandler(alertRules *models.AlertRuleModel, addresses *models.AddressModel) *AlertRuleHandler { - return &AlertRuleHandler{alertRules: alertRules, addresses: addresses} +func NewAlertRuleHandler(alertRules *models.AlertRuleModel, addresses *models.AddressModel, users *models.UserModel) *AlertRuleHandler { + return &AlertRuleHandler{alertRules: alertRules, addresses: addresses, users: users} } func (h *AlertRuleHandler) Create(w http.ResponseWriter, r *http.Request) { @@ -132,6 +133,28 @@ func (h *AlertRuleHandler) Create(w http.ResponseWriter, r *http.Request) { return } + user, userErr := h.users.GetByID(r.Context(), userID) + if userErr != nil || user == nil { + log.Printf("Failed to get user %s for tier check: %v", userID, userErr) + writeError(w, http.StatusInternalServerError, "INTERNAL_ERROR", "Failed to verify account") + return + } + + limits := domain.GetTierLimits(user.SubscriptionTier) + if !limits.IsUnlimitedAlertTypes() { + typeCount, countErr := h.alertRules.CountDistinctTypesByAddress(r.Context(), addressID) + if countErr != nil { + log.Printf("Failed to count alert types for address %d: %v", addressID, countErr) + writeError(w, http.StatusInternalServerError, "INTERNAL_ERROR", "Failed to create alert rule") + return + } + if typeCount >= limits.MaxAlertTypes { + writeError(w, http.StatusForbidden, "TIER_LIMIT_REACHED", + fmt.Sprintf("Your %s plan allows %d alert type(s) per address. Upgrade for more.", user.SubscriptionTier, limits.MaxAlertTypes)) + return + } + } + newAlert, err := h.alertRules.Create(r.Context(), addressID, alertType, threshold, minimum, maximum) if err != nil { log.Printf("Error creating alert rule: %v", err) diff --git a/backend/internal/handlers/notification_config.go b/backend/internal/handlers/notification_config.go index d48b9dd..04203ab 100644 --- a/backend/internal/handlers/notification_config.go +++ b/backend/internal/handlers/notification_config.go @@ -18,11 +18,12 @@ var emailRe = regexp.MustCompile(`^[^\s@]+@[^\s@]+\.[^\s@]+$`) type NotificationConfigHandler struct { configs *models.NotificationConfigModel + users *models.UserModel cfg *config.Config } -func NewNotificationConfigHandler(configs *models.NotificationConfigModel, cfg *config.Config) *NotificationConfigHandler { - return &NotificationConfigHandler{configs: configs, cfg: cfg} +func NewNotificationConfigHandler(configs *models.NotificationConfigModel, users *models.UserModel, cfg *config.Config) *NotificationConfigHandler { + return &NotificationConfigHandler{configs: configs, users: users, cfg: cfg} } func (h *NotificationConfigHandler) GetConfig(w http.ResponseWriter, r *http.Request) { @@ -95,6 +96,26 @@ func (h *NotificationConfigHandler) UpdateConfig(w http.ResponseWriter, r *http. return } + user, userErr := h.users.GetByID(r.Context(), userID) + if userErr != nil || user == nil { + log.Printf("Failed to get user %s for tier check: %v", userID, userErr) + writeError(w, http.StatusInternalServerError, "INTERNAL_ERROR", "Failed to verify account") + return + } + + limits := domain.GetTierLimits(user.SubscriptionTier) + + if !limits.ChannelAllowed("discord") { + body.DiscordWebhookURL = nil + } + if !limits.ChannelAllowed("telegram") { + body.TelegramBotToken = nil + body.TelegramChatID = nil + } + if !limits.ChannelAllowed("slack") { + body.SlackWebhookURL = nil + } + enabled := true if body.NotificationEnabled != nil { enabled = *body.NotificationEnabled diff --git a/backend/internal/handlers/stripe.go b/backend/internal/handlers/stripe.go index de48e8f..e8c7aa8 100644 --- a/backend/internal/handlers/stripe.go +++ b/backend/internal/handlers/stripe.go @@ -2,6 +2,7 @@ package handlers import ( "encoding/json" + "fmt" "io" "log" "net/http" @@ -12,6 +13,7 @@ import ( "github.com/stripe/stripe-go/v82/webhook" "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" ) @@ -28,10 +30,45 @@ func NewStripeHandler(users *models.UserModel, cfg *config.Config) *StripeHandle return &StripeHandler{users: users, cfg: cfg} } -// CreateCheckoutSession creates a Stripe Checkout session for the monthly subscription. +func (h *StripeHandler) priceIDForTier(tier domain.SubscriptionTier) (string, error) { + switch tier { + case domain.TierPremium: + return h.cfg.StripePriceIDPremium, nil + case domain.TierPro: + return h.cfg.StripePriceIDPro, nil + default: + return "", fmt.Errorf("no Stripe price for tier %q", tier) //nolint:err113 + } +} + +// CreateCheckoutSession creates a Stripe Checkout session for the selected tier. func (h *StripeHandler) CreateCheckoutSession(w http.ResponseWriter, r *http.Request) { userID := middleware.GetUserID(r.Context()) + var body struct { + Tier string `json:"tier"` + } + if err := json.NewDecoder(r.Body).Decode(&body); err != nil { + writeError(w, http.StatusBadRequest, "BAD_REQUEST", "Invalid request body") + return + } + + if body.Tier == "" { + body.Tier = "premium" + } + + tier := domain.SubscriptionTier(body.Tier) + if tier != domain.TierPremium && tier != domain.TierPro { + writeError(w, http.StatusBadRequest, "VALIDATION_ERROR", "Tier must be 'premium' or 'pro'") + return + } + + priceID, err := h.priceIDForTier(tier) + if err != nil { + writeError(w, http.StatusBadRequest, "VALIDATION_ERROR", err.Error()) + return + } + user, err := h.users.GetByID(r.Context(), userID) if err != nil || user == nil { log.Printf("Failed to get user %s: %v", userID, err) @@ -43,7 +80,7 @@ func (h *StripeHandler) CreateCheckoutSession(w http.ResponseWriter, r *http.Req Mode: stripe.String(string(stripe.CheckoutSessionModeSubscription)), LineItems: []*stripe.CheckoutSessionLineItemParams{ { - Price: stripe.String(h.cfg.StripePriceID), + Price: stripe.String(priceID), Quantity: stripe.Int64(1), }, }, @@ -53,6 +90,8 @@ func (h *StripeHandler) CreateCheckoutSession(w http.ResponseWriter, r *http.Req CustomerEmail: stripe.String(user.Email), } + params.AddMetadata("tier", string(tier)) + if user.StripeCustomerID != nil && *user.StripeCustomerID != "" { params.Customer = user.StripeCustomerID params.CustomerEmail = nil @@ -80,14 +119,14 @@ func (h *StripeHandler) GetSubscriptionStatus(w http.ResponseWriter, r *http.Req } writeJSON(w, http.StatusOK, map[string]any{ - "subscription_status": user.SubscriptionStatus, + "subscription_status": user.SubscriptionStatus, + "subscription_tier": user.SubscriptionTier, "subscription_created_at": user.SubscriptionCreatedAt, }) } // VerifyCheckoutSession retrieves a completed checkout session from Stripe, // confirms payment, and activates the user's subscription in the database. -// This is the primary activation path; webhooks serve as a backup. func (h *StripeHandler) VerifyCheckoutSession(w http.ResponseWriter, r *http.Request) { userID := middleware.GetUserID(r.Context()) @@ -116,6 +155,11 @@ func (h *StripeHandler) VerifyCheckoutSession(w http.ResponseWriter, r *http.Req return } + tier := domain.TierPremium + if t, ok := s.Metadata["tier"]; ok && domain.IsValidTier(t) { + tier = domain.SubscriptionTier(t) + } + customerID := "" if s.Customer != nil { customerID = s.Customer.ID @@ -131,13 +175,33 @@ func (h *StripeHandler) VerifyCheckoutSession(w http.ResponseWriter, r *http.Req } } if subscriptionID != "" && customerID != "" { - if err := h.users.ActivateSubscription(r.Context(), customerID, subscriptionID, "active"); err != nil { + if err := h.users.ActivateSubscription(r.Context(), customerID, subscriptionID, "active", tier); err != nil { log.Printf("VerifyCheckout: failed to activate subscription: %v", err) } } - log.Printf("Checkout verified for user %s, customer %s, subscription %s", userID, customerID, subscriptionID) - writeJSON(w, http.StatusOK, map[string]string{"subscription_status": "active"}) + log.Printf("Checkout verified for user %s, customer %s, subscription %s, tier %s", userID, customerID, subscriptionID, tier) + writeJSON(w, http.StatusOK, map[string]string{ + "subscription_status": "active", + "subscription_tier": string(tier), + }) +} + +// ActivateFreeTier sets the user to the free tier without Stripe involvement. +func (h *StripeHandler) ActivateFreeTier(w http.ResponseWriter, r *http.Request) { + userID := middleware.GetUserID(r.Context()) + + if err := h.users.ActivateFreeTier(r.Context(), userID); err != nil { + log.Printf("ActivateFreeTier: failed for user %s: %v", userID, err) + writeError(w, http.StatusInternalServerError, "INTERNAL_ERROR", "Failed to activate free tier") + return + } + + log.Printf("Free tier activated for user %s", userID) + writeJSON(w, http.StatusOK, map[string]string{ + "subscription_status": "active", + "subscription_tier": "free", + }) } // CreatePortalSession creates a Stripe Billing Portal session so the user can @@ -173,7 +237,6 @@ func (h *StripeHandler) CreatePortalSession(w http.ResponseWriter, r *http.Reque } // HandleWebhook processes incoming Stripe webhook events. -// This endpoint must NOT require authentication (Stripe calls it directly). func (h *StripeHandler) HandleWebhook(w http.ResponseWriter, r *http.Request) { payload, err := io.ReadAll(io.LimitReader(r.Body, webhookMaxBodyBytes)) if err != nil { @@ -217,6 +280,11 @@ func (h *StripeHandler) handleCheckoutCompleted(r *http.Request, event stripe.Ev return } + tier := domain.TierPremium + if t, ok := session.Metadata["tier"]; ok && domain.IsValidTier(t) { + tier = domain.SubscriptionTier(t) + } + customerID := "" if session.Customer != nil { customerID = session.Customer.ID @@ -233,12 +301,12 @@ func (h *StripeHandler) handleCheckoutCompleted(r *http.Request, event stripe.Ev } if subscriptionID != "" && customerID != "" { - if err := h.users.ActivateSubscription(r.Context(), customerID, subscriptionID, "active"); err != nil { + if err := h.users.ActivateSubscription(r.Context(), customerID, subscriptionID, "active", tier); err != nil { log.Printf("Failed to activate subscription: %v", err) } } - log.Printf("Checkout completed for user %s, customer %s, subscription %s", userID, customerID, subscriptionID) + log.Printf("Checkout completed for user %s, customer %s, subscription %s, tier %s", userID, customerID, subscriptionID, tier) } func (h *StripeHandler) handleSubscriptionUpdated(r *http.Request, event stripe.Event) { @@ -257,7 +325,7 @@ func (h *StripeHandler) handleSubscriptionUpdated(r *http.Request, event stripe. } status := string(sub.Status) - if err := h.users.ActivateSubscription(r.Context(), customerID, sub.ID, status); err != nil { + if err := h.users.UpdateSubscriptionStatus(r.Context(), customerID, status); err != nil { log.Printf("Failed to update subscription status: %v", err) } diff --git a/backend/internal/middleware/auth.go b/backend/internal/middleware/auth.go index b9f3306..ba3585a 100644 --- a/backend/internal/middleware/auth.go +++ b/backend/internal/middleware/auth.go @@ -17,6 +17,7 @@ type contextKey string const ( UserIDKey contextKey = "user_id" UserEmailKey contextKey = "user_email" + UserTierKey contextKey = "user_tier" ) type errorResponse struct { @@ -94,6 +95,7 @@ func Authenticate(userModel *models.UserModel) func(http.Handler) http.Handler { ctx := context.WithValue(r.Context(), UserIDKey, user.ID) ctx = context.WithValue(ctx, UserEmailKey, email) + ctx = context.WithValue(ctx, UserTierKey, string(user.SubscriptionTier)) next.ServeHTTP(w, r.WithContext(ctx)) }) @@ -150,3 +152,10 @@ func GetUserEmail(ctx context.Context) string { } return "" } + +func GetUserTier(ctx context.Context) string { + if v, ok := ctx.Value(UserTierKey).(string); ok { + return v + } + return "free" +} diff --git a/backend/internal/models/address.go b/backend/internal/models/address.go index 75ffdab..fbd6d45 100644 --- a/backend/internal/models/address.go +++ b/backend/internal/models/address.go @@ -125,6 +125,15 @@ func (m *AddressModel) UpdateLabel(ctx context.Context, id int, userID string, l return &a, nil } +func (m *AddressModel) CountByUser(ctx context.Context, userID string) (int, error) { + var count int + err := m.pool.QueryRow(ctx, + `SELECT COUNT(*) FROM addresses WHERE user_id = $1`, + userID, + ).Scan(&count) + return count, err +} + func (m *AddressModel) Remove(ctx context.Context, id int, userID string) (bool, error) { tag, err := m.pool.Exec(ctx, `DELETE FROM addresses WHERE id = $1 AND user_id = $2`, diff --git a/backend/internal/models/alert_rule.go b/backend/internal/models/alert_rule.go index 4f0a834..e4d4278 100644 --- a/backend/internal/models/alert_rule.go +++ b/backend/internal/models/alert_rule.go @@ -121,6 +121,15 @@ func (m *AlertRuleModel) UpdateThresholds(ctx context.Context, id int, minimum, return &r, nil } +func (m *AlertRuleModel) CountDistinctTypesByAddress(ctx context.Context, addressID int) (int, error) { + var count int + err := m.pool.QueryRow(ctx, + `SELECT COUNT(DISTINCT type) FROM alert_rules WHERE address_id = $1`, + addressID, + ).Scan(&count) + return count, err +} + func (m *AlertRuleModel) Remove(ctx context.Context, id int) (bool, error) { tag, err := m.pool.Exec(ctx, `DELETE FROM alert_rules WHERE id = $1`, diff --git a/backend/internal/models/user.go b/backend/internal/models/user.go index a271bd0..1481c68 100644 --- a/backend/internal/models/user.go +++ b/backend/internal/models/user.go @@ -19,14 +19,14 @@ func NewUserModel(pool *pgxpool.Pool) *UserModel { const userColumns = `id, firebase_uid, email, display_name, stripe_customer_id, stripe_subscription_id, subscription_status, - subscription_created_at, created_at, updated_at` + subscription_tier, subscription_created_at, created_at, updated_at` func scanUser(row pgx.Row) (*domain.User, error) { var u domain.User err := row.Scan( &u.ID, &u.FirebaseUID, &u.Email, &u.DisplayName, &u.StripeCustomerID, &u.StripeSubscriptionID, &u.SubscriptionStatus, - &u.SubscriptionCreatedAt, &u.CreatedAt, &u.UpdatedAt, + &u.SubscriptionTier, &u.SubscriptionCreatedAt, &u.CreatedAt, &u.UpdatedAt, ) if err != nil { if errors.Is(err, pgx.ErrNoRows) { @@ -66,15 +66,37 @@ func (m *UserModel) UpdateStripeCustomer(ctx context.Context, userID, stripeCust return err } -func (m *UserModel) ActivateSubscription(ctx context.Context, stripeCustomerID, subscriptionID, status string) error { +func (m *UserModel) ActivateSubscription(ctx context.Context, stripeCustomerID, subscriptionID, status string, tier domain.SubscriptionTier) error { _, err := m.pool.Exec(ctx, `UPDATE users SET stripe_subscription_id = $2, subscription_status = $3, + subscription_tier = $4, subscription_created_at = COALESCE(subscription_created_at, NOW()), updated_at = NOW() WHERE stripe_customer_id = $1`, - stripeCustomerID, subscriptionID, status, + stripeCustomerID, subscriptionID, status, string(tier), + ) + return err +} + +func (m *UserModel) UpdateSubscriptionTier(ctx context.Context, userID string, tier domain.SubscriptionTier) error { + _, err := m.pool.Exec(ctx, + `UPDATE users SET subscription_tier = $2, updated_at = NOW() WHERE id = $1`, + userID, string(tier), + ) + return err +} + +func (m *UserModel) ActivateFreeTier(ctx context.Context, userID string) error { + _, err := m.pool.Exec(ctx, + `UPDATE users + SET subscription_status = 'active', + subscription_tier = 'free', + subscription_created_at = COALESCE(subscription_created_at, NOW()), + updated_at = NOW() + WHERE id = $1`, + userID, ) return err } diff --git a/frontend/src/api/stripe.jsx b/frontend/src/api/stripe.jsx index 701ad2e..d4c1503 100644 --- a/frontend/src/api/stripe.jsx +++ b/frontend/src/api/stripe.jsx @@ -1,15 +1,29 @@ import { getAuthHeaders } from "./authHeaders"; import { API_BASE } from "./config"; -export async function createCheckoutSession() { +export async function createCheckoutSession(tier = "premium") { const headers = await getAuthHeaders(); const res = await fetch(`${API_BASE}/stripe/create-checkout-session`, { + method: "POST", + headers: { ...headers, "Content-Type": "application/json" }, + body: JSON.stringify({ tier }), + }); + if (!res.ok) { + const data = await res.json(); + throw new Error(data.message || "Failed to create checkout session"); + } + return res.json(); +} + +export async function activateFreeTier() { + const headers = await getAuthHeaders(); + const res = await fetch(`${API_BASE}/stripe/activate-free`, { method: "POST", headers, }); if (!res.ok) { const data = await res.json(); - throw new Error(data.message || "Failed to create checkout session"); + throw new Error(data.message || "Failed to activate free tier"); } return res.json(); } diff --git a/frontend/src/components/TierPicker.css b/frontend/src/components/TierPicker.css new file mode 100644 index 0000000..c783237 --- /dev/null +++ b/frontend/src/components/TierPicker.css @@ -0,0 +1,120 @@ +.tier-picker { + display: grid; + grid-template-columns: repeat(3, 1fr); + gap: 1rem; +} + +.tier-picker__card { + position: relative; + display: flex; + flex-direction: column; + padding: 1.5rem 1.25rem; + background-color: var(--color-bg-elevated); + border: 1px solid var(--color-border); + border-radius: var(--radius-xl); + text-align: center; +} + +.tier-picker__card--highlighted { + border-color: var(--color-primary); + box-shadow: 0 0 0 1px var(--color-primary); +} + +.tier-picker__card--selected { + border-color: var(--color-success); + box-shadow: 0 0 0 1px var(--color-success); +} + +.tier-picker__badge { + position: absolute; + top: -10px; + left: 50%; + transform: translateX(-50%); + background-color: var(--color-primary); + color: white; + font-size: 0.7rem; + font-weight: 700; + padding: 2px 12px; + border-radius: 10px; + white-space: nowrap; +} + +.tier-picker__name { + margin-bottom: 0.5rem; + font-size: 1.05rem; + color: var(--color-text); +} + +.tier-picker__price { + margin-bottom: 1.25rem; +} + +.tier-picker__amount { + font-size: 2rem; + font-weight: 700; + color: white; +} + +.tier-picker__period { + font-size: 0.85rem; + color: var(--color-text-dimmed); +} + +.tier-picker__features { + list-style: none; + padding: 0; + margin: 0 0 1.25rem; + text-align: left; + flex-grow: 1; +} + +.tier-picker__feature { + padding: 0.3rem 0; + font-size: 0.82rem; + color: var(--color-text-label); +} + +.tier-picker__feature--disabled { + color: var(--color-text-dimmed); + opacity: 0.55; +} + +.tier-picker__check { + color: var(--color-primary); + font-weight: bold; + margin-right: 0.4rem; +} + +.tier-picker__dash { + margin-right: 0.4rem; +} + +.tier-picker__btn { + width: 100%; + padding: 0.6rem; + background-color: var(--color-bg-card); + border: 1px solid var(--color-border); + color: var(--color-text); + cursor: pointer; + border-radius: var(--radius-md); + font-weight: 600; +} + +.tier-picker__btn:hover { + background-color: var(--color-primary); + border-color: var(--color-primary); +} + +.tier-picker__btn--selected { + background-color: var(--color-success); + border-color: var(--color-success); + color: white; +} + +@media (max-width: 768px) { + .tier-picker { + grid-template-columns: 1fr; + max-width: 400px; + margin: 0 auto; + } +} diff --git a/frontend/src/components/TierPicker.jsx b/frontend/src/components/TierPicker.jsx new file mode 100644 index 0000000..3e3d556 --- /dev/null +++ b/frontend/src/components/TierPicker.jsx @@ -0,0 +1,94 @@ +import "./TierPicker.css"; + +const TIERS = [ + { + id: "free", + name: "Free Trial Access", + price: "$0", + period: "", + features: [ + "Monitor 1 blockchain address", + "Daily email digest alert", + "1 transaction alert type per monitored address", + ], + disabledFeatures: [ + "Real-time Discord alerts", + "Real-time Telegram alerts", + "Real-time Slack alerts", + ], + }, + { + id: "premium", + name: "Premium Access", + price: "$1.99", + period: "/month", + features: [ + "Monitor 3 blockchain addresses", + "Daily email digest alert", + "Real-time Discord alerts", + "Real-time Telegram alerts", + "2 transaction alert types per monitored address", + ], + disabledFeatures: [ + "Real-time Slack alerts", + ], + highlighted: true, + }, + { + id: "pro", + name: "Pro Access", + price: "$11.99", + period: "/month", + features: [ + "Monitor unlimited blockchain addresses", + "Daily email digest alert", + "Real-time Discord alerts", + "Real-time Telegram alerts", + "Real-time Slack alerts", + "Unlimited transaction alert types per monitored address", + ], + disabledFeatures: [], + }, +]; + +export default function TierPicker({ onSelect, selectedTier }) { + return ( +
+ {TIERS.map((tier) => ( +
+ {tier.highlighted && ( +
Most Popular
+ )} +

{tier.name}

+
+ {tier.price} + {tier.period && ( + {tier.period} + )} +
+
    + {tier.features.map((f) => ( +
  • + {f} +
  • + ))} + {tier.disabledFeatures.map((f) => ( +
  • + {f} +
  • + ))} +
+ +
+ ))} +
+ ); +} diff --git a/frontend/src/components/UpgradeBanner.css b/frontend/src/components/UpgradeBanner.css new file mode 100644 index 0000000..3ae45f1 --- /dev/null +++ b/frontend/src/components/UpgradeBanner.css @@ -0,0 +1,32 @@ +.upgrade-banner { + display: flex; + align-items: center; + justify-content: space-between; + gap: 1rem; + padding: 0.75rem 1rem; + margin-bottom: 1rem; + border-radius: var(--radius-md); + background-color: #3a2e00; + border: 1px solid #aa7700; + color: #ffcc44; +} + +.upgrade-banner__message { + font-size: 0.88rem; +} + +.upgrade-banner__link { + background: none; + border: 1px solid #aa7700; + color: #ffcc44; + padding: 0.25rem 0.75rem; + border-radius: var(--radius-md); + cursor: pointer; + font-size: 0.82rem; + font-weight: 600; + white-space: nowrap; +} + +.upgrade-banner__link:hover { + background-color: #4a3800; +} diff --git a/frontend/src/components/UpgradeBanner.jsx b/frontend/src/components/UpgradeBanner.jsx new file mode 100644 index 0000000..3b66f28 --- /dev/null +++ b/frontend/src/components/UpgradeBanner.jsx @@ -0,0 +1,18 @@ +import { useNavigate } from "react-router-dom"; +import "./UpgradeBanner.css"; + +export default function UpgradeBanner({ message, linkTo = "/account" }) { + const navigate = useNavigate(); + + return ( +
+ {message} + +
+ ); +} diff --git a/frontend/src/contexts/AuthContext.jsx b/frontend/src/contexts/AuthContext.jsx index 586d654..03e1082 100644 --- a/frontend/src/contexts/AuthContext.jsx +++ b/frontend/src/contexts/AuthContext.jsx @@ -1,10 +1,10 @@ /** * AuthContext - Firebase Authentication State Management * - * Provides authentication state and methods throughout the app + * Provides authentication state, tier info, and methods throughout the app */ -import { createContext, useContext, useEffect, useState } from "react"; +import { createContext, useContext, useEffect, useState, useCallback } from "react"; import { createUserWithEmailAndPassword, signInWithEmailAndPassword, @@ -12,9 +12,16 @@ import { onAuthStateChanged, } from "firebase/auth"; import { auth } from "../firebase/config"; +import { getAccount } from "../api/account"; const AuthContext = createContext(); +const DEFAULT_TIER_LIMITS = { + max_addresses: 1, + max_alert_types: 1, + allowed_channels: ["email"], +}; + /** * Hook to access auth context * @returns {Object} Auth context value @@ -28,16 +35,28 @@ export function useAuth() { } /** - * AuthProvider - Wraps app and provides auth state + * AuthProvider - Wraps app and provides auth state + tier info */ export function AuthProvider({ children }) { const [currentUser, setCurrentUser] = useState(null); const [loading, setLoading] = useState(true); const [error, setError] = useState(null); - /** - * Sign up with email and password - */ + const [userTier, setUserTier] = useState("free"); + const [tierLimits, setTierLimits] = useState(DEFAULT_TIER_LIMITS); + const [addressCount, setAddressCount] = useState(0); + + const refreshAccount = useCallback(async () => { + try { + const data = await getAccount(); + setUserTier(data.subscription_tier || "free"); + setTierLimits(data.tier_limits || DEFAULT_TIER_LIMITS); + setAddressCount(data.address_count || 0); + } catch { + // account fetch can fail during onboarding before subscription is active + } + }, []); + async function signup(email, password) { try { setError(null); @@ -53,9 +72,6 @@ export function AuthProvider({ children }) { } } - /** - * Log in with email and password - */ async function login(email, password) { try { setError(null); @@ -71,12 +87,12 @@ export function AuthProvider({ children }) { } } - /** - * Log out current user - */ async function logout() { try { setError(null); + setUserTier("free"); + setTierLimits(DEFAULT_TIER_LIMITS); + setAddressCount(0); await signOut(auth); } catch (err) { setError(err.message); @@ -84,18 +100,16 @@ export function AuthProvider({ children }) { } } - /** - * Listen for auth state changes - */ useEffect(() => { const unsubscribe = onAuthStateChanged(auth, (user) => { setCurrentUser(user); setLoading(false); + if (user) { + refreshAccount(); + } }); - - // Cleanup subscription return unsubscribe; - }, []); + }, [refreshAccount]); const value = { currentUser, @@ -104,6 +118,10 @@ export function AuthProvider({ children }) { logout, error, loading, + userTier, + tierLimits, + addressCount, + refreshAccount, }; return ( diff --git a/frontend/src/pages/addresses/Addresses.jsx b/frontend/src/pages/addresses/Addresses.jsx index b0255ac..a23144c 100644 --- a/frontend/src/pages/addresses/Addresses.jsx +++ b/frontend/src/pages/addresses/Addresses.jsx @@ -1,15 +1,22 @@ import { useState, useEffect } from "react"; +import { useAuth } from "../../contexts/AuthContext"; import AddressForm from "../../components/AddressForm"; +import UpgradeBanner from "../../components/UpgradeBanner"; import { getAddresses, createAddress, deleteAddress, updateAddress } from "../../api/addresses"; import "./Addresses.css"; export default function Addresses() { + const { tierLimits, refreshAccount } = useAuth(); const [addresses, setAddresses] = useState([]); const [loading, setLoading] = useState(true); const [error, setError] = useState(null); const [editingId, setEditingId] = useState(null); const [editLabel, setEditLabel] = useState(""); + const maxAddresses = tierLimits.max_addresses; + const isUnlimited = maxAddresses === -1; + const atLimit = !isUnlimited && addresses.length >= maxAddresses; + useEffect(() => { async function fetchAddresses() { try { @@ -32,6 +39,7 @@ export default function Addresses() { const newAddress = await createAddress(data); setAddresses((prev) => [...prev, newAddress]); setError(null); + refreshAccount(); } catch (err) { setError(err.message); console.error("Failed to create address:", err); @@ -47,6 +55,7 @@ export default function Addresses() { await deleteAddress(id); setAddresses((prev) => prev.filter((a) => a.id !== id)); setError(null); + refreshAccount(); } catch (err) { setError(err.message); console.error("Failed to delete address:", err); @@ -80,9 +89,17 @@ export default function Addresses() {

Add Addresses to Track

-
- -
+ {atLimit && ( + + )} + + {!atLimit && ( +
+ +
+ )}

Existing Tracked Addresses

diff --git a/frontend/src/pages/alerts/Alerts.css b/frontend/src/pages/alerts/Alerts.css index c89c7c7..bcae534 100644 --- a/frontend/src/pages/alerts/Alerts.css +++ b/frontend/src/pages/alerts/Alerts.css @@ -125,6 +125,20 @@ margin-top: 1rem; } +/* Locked / disabled notification channel */ +.alerts__channel-locked { + opacity: 0.55; + position: relative; +} + +.alerts__channel-locked .upgrade-banner { + opacity: 1; +} + +.alerts__channel-locked input { + pointer-events: none; +} + /* ── Alerts Page Responsive ──────────────────────────────── */ /* Laptop: tighten the gap */ diff --git a/frontend/src/pages/alerts/Alerts.jsx b/frontend/src/pages/alerts/Alerts.jsx index 34daa8b..b5315e1 100644 --- a/frontend/src/pages/alerts/Alerts.jsx +++ b/frontend/src/pages/alerts/Alerts.jsx @@ -1,7 +1,9 @@ import { useState, useEffect } from "react"; +import { useAuth } from "../../contexts/AuthContext"; import AlertForm from "../../components/AlertForm"; import Button from "../../components/Button"; import Input from "../../components/Input"; +import UpgradeBanner from "../../components/UpgradeBanner"; import { getAddresses } from "../../api/addresses"; import { getAlerts, @@ -20,6 +22,8 @@ import { import "./Alerts.css"; export default function Alerts() { + const { tierLimits, userTier } = useAuth(); + const [addresses, setAddresses] = useState([]); const [selectedAddressId, setSelectedAddressId] = useState(null); const [alerts, setAlerts] = useState([]); @@ -43,6 +47,21 @@ export default function Alerts() { const [openAccordions, setOpenAccordions] = useState({}); const [thresholdEdits, setThresholdEdits] = useState({}); + const allowedChannels = tierLimits.allowed_channels || ["email"]; + const canTelegram = allowedChannels.includes("telegram"); + const canDiscord = allowedChannels.includes("discord"); + const canSlack = allowedChannels.includes("slack"); + + const maxAlertTypes = tierLimits.max_alert_types; + const isUnlimitedAlerts = maxAlertTypes === -1; + + function distinctAlertTypeCount() { + const types = new Set(alerts.map((a) => a.type)); + return types.size; + } + + const atAlertLimit = !isUnlimitedAlerts && distinctAlertTypeCount() >= maxAlertTypes; + function toggleAccordion(alertId) { setOpenAccordions((prev) => ({ ...prev, [alertId]: !prev[alertId] })); } @@ -209,11 +228,11 @@ export default function Alerts() { const config = { notification_enabled: notificationEnabled, - discord_webhook_url: discordWebhookUrl || null, - telegram_bot_token: telegramBotToken || null, - telegram_chat_id: telegramChatId || null, + discord_webhook_url: canDiscord ? (discordWebhookUrl || null) : null, + telegram_bot_token: canTelegram ? (telegramBotToken || null) : null, + telegram_chat_id: canTelegram ? (telegramChatId || null) : null, email: email || null, - slack_webhook_url: slackWebhookUrl || null, + slack_webhook_url: canSlack ? (slackWebhookUrl || null) : null, }; await updateNotificationConfig(config); @@ -357,10 +376,18 @@ export default function Alerts() {
-
-

Create New Alert

- -
+ {atAlertLimit && userTier !== "pro" && ( + + )} + + {!atAlertLimit && ( +
+

Create New Alert

+ +
+ )}

Active Alert Rules

@@ -540,19 +567,24 @@ export default function Alerts() { {notificationEnabled && ( <> {/* Telegram */} -
+

Telegram

+ {!canTelegram && ( + + )} {/* Discord */} -
+

Discord

+ {!canDiscord && ( + + )}
{/* Slack */} -
+

Slack

+ {!canSlack && ( + + )}
({ ...prev, [field]: value })); } + const tierLimits = TIER_LIMITS[data.selectedTier] || TIER_LIMITS.free; + useEffect(() => { if (!currentUser) return; getAddresses() @@ -67,7 +78,6 @@ export default function Subscribe() { .catch(() => { }); }, [currentUser, navigate]); - // Handle Stripe redirect back from checkout useEffect(() => { if (!currentUser) return; const payment = searchParams.get("payment"); @@ -77,23 +87,32 @@ export default function Subscribe() { setLoading(true); verifyCheckoutSession(sessionId) .then(() => { - setStep(2); + setStep(3); }) .catch((err) => { setError("Payment verification failed: " + err.message); - setStep(1); + setStep(2); }) .finally(() => setLoading(false)); } else if (payment === "cancelled") { setSearchParams({}, { replace: true }); - setStep(1); + setStep(2); setError("Payment was cancelled. Please try again."); } }, [currentUser, searchParams, setSearchParams]); // ── Step handlers ───────────────────────────────────────────────────────── - async function handleStep1() { + function handleStep1() { + setError(""); + if (!data.selectedTier) { + setError("Please select a plan to continue"); + return; + } + setStep(2); + } + + async function handleStep2() { setError(""); if (!data.email || !data.password || !data.confirmPassword) { setError("Please fill in all fields"); @@ -112,12 +131,19 @@ export default function Subscribe() { if (!currentUser) { await signup(data.email, data.password); } - const status = await getSubscriptionStatus(); - if (status.subscription_status === "active" || status.subscription_status === "trialing") { - setStep(2); + + if (data.selectedTier === "free") { + await activateFreeTier(); + setStep(3); return; } - const { url } = await createCheckoutSession(); + + const status = await getSubscriptionStatus(); + if (status.subscription_status === "active" || status.subscription_status === "trialing") { + setStep(3); + return; + } + const { url } = await createCheckoutSession(data.selectedTier); window.location.href = url; } catch (err) { if (err.code === "auth/email-already-in-use") { @@ -133,7 +159,7 @@ export default function Subscribe() { } } - async function handleStep2() { + async function handleStep3() { setError(""); if (!data.walletAddress) { setError("Please enter a wallet address"); @@ -150,7 +176,7 @@ export default function Subscribe() { label: data.walletLabel || undefined, }); set("createdAddressId", created.id); - setStep(3); + setStep(4); } catch (err) { setError(err.message); } finally { @@ -158,7 +184,7 @@ export default function Subscribe() { } } - async function handleStep3() { + async function handleStep4() { setError(""); const rules = []; if (data.alertIncomingTx) rules.push({ type: "incoming_tx" }); @@ -179,7 +205,7 @@ export default function Subscribe() { } if (rules.length === 0) { - setStep(4); + setStep(5); return; } @@ -191,7 +217,7 @@ export default function Subscribe() { created.push(result); } set("alertsCreated", created); - setStep(4); + setStep(5); } catch (err) { setError(err.message); } finally { @@ -199,12 +225,12 @@ export default function Subscribe() { } } - async function handleStep4() { + async function handleStep5() { setError(""); const hasAny = data.discordWebhookUrl || data.slackWebhookUrl || data.notificationEmail; if (!hasAny) { - setStep(5); + setStep(6); return; } try { @@ -216,7 +242,7 @@ export default function Subscribe() { email: data.notificationEmail || undefined, }); set("notificationConfigured", true); - setStep(5); + setStep(6); } catch (err) { setError(err.message); } finally { @@ -237,6 +263,21 @@ export default function Subscribe() { } } + // ── Alert type limit helpers ────────────────────────────────────────────── + + function countSelectedAlerts() { + let count = 0; + if (data.alertIncomingTx) count++; + if (data.alertOutgoingTx) count++; + if (data.alertLargeTransfer) count++; + if (data.alertBalanceBelow) count++; + return count; + } + + function canSelectMoreAlerts() { + return tierLimits.maxAlertTypes > countSelectedAlerts(); + } + // ── Progress bar ────────────────────────────────────────────────────────── function ProgressBar() { @@ -284,7 +325,22 @@ export default function Subscribe() { // ── Step content ────────────────────────────────────────────────────────── - function Step1() { + function StepChoosePlan() { + return ( + <> +

Choose your plan

+

+ Select the plan that works best for you. You can upgrade anytime. +

+ set("selectedTier", tier)} + selectedTier={data.selectedTier} + /> + + ); + } + + function StepCreateAccount() { return ( <>

Create your account

@@ -317,7 +373,7 @@ export default function Subscribe() { ); } - function Step2() { + function StepAddWallet() { return ( <>

Add a wallet address

@@ -343,28 +399,35 @@ export default function Subscribe() { ); } - function Step3() { + function StepAlertRules() { + const atLimit = !canSelectMoreAlerts(); + const maxTypes = tierLimits.maxAlertTypes; + return ( <>

Configure alert rules

- Choose which events trigger notifications. You can change these later. + Choose which events trigger notifications ({countSelectedAlerts()}/{maxTypes} selected). + You can change these later.

set("alertIncomingTx", v)} label="Incoming transaction" + disabled={!data.alertIncomingTx && atLimit} /> set("alertOutgoingTx", v)} label="Outgoing transaction" + disabled={!data.alertOutgoingTx && atLimit} /> set("alertLargeTransfer", v)} label="Large transfer" + disabled={!data.alertLargeTransfer && atLimit} > {data.alertLargeTransfer && (
@@ -384,6 +447,7 @@ export default function Subscribe() { checked={data.alertBalanceBelow} onChange={(v) => set("alertBalanceBelow", v)} label="Balance below" + disabled={!data.alertBalanceBelow && atLimit} > {data.alertBalanceBelow && (
@@ -399,11 +463,21 @@ export default function Subscribe() {
)} + + {atLimit && data.selectedTier !== "pro" && ( +

+ Your {data.selectedTier} plan allows {maxTypes} alert type{maxTypes !== 1 ? "s" : ""} per address. Upgrade for more. +

+ )} ); } - function Step4() { + function StepNotifications() { + const channels = tierLimits.channels; + const canDiscord = channels.includes("discord"); + const canSlack = channels.includes("slack"); + return ( <>

Set up notifications

@@ -411,7 +485,17 @@ export default function Subscribe() { Add at least one channel so you receive alerts. All fields are optional.

-
+ set("notificationEmail", v)} + disabled={loading} + placeholder="you@example.com" + className="form-field--last" + /> + + -
+ - - set("notificationEmail", v)} - disabled={loading} - placeholder="you@example.com" - className="form-field--last" - /> ); } - function Step5() { + function StepDone() { const alertCount = data.alertsCreated.length; const hasNotif = data.notificationConfigured; @@ -479,6 +559,12 @@ export default function Subscribe() {

Summary