add subscription tiers

This commit is contained in:
KS Jannette
2026-03-10 15:22:48 -04:00
parent 7b45c5b2ab
commit f9fa7def2b
26 changed files with 944 additions and 160 deletions

View File

@@ -55,14 +55,14 @@ func main() {
cfg.ResendAPIKey, cfg.EmailFrom, alertEventModel, notifConfigModel, cfg.ResendAPIKey, cfg.EmailFrom, alertEventModel, notifConfigModel,
) )
addressHandler := handlers.NewAddressHandler(addressModel) addressHandler := handlers.NewAddressHandler(addressModel, userModel)
alertRuleHandler := handlers.NewAlertRuleHandler(alertRuleModel, addressModel) alertRuleHandler := handlers.NewAlertRuleHandler(alertRuleModel, addressModel, userModel)
alertEventHandler := handlers.NewAlertEventHandler(alertEventModel) alertEventHandler := handlers.NewAlertEventHandler(alertEventModel)
notifConfigHandler := handlers.NewNotificationConfigHandler(notifConfigModel, cfg) notifConfigHandler := handlers.NewNotificationConfigHandler(notifConfigModel, userModel, cfg)
emailDigestHandler := handlers.NewEmailDigestHandler(emailDigestSvc, notifConfigModel) emailDigestHandler := handlers.NewEmailDigestHandler(emailDigestSvc, notifConfigModel)
statusHandler := handlers.NewStatusHandler(checkpointModel) statusHandler := handlers.NewStatusHandler(checkpointModel)
stripeHandler := handlers.NewStripeHandler(userModel, cfg) stripeHandler := handlers.NewStripeHandler(userModel, cfg)
accountHandler := handlers.NewAccountHandler(userModel, cfg) accountHandler := handlers.NewAccountHandler(userModel, addressModel, cfg)
authenticate := middleware.Authenticate(userModel) authenticate := middleware.Authenticate(userModel)
requireSub := middleware.RequireSubscription(userModel) requireSub := middleware.RequireSubscription(userModel)
@@ -91,6 +91,8 @@ func main() {
authenticate(http.HandlerFunc(stripeHandler.VerifyCheckoutSession))) authenticate(http.HandlerFunc(stripeHandler.VerifyCheckoutSession)))
mux.Handle("POST "+b+"/stripe/create-portal-session", mux.Handle("POST "+b+"/stripe/create-portal-session",
authenticate(http.HandlerFunc(stripeHandler.CreatePortalSession))) authenticate(http.HandlerFunc(stripeHandler.CreatePortalSession)))
mux.Handle("POST "+b+"/stripe/activate-free",
authenticate(http.HandlerFunc(stripeHandler.ActivateFreeTier)))
// Account route (auth required, NO subscription required) // Account route (auth required, NO subscription required)
mux.Handle("GET "+b+"/user/account", mux.Handle("GET "+b+"/user/account",

View File

@@ -0,0 +1 @@
ALTER TABLE users ADD COLUMN IF NOT EXISTS subscription_tier VARCHAR(20) DEFAULT 'free';

View File

@@ -14,6 +14,7 @@ CREATE TABLE users (
stripe_customer_id VARCHAR(255), stripe_customer_id VARCHAR(255),
stripe_subscription_id VARCHAR(255), stripe_subscription_id VARCHAR(255),
subscription_status VARCHAR(50) DEFAULT 'none', subscription_status VARCHAR(50) DEFAULT 'none',
subscription_tier VARCHAR(20) DEFAULT 'free',
subscription_created_at TIMESTAMP, subscription_created_at TIMESTAMP,
created_at TIMESTAMP DEFAULT NOW(), created_at TIMESTAMP DEFAULT NOW(),
updated_at TIMESTAMP DEFAULT NOW() updated_at TIMESTAMP DEFAULT NOW()

View File

@@ -32,11 +32,12 @@ type Config struct {
ResendAPIKey string ResendAPIKey string
EmailFrom string EmailFrom string
DigestIntervalHours int DigestIntervalHours int
StripeSecretKey string StripeSecretKey string
StripeWebhookSecret string StripeWebhookSecret string
StripePriceID string StripePriceIDPremium string
StripePublishableKey string StripePriceIDPro string
FrontendURL string StripePublishableKey string
FrontendURL string
} }
// Load reads configuration from environment variables and returns a Config. // Load reads configuration from environment variables and returns a Config.
@@ -57,10 +58,11 @@ func Load() (*Config, error) {
ResendAPIKey: os.Getenv("RESEND_API_KEY"), ResendAPIKey: os.Getenv("RESEND_API_KEY"),
EmailFrom: getEnv("EMAIL_FROM", "Koin Ping <alerts@koinping.com>"), EmailFrom: getEnv("EMAIL_FROM", "Koin Ping <alerts@koinping.com>"),
DigestIntervalHours: getEnvInt("DIGEST_INTERVAL_HOURS", defaultDigestIntervalHours), DigestIntervalHours: getEnvInt("DIGEST_INTERVAL_HOURS", defaultDigestIntervalHours),
StripeSecretKey: os.Getenv("STRIPE_SECRET_KEY"), StripeSecretKey: os.Getenv("STRIPE_SECRET_KEY"),
StripeWebhookSecret: os.Getenv("STRIPE_WEBHOOK_SECRET"), StripeWebhookSecret: os.Getenv("STRIPE_WEBHOOK_SECRET"),
StripePriceID: os.Getenv("STRIPE_PRICE_ID"), StripePriceIDPremium: os.Getenv("STRIPE_PRICE_ID_PREMIUM"),
StripePublishableKey: os.Getenv("STRIPE_PUBLISHABLE_KEY"), StripePriceIDPro: os.Getenv("STRIPE_PRICE_ID_PRO"),
StripePublishableKey: os.Getenv("STRIPE_PUBLISHABLE_KEY"),
FrontendURL: getEnv("FRONTEND_URL", "http://localhost:3000"), FrontendURL: getEnv("FRONTEND_URL", "http://localhost:3000"),
} }

View File

@@ -2,17 +2,77 @@ package domain
import "time" 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 { type User struct {
ID string `json:"id"` ID string `json:"id"`
FirebaseUID string `json:"-"` FirebaseUID string `json:"-"`
Email string `json:"email"` Email string `json:"email"`
DisplayName *string `json:"display_name"` //nolint:tagliatelle DisplayName *string `json:"display_name"` //nolint:tagliatelle
StripeCustomerID *string `json:"-"` StripeCustomerID *string `json:"-"`
StripeSubscriptionID *string `json:"-"` StripeSubscriptionID *string `json:"-"`
SubscriptionStatus string `json:"subscription_status"` //nolint:tagliatelle SubscriptionStatus string `json:"subscription_status"` //nolint:tagliatelle
SubscriptionCreatedAt *time.Time `json:"subscription_created_at,omitempty"` //nolint:tagliatelle SubscriptionTier SubscriptionTier `json:"subscription_tier"` //nolint:tagliatelle
CreatedAt time.Time `json:"created_at"` //nolint:tagliatelle SubscriptionCreatedAt *time.Time `json:"subscription_created_at,omitempty"` //nolint:tagliatelle
UpdatedAt time.Time `json:"updated_at"` //nolint:tagliatelle CreatedAt time.Time `json:"created_at"` //nolint:tagliatelle
UpdatedAt time.Time `json:"updated_at"` //nolint:tagliatelle
} }
type Address struct { type Address struct {

View File

@@ -6,29 +6,40 @@ import (
"time" "time"
"github.com/kjannette/koin-ping/backend/internal/config" "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/middleware"
"github.com/kjannette/koin-ping/backend/internal/models" "github.com/kjannette/koin-ping/backend/internal/models"
) )
type AccountHandler struct { type AccountHandler struct {
users *models.UserModel users *models.UserModel
cfg *config.Config addresses *models.AddressModel
cfg *config.Config
} }
func NewAccountHandler(users *models.UserModel, cfg *config.Config) *AccountHandler { func NewAccountHandler(users *models.UserModel, addresses *models.AddressModel, cfg *config.Config) *AccountHandler {
return &AccountHandler{users: users, cfg: cfg} return &AccountHandler{users: users, addresses: addresses, cfg: cfg}
} }
type accountResponse struct { type accountResponse struct {
UserID string `json:"user_id"` UserID string `json:"user_id"`
Email string `json:"email"` Email string `json:"email"`
UserName string `json:"user_name"` UserName string `json:"user_name"`
SubscriptionStatus string `json:"subscription_status"` SubscriptionStatus string `json:"subscription_status"`
SubscriptionPlan string `json:"subscription_plan"` SubscriptionTier string `json:"subscription_tier"`
MemberSince *string `json:"member_since,omitempty"` SubscriptionPlan string `json:"subscription_plan"`
NextBillingDate *string `json:"next_billing_date,omitempty"` TierLimits domain.TierLimits `json:"tier_limits"`
CancelAtPeriodEnd bool `json:"cancel_at_period_end"` AddressCount int `json:"address_count"`
PeriodEndDate *string `json:"period_end_date,omitempty"` 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) { 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 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{ resp := accountResponse{
UserID: user.ID, UserID: user.ID,
Email: email, Email: email,
UserName: email, UserName: email,
SubscriptionStatus: user.SubscriptionStatus, SubscriptionStatus: user.SubscriptionStatus,
SubscriptionPlan: "monthly/$1.99", SubscriptionTier: string(user.SubscriptionTier),
SubscriptionPlan: planLabel,
TierLimits: domain.GetTierLimits(user.SubscriptionTier),
AddressCount: addrCount,
} }
if user.SubscriptionCreatedAt != nil { if user.SubscriptionCreatedAt != nil {
@@ -55,9 +80,6 @@ func (h *AccountHandler) GetAccount(w http.ResponseWriter, r *http.Request) {
resp.MemberSince = &t 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.NextBillingDate = nil
resp.CancelAtPeriodEnd = false resp.CancelAtPeriodEnd = false
resp.PeriodEndDate = nil resp.PeriodEndDate = nil

View File

@@ -3,6 +3,7 @@ package handlers
import ( import (
"encoding/json" "encoding/json"
"fmt"
"log" "log"
"net/http" "net/http"
"regexp" "regexp"
@@ -17,10 +18,11 @@ var ethAddressRe = regexp.MustCompile(`^0x[a-fA-F0-9]{40}$`)
type AddressHandler struct { type AddressHandler struct {
addresses *models.AddressModel addresses *models.AddressModel
users *models.UserModel
} }
func NewAddressHandler(addresses *models.AddressModel) *AddressHandler { func NewAddressHandler(addresses *models.AddressModel, users *models.UserModel) *AddressHandler {
return &AddressHandler{addresses: addresses} return &AddressHandler{addresses: addresses, users: users}
} }
func (h *AddressHandler) Create(w http.ResponseWriter, r *http.Request) { 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 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) log.Printf("User %s creating address: %s", userID, body.Address)
addr, err := h.addresses.Create(r.Context(), userID, body.Address, body.Label) addr, err := h.addresses.Create(r.Context(), userID, body.Address, body.Label)

View File

@@ -19,10 +19,11 @@ var errThresholdFormat = errors.New("unsupported threshold format")
type AlertRuleHandler struct { type AlertRuleHandler struct {
alertRules *models.AlertRuleModel alertRules *models.AlertRuleModel
addresses *models.AddressModel addresses *models.AddressModel
users *models.UserModel
} }
func NewAlertRuleHandler(alertRules *models.AlertRuleModel, addresses *models.AddressModel) *AlertRuleHandler { func NewAlertRuleHandler(alertRules *models.AlertRuleModel, addresses *models.AddressModel, users *models.UserModel) *AlertRuleHandler {
return &AlertRuleHandler{alertRules: alertRules, addresses: addresses} return &AlertRuleHandler{alertRules: alertRules, addresses: addresses, users: users}
} }
func (h *AlertRuleHandler) Create(w http.ResponseWriter, r *http.Request) { 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 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) newAlert, err := h.alertRules.Create(r.Context(), addressID, alertType, threshold, minimum, maximum)
if err != nil { if err != nil {
log.Printf("Error creating alert rule: %v", err) log.Printf("Error creating alert rule: %v", err)

View File

@@ -18,11 +18,12 @@ var emailRe = regexp.MustCompile(`^[^\s@]+@[^\s@]+\.[^\s@]+$`)
type NotificationConfigHandler struct { type NotificationConfigHandler struct {
configs *models.NotificationConfigModel configs *models.NotificationConfigModel
users *models.UserModel
cfg *config.Config cfg *config.Config
} }
func NewNotificationConfigHandler(configs *models.NotificationConfigModel, cfg *config.Config) *NotificationConfigHandler { func NewNotificationConfigHandler(configs *models.NotificationConfigModel, users *models.UserModel, cfg *config.Config) *NotificationConfigHandler {
return &NotificationConfigHandler{configs: configs, cfg: cfg} return &NotificationConfigHandler{configs: configs, users: users, cfg: cfg}
} }
func (h *NotificationConfigHandler) GetConfig(w http.ResponseWriter, r *http.Request) { func (h *NotificationConfigHandler) GetConfig(w http.ResponseWriter, r *http.Request) {
@@ -95,6 +96,26 @@ func (h *NotificationConfigHandler) UpdateConfig(w http.ResponseWriter, r *http.
return 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 enabled := true
if body.NotificationEnabled != nil { if body.NotificationEnabled != nil {
enabled = *body.NotificationEnabled enabled = *body.NotificationEnabled

View File

@@ -2,6 +2,7 @@ package handlers
import ( import (
"encoding/json" "encoding/json"
"fmt"
"io" "io"
"log" "log"
"net/http" "net/http"
@@ -12,6 +13,7 @@ import (
"github.com/stripe/stripe-go/v82/webhook" "github.com/stripe/stripe-go/v82/webhook"
"github.com/kjannette/koin-ping/backend/internal/config" "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/middleware"
"github.com/kjannette/koin-ping/backend/internal/models" "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} 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) { func (h *StripeHandler) CreateCheckoutSession(w http.ResponseWriter, r *http.Request) {
userID := middleware.GetUserID(r.Context()) 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) user, err := h.users.GetByID(r.Context(), userID)
if err != nil || user == nil { if err != nil || user == nil {
log.Printf("Failed to get user %s: %v", userID, err) 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)), Mode: stripe.String(string(stripe.CheckoutSessionModeSubscription)),
LineItems: []*stripe.CheckoutSessionLineItemParams{ LineItems: []*stripe.CheckoutSessionLineItemParams{
{ {
Price: stripe.String(h.cfg.StripePriceID), Price: stripe.String(priceID),
Quantity: stripe.Int64(1), Quantity: stripe.Int64(1),
}, },
}, },
@@ -53,6 +90,8 @@ func (h *StripeHandler) CreateCheckoutSession(w http.ResponseWriter, r *http.Req
CustomerEmail: stripe.String(user.Email), CustomerEmail: stripe.String(user.Email),
} }
params.AddMetadata("tier", string(tier))
if user.StripeCustomerID != nil && *user.StripeCustomerID != "" { if user.StripeCustomerID != nil && *user.StripeCustomerID != "" {
params.Customer = user.StripeCustomerID params.Customer = user.StripeCustomerID
params.CustomerEmail = nil params.CustomerEmail = nil
@@ -80,14 +119,14 @@ func (h *StripeHandler) GetSubscriptionStatus(w http.ResponseWriter, r *http.Req
} }
writeJSON(w, http.StatusOK, map[string]any{ writeJSON(w, http.StatusOK, map[string]any{
"subscription_status": user.SubscriptionStatus, "subscription_status": user.SubscriptionStatus,
"subscription_tier": user.SubscriptionTier,
"subscription_created_at": user.SubscriptionCreatedAt, "subscription_created_at": user.SubscriptionCreatedAt,
}) })
} }
// VerifyCheckoutSession retrieves a completed checkout session from Stripe, // VerifyCheckoutSession retrieves a completed checkout session from Stripe,
// confirms payment, and activates the user's subscription in the database. // 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) { func (h *StripeHandler) VerifyCheckoutSession(w http.ResponseWriter, r *http.Request) {
userID := middleware.GetUserID(r.Context()) userID := middleware.GetUserID(r.Context())
@@ -116,6 +155,11 @@ func (h *StripeHandler) VerifyCheckoutSession(w http.ResponseWriter, r *http.Req
return return
} }
tier := domain.TierPremium
if t, ok := s.Metadata["tier"]; ok && domain.IsValidTier(t) {
tier = domain.SubscriptionTier(t)
}
customerID := "" customerID := ""
if s.Customer != nil { if s.Customer != nil {
customerID = s.Customer.ID customerID = s.Customer.ID
@@ -131,13 +175,33 @@ func (h *StripeHandler) VerifyCheckoutSession(w http.ResponseWriter, r *http.Req
} }
} }
if subscriptionID != "" && customerID != "" { 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("VerifyCheckout: failed to activate subscription: %v", err)
} }
} }
log.Printf("Checkout verified for user %s, customer %s, subscription %s", userID, customerID, subscriptionID) 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"}) 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 // 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. // 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) { func (h *StripeHandler) HandleWebhook(w http.ResponseWriter, r *http.Request) {
payload, err := io.ReadAll(io.LimitReader(r.Body, webhookMaxBodyBytes)) payload, err := io.ReadAll(io.LimitReader(r.Body, webhookMaxBodyBytes))
if err != nil { if err != nil {
@@ -217,6 +280,11 @@ func (h *StripeHandler) handleCheckoutCompleted(r *http.Request, event stripe.Ev
return return
} }
tier := domain.TierPremium
if t, ok := session.Metadata["tier"]; ok && domain.IsValidTier(t) {
tier = domain.SubscriptionTier(t)
}
customerID := "" customerID := ""
if session.Customer != nil { if session.Customer != nil {
customerID = session.Customer.ID customerID = session.Customer.ID
@@ -233,12 +301,12 @@ func (h *StripeHandler) handleCheckoutCompleted(r *http.Request, event stripe.Ev
} }
if subscriptionID != "" && customerID != "" { 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("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) { 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) 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) log.Printf("Failed to update subscription status: %v", err)
} }

View File

@@ -17,6 +17,7 @@ type contextKey string
const ( const (
UserIDKey contextKey = "user_id" UserIDKey contextKey = "user_id"
UserEmailKey contextKey = "user_email" UserEmailKey contextKey = "user_email"
UserTierKey contextKey = "user_tier"
) )
type errorResponse struct { 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(r.Context(), UserIDKey, user.ID)
ctx = context.WithValue(ctx, UserEmailKey, email) ctx = context.WithValue(ctx, UserEmailKey, email)
ctx = context.WithValue(ctx, UserTierKey, string(user.SubscriptionTier))
next.ServeHTTP(w, r.WithContext(ctx)) next.ServeHTTP(w, r.WithContext(ctx))
}) })
@@ -150,3 +152,10 @@ func GetUserEmail(ctx context.Context) string {
} }
return "" return ""
} }
func GetUserTier(ctx context.Context) string {
if v, ok := ctx.Value(UserTierKey).(string); ok {
return v
}
return "free"
}

View File

@@ -125,6 +125,15 @@ func (m *AddressModel) UpdateLabel(ctx context.Context, id int, userID string, l
return &a, nil 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) { func (m *AddressModel) Remove(ctx context.Context, id int, userID string) (bool, error) {
tag, err := m.pool.Exec(ctx, tag, err := m.pool.Exec(ctx,
`DELETE FROM addresses WHERE id = $1 AND user_id = $2`, `DELETE FROM addresses WHERE id = $1 AND user_id = $2`,

View File

@@ -121,6 +121,15 @@ func (m *AlertRuleModel) UpdateThresholds(ctx context.Context, id int, minimum,
return &r, nil 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) { func (m *AlertRuleModel) Remove(ctx context.Context, id int) (bool, error) {
tag, err := m.pool.Exec(ctx, tag, err := m.pool.Exec(ctx,
`DELETE FROM alert_rules WHERE id = $1`, `DELETE FROM alert_rules WHERE id = $1`,

View File

@@ -19,14 +19,14 @@ func NewUserModel(pool *pgxpool.Pool) *UserModel {
const userColumns = `id, firebase_uid, email, display_name, const userColumns = `id, firebase_uid, email, display_name,
stripe_customer_id, stripe_subscription_id, subscription_status, 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) { func scanUser(row pgx.Row) (*domain.User, error) {
var u domain.User var u domain.User
err := row.Scan( err := row.Scan(
&u.ID, &u.FirebaseUID, &u.Email, &u.DisplayName, &u.ID, &u.FirebaseUID, &u.Email, &u.DisplayName,
&u.StripeCustomerID, &u.StripeSubscriptionID, &u.SubscriptionStatus, &u.StripeCustomerID, &u.StripeSubscriptionID, &u.SubscriptionStatus,
&u.SubscriptionCreatedAt, &u.CreatedAt, &u.UpdatedAt, &u.SubscriptionTier, &u.SubscriptionCreatedAt, &u.CreatedAt, &u.UpdatedAt,
) )
if err != nil { if err != nil {
if errors.Is(err, pgx.ErrNoRows) { if errors.Is(err, pgx.ErrNoRows) {
@@ -66,15 +66,37 @@ func (m *UserModel) UpdateStripeCustomer(ctx context.Context, userID, stripeCust
return err 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, _, err := m.pool.Exec(ctx,
`UPDATE users `UPDATE users
SET stripe_subscription_id = $2, SET stripe_subscription_id = $2,
subscription_status = $3, subscription_status = $3,
subscription_tier = $4,
subscription_created_at = COALESCE(subscription_created_at, NOW()), subscription_created_at = COALESCE(subscription_created_at, NOW()),
updated_at = NOW() updated_at = NOW()
WHERE stripe_customer_id = $1`, 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 return err
} }

View File

@@ -1,15 +1,29 @@
import { getAuthHeaders } from "./authHeaders"; import { getAuthHeaders } from "./authHeaders";
import { API_BASE } from "./config"; import { API_BASE } from "./config";
export async function createCheckoutSession() { export async function createCheckoutSession(tier = "premium") {
const headers = await getAuthHeaders(); const headers = await getAuthHeaders();
const res = await fetch(`${API_BASE}/stripe/create-checkout-session`, { 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", method: "POST",
headers, headers,
}); });
if (!res.ok) { if (!res.ok) {
const data = await res.json(); 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(); return res.json();
} }

View File

@@ -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;
}
}

View File

@@ -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 (
<div className="tier-picker">
{TIERS.map((tier) => (
<div
key={tier.id}
className={`tier-picker__card${tier.highlighted ? " tier-picker__card--highlighted" : ""}${selectedTier === tier.id ? " tier-picker__card--selected" : ""}`}
>
{tier.highlighted && (
<div className="tier-picker__badge">Most Popular</div>
)}
<h3 className="tier-picker__name">{tier.name}</h3>
<div className="tier-picker__price">
<span className="tier-picker__amount">{tier.price}</span>
{tier.period && (
<span className="tier-picker__period">{tier.period}</span>
)}
</div>
<ul className="tier-picker__features">
{tier.features.map((f) => (
<li key={f} className="tier-picker__feature">
<span className="tier-picker__check">&#10003;</span> {f}
</li>
))}
{tier.disabledFeatures.map((f) => (
<li key={f} className="tier-picker__feature tier-picker__feature--disabled">
<span className="tier-picker__dash">&mdash;</span> {f}
</li>
))}
</ul>
<button
className={`btn tier-picker__btn${selectedTier === tier.id ? " tier-picker__btn--selected" : ""}`}
onClick={() => onSelect(tier.id)}
>
{selectedTier === tier.id ? "Selected" : "Select"}
</button>
</div>
))}
</div>
);
}

View File

@@ -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;
}

View File

@@ -0,0 +1,18 @@
import { useNavigate } from "react-router-dom";
import "./UpgradeBanner.css";
export default function UpgradeBanner({ message, linkTo = "/account" }) {
const navigate = useNavigate();
return (
<div className="upgrade-banner">
<span className="upgrade-banner__message">{message}</span>
<button
className="upgrade-banner__link"
onClick={() => navigate(linkTo)}
>
Upgrade
</button>
</div>
);
}

View File

@@ -1,10 +1,10 @@
/** /**
* AuthContext - Firebase Authentication State Management * 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 { import {
createUserWithEmailAndPassword, createUserWithEmailAndPassword,
signInWithEmailAndPassword, signInWithEmailAndPassword,
@@ -12,9 +12,16 @@ import {
onAuthStateChanged, onAuthStateChanged,
} from "firebase/auth"; } from "firebase/auth";
import { auth } from "../firebase/config"; import { auth } from "../firebase/config";
import { getAccount } from "../api/account";
const AuthContext = createContext(); const AuthContext = createContext();
const DEFAULT_TIER_LIMITS = {
max_addresses: 1,
max_alert_types: 1,
allowed_channels: ["email"],
};
/** /**
* Hook to access auth context * Hook to access auth context
* @returns {Object} Auth context value * @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 }) { export function AuthProvider({ children }) {
const [currentUser, setCurrentUser] = useState(null); const [currentUser, setCurrentUser] = useState(null);
const [loading, setLoading] = useState(true); const [loading, setLoading] = useState(true);
const [error, setError] = useState(null); const [error, setError] = useState(null);
/** const [userTier, setUserTier] = useState("free");
* Sign up with email and password 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) { async function signup(email, password) {
try { try {
setError(null); setError(null);
@@ -53,9 +72,6 @@ export function AuthProvider({ children }) {
} }
} }
/**
* Log in with email and password
*/
async function login(email, password) { async function login(email, password) {
try { try {
setError(null); setError(null);
@@ -71,12 +87,12 @@ export function AuthProvider({ children }) {
} }
} }
/**
* Log out current user
*/
async function logout() { async function logout() {
try { try {
setError(null); setError(null);
setUserTier("free");
setTierLimits(DEFAULT_TIER_LIMITS);
setAddressCount(0);
await signOut(auth); await signOut(auth);
} catch (err) { } catch (err) {
setError(err.message); setError(err.message);
@@ -84,18 +100,16 @@ export function AuthProvider({ children }) {
} }
} }
/**
* Listen for auth state changes
*/
useEffect(() => { useEffect(() => {
const unsubscribe = onAuthStateChanged(auth, (user) => { const unsubscribe = onAuthStateChanged(auth, (user) => {
setCurrentUser(user); setCurrentUser(user);
setLoading(false); setLoading(false);
if (user) {
refreshAccount();
}
}); });
// Cleanup subscription
return unsubscribe; return unsubscribe;
}, []); }, [refreshAccount]);
const value = { const value = {
currentUser, currentUser,
@@ -104,6 +118,10 @@ export function AuthProvider({ children }) {
logout, logout,
error, error,
loading, loading,
userTier,
tierLimits,
addressCount,
refreshAccount,
}; };
return ( return (

View File

@@ -1,15 +1,22 @@
import { useState, useEffect } from "react"; import { useState, useEffect } from "react";
import { useAuth } from "../../contexts/AuthContext";
import AddressForm from "../../components/AddressForm"; import AddressForm from "../../components/AddressForm";
import UpgradeBanner from "../../components/UpgradeBanner";
import { getAddresses, createAddress, deleteAddress, updateAddress } from "../../api/addresses"; import { getAddresses, createAddress, deleteAddress, updateAddress } from "../../api/addresses";
import "./Addresses.css"; import "./Addresses.css";
export default function Addresses() { export default function Addresses() {
const { tierLimits, refreshAccount } = useAuth();
const [addresses, setAddresses] = useState([]); const [addresses, setAddresses] = useState([]);
const [loading, setLoading] = useState(true); const [loading, setLoading] = useState(true);
const [error, setError] = useState(null); const [error, setError] = useState(null);
const [editingId, setEditingId] = useState(null); const [editingId, setEditingId] = useState(null);
const [editLabel, setEditLabel] = useState(""); const [editLabel, setEditLabel] = useState("");
const maxAddresses = tierLimits.max_addresses;
const isUnlimited = maxAddresses === -1;
const atLimit = !isUnlimited && addresses.length >= maxAddresses;
useEffect(() => { useEffect(() => {
async function fetchAddresses() { async function fetchAddresses() {
try { try {
@@ -32,6 +39,7 @@ export default function Addresses() {
const newAddress = await createAddress(data); const newAddress = await createAddress(data);
setAddresses((prev) => [...prev, newAddress]); setAddresses((prev) => [...prev, newAddress]);
setError(null); setError(null);
refreshAccount();
} catch (err) { } catch (err) {
setError(err.message); setError(err.message);
console.error("Failed to create address:", err); console.error("Failed to create address:", err);
@@ -47,6 +55,7 @@ export default function Addresses() {
await deleteAddress(id); await deleteAddress(id);
setAddresses((prev) => prev.filter((a) => a.id !== id)); setAddresses((prev) => prev.filter((a) => a.id !== id));
setError(null); setError(null);
refreshAccount();
} catch (err) { } catch (err) {
setError(err.message); setError(err.message);
console.error("Failed to delete address:", err); console.error("Failed to delete address:", err);
@@ -80,9 +89,17 @@ export default function Addresses() {
<div className="page"> <div className="page">
<h1>Add Addresses to Track</h1> <h1>Add Addresses to Track</h1>
<div className="mb-xl"> {atLimit && (
<AddressForm onSubmit={handleAddressSubmit} /> <UpgradeBanner
</div> message={`Your plan allows ${maxAddresses} address${maxAddresses !== 1 ? "es" : ""}. Upgrade to track more.`}
/>
)}
{!atLimit && (
<div className="mb-xl">
<AddressForm onSubmit={handleAddressSubmit} />
</div>
)}
<div> <div>
<h2>Existing Tracked Addresses</h2> <h2>Existing Tracked Addresses</h2>

View File

@@ -125,6 +125,20 @@
margin-top: 1rem; 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 ──────────────────────────────── */ /* ── Alerts Page Responsive ──────────────────────────────── */
/* Laptop: tighten the gap */ /* Laptop: tighten the gap */

View File

@@ -1,7 +1,9 @@
import { useState, useEffect } from "react"; import { useState, useEffect } from "react";
import { useAuth } from "../../contexts/AuthContext";
import AlertForm from "../../components/AlertForm"; import AlertForm from "../../components/AlertForm";
import Button from "../../components/Button"; import Button from "../../components/Button";
import Input from "../../components/Input"; import Input from "../../components/Input";
import UpgradeBanner from "../../components/UpgradeBanner";
import { getAddresses } from "../../api/addresses"; import { getAddresses } from "../../api/addresses";
import { import {
getAlerts, getAlerts,
@@ -20,6 +22,8 @@ import {
import "./Alerts.css"; import "./Alerts.css";
export default function Alerts() { export default function Alerts() {
const { tierLimits, userTier } = useAuth();
const [addresses, setAddresses] = useState([]); const [addresses, setAddresses] = useState([]);
const [selectedAddressId, setSelectedAddressId] = useState(null); const [selectedAddressId, setSelectedAddressId] = useState(null);
const [alerts, setAlerts] = useState([]); const [alerts, setAlerts] = useState([]);
@@ -43,6 +47,21 @@ export default function Alerts() {
const [openAccordions, setOpenAccordions] = useState({}); const [openAccordions, setOpenAccordions] = useState({});
const [thresholdEdits, setThresholdEdits] = 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) { function toggleAccordion(alertId) {
setOpenAccordions((prev) => ({ ...prev, [alertId]: !prev[alertId] })); setOpenAccordions((prev) => ({ ...prev, [alertId]: !prev[alertId] }));
} }
@@ -209,11 +228,11 @@ export default function Alerts() {
const config = { const config = {
notification_enabled: notificationEnabled, notification_enabled: notificationEnabled,
discord_webhook_url: discordWebhookUrl || null, discord_webhook_url: canDiscord ? (discordWebhookUrl || null) : null,
telegram_bot_token: telegramBotToken || null, telegram_bot_token: canTelegram ? (telegramBotToken || null) : null,
telegram_chat_id: telegramChatId || null, telegram_chat_id: canTelegram ? (telegramChatId || null) : null,
email: email || null, email: email || null,
slack_webhook_url: slackWebhookUrl || null, slack_webhook_url: canSlack ? (slackWebhookUrl || null) : null,
}; };
await updateNotificationConfig(config); await updateNotificationConfig(config);
@@ -357,10 +376,18 @@ export default function Alerts() {
</div> </div>
</div> </div>
<div className="mb-xl"> {atAlertLimit && userTier !== "pro" && (
<h3>Create New Alert</h3> <UpgradeBanner
<AlertForm onSubmit={handleAlertSubmit} /> message={`Your ${userTier} plan allows ${maxAlertTypes} alert type${maxAlertTypes !== 1 ? "s" : ""} per address. Upgrade for more.`}
</div> />
)}
{!atAlertLimit && (
<div className="mb-xl">
<h3>Create New Alert</h3>
<AlertForm onSubmit={handleAlertSubmit} />
</div>
)}
<div> <div>
<h3>Active Alert Rules</h3> <h3>Active Alert Rules</h3>
@@ -540,19 +567,24 @@ export default function Alerts() {
{notificationEnabled && ( {notificationEnabled && (
<> <>
{/* Telegram */} {/* Telegram */}
<div className="section"> <div className={`section${!canTelegram ? " alerts__channel-locked" : ""}`}>
<h3 className="mt-0 mb-md">Telegram</h3> <h3 className="mt-0 mb-md">Telegram</h3>
{!canTelegram && (
<UpgradeBanner message="Upgrade to Premium to enable Telegram alerts" />
)}
<Input <Input
label="Bot Token" label="Bot Token"
value={telegramBotToken} value={telegramBotToken}
onChange={setTelegramBotToken} onChange={setTelegramBotToken}
placeholder="123456789:ABCdefGHIjklMNOpqrSTUvwxYZ" placeholder="123456789:ABCdefGHIjklMNOpqrSTUvwxYZ"
disabled={!canTelegram}
/> />
<Input <Input
label="Chat ID" label="Chat ID"
value={telegramChatId} value={telegramChatId}
onChange={setTelegramChatId} onChange={setTelegramChatId}
placeholder="-1001234567890" placeholder="-1001234567890"
disabled={!canTelegram}
/> />
<a <a
href="https://core.telegram.org/bots#how-do-i-create-a-bot" href="https://core.telegram.org/bots#how-do-i-create-a-bot"
@@ -597,13 +629,17 @@ export default function Alerts() {
</div> </div>
{/* Discord */} {/* Discord */}
<div className="section"> <div className={`section${!canDiscord ? " alerts__channel-locked" : ""}`}>
<h3 className="mt-0 mb-md">Discord</h3> <h3 className="mt-0 mb-md">Discord</h3>
{!canDiscord && (
<UpgradeBanner message="Upgrade to Premium to enable Discord alerts" />
)}
<Input <Input
label="Discord Webhook URL" label="Discord Webhook URL"
value={discordWebhookUrl} value={discordWebhookUrl}
onChange={setDiscordWebhookUrl} onChange={setDiscordWebhookUrl}
placeholder="https://discord.com/api/webhooks/..." placeholder="https://discord.com/api/webhooks/..."
disabled={!canDiscord}
/> />
<a <a
href="https://support.discord.com/hc/en-us/articles/228383668-Intro-to-Webhooks" href="https://support.discord.com/hc/en-us/articles/228383668-Intro-to-Webhooks"
@@ -616,13 +652,17 @@ export default function Alerts() {
</div> </div>
{/* Slack */} {/* Slack */}
<div className="section"> <div className={`section${!canSlack ? " alerts__channel-locked" : ""}`}>
<h3 className="mt-0 mb-md">Slack</h3> <h3 className="mt-0 mb-md">Slack</h3>
{!canSlack && (
<UpgradeBanner message="Upgrade to Pro to enable Slack alerts" />
)}
<Input <Input
label="Slack Webhook URL" label="Slack Webhook URL"
value={slackWebhookUrl} value={slackWebhookUrl}
onChange={setSlackWebhookUrl} onChange={setSlackWebhookUrl}
placeholder="https://hooks.slack.com/services/..." placeholder="https://hooks.slack.com/services/..."
disabled={!canSlack}
/> />
<a <a
href="https://api.slack.com/messaging/webhooks" href="https://api.slack.com/messaging/webhooks"

View File

@@ -15,6 +15,10 @@
padding: 0 1rem; padding: 0 1rem;
} }
.subscribe__container--wide {
max-width: 920px;
}
.subscribe__title { .subscribe__title {
text-align: center; text-align: center;
margin-bottom: 2rem; margin-bottom: 2rem;
@@ -237,6 +241,29 @@
font-style: italic; font-style: italic;
} }
/* Tier limit hints */
.subscribe__tier-hint {
font-size: 0.82rem;
color: #ffcc44;
margin-top: 0.5rem;
}
/* Disabled channel group */
.subscribe__channel-disabled {
opacity: 0.5;
pointer-events: none;
}
.subscribe__channel-disabled .subscribe__tier-hint {
pointer-events: auto;
opacity: 1;
}
/* Disabled checkbox row */
.checkbox-row--disabled {
opacity: 0.5;
}
/* ── Subscribe Responsive ────────────────────────────────── */ /* ── Subscribe Responsive ────────────────────────────────── */
@media (max-width: 480px) { @media (max-width: 480px) {

View File

@@ -7,12 +7,20 @@ import {
updateNotificationConfig, updateNotificationConfig,
testNotificationChannels, testNotificationChannels,
} from "../../api/notificationConfig"; } from "../../api/notificationConfig";
import { createCheckoutSession, getSubscriptionStatus, verifyCheckoutSession } from "../../api/stripe"; import { createCheckoutSession, getSubscriptionStatus, verifyCheckoutSession, activateFreeTier } from "../../api/stripe";
import Input from "../../components/Input"; import Input from "../../components/Input";
import Button from "../../components/Button"; import Button from "../../components/Button";
import TierPicker from "../../components/TierPicker";
import "./Subscribe.css"; import "./Subscribe.css";
const TIER_LIMITS = {
free: { maxAlertTypes: 1, channels: ["email"] },
premium: { maxAlertTypes: 2, channels: ["email", "discord", "telegram"] },
pro: { maxAlertTypes: 4, channels: ["email", "discord", "telegram", "slack"] },
};
const STEPS = [ const STEPS = [
"Choose Plan",
"Create Account", "Create Account",
"Add Wallet", "Add Wallet",
"Alert Rules", "Alert Rules",
@@ -33,6 +41,7 @@ export default function Subscribe() {
const [testLoading, setTestLoading] = useState(false); const [testLoading, setTestLoading] = useState(false);
const [data, setData] = useState({ const [data, setData] = useState({
selectedTier: "",
email: "", email: "",
password: "", password: "",
confirmPassword: "", confirmPassword: "",
@@ -56,6 +65,8 @@ export default function Subscribe() {
setData((prev) => ({ ...prev, [field]: value })); setData((prev) => ({ ...prev, [field]: value }));
} }
const tierLimits = TIER_LIMITS[data.selectedTier] || TIER_LIMITS.free;
useEffect(() => { useEffect(() => {
if (!currentUser) return; if (!currentUser) return;
getAddresses() getAddresses()
@@ -67,7 +78,6 @@ export default function Subscribe() {
.catch(() => { }); .catch(() => { });
}, [currentUser, navigate]); }, [currentUser, navigate]);
// Handle Stripe redirect back from checkout
useEffect(() => { useEffect(() => {
if (!currentUser) return; if (!currentUser) return;
const payment = searchParams.get("payment"); const payment = searchParams.get("payment");
@@ -77,23 +87,32 @@ export default function Subscribe() {
setLoading(true); setLoading(true);
verifyCheckoutSession(sessionId) verifyCheckoutSession(sessionId)
.then(() => { .then(() => {
setStep(2); setStep(3);
}) })
.catch((err) => { .catch((err) => {
setError("Payment verification failed: " + err.message); setError("Payment verification failed: " + err.message);
setStep(1); setStep(2);
}) })
.finally(() => setLoading(false)); .finally(() => setLoading(false));
} else if (payment === "cancelled") { } else if (payment === "cancelled") {
setSearchParams({}, { replace: true }); setSearchParams({}, { replace: true });
setStep(1); setStep(2);
setError("Payment was cancelled. Please try again."); setError("Payment was cancelled. Please try again.");
} }
}, [currentUser, searchParams, setSearchParams]); }, [currentUser, searchParams, setSearchParams]);
// ── Step handlers ───────────────────────────────────────────────────────── // ── 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(""); setError("");
if (!data.email || !data.password || !data.confirmPassword) { if (!data.email || !data.password || !data.confirmPassword) {
setError("Please fill in all fields"); setError("Please fill in all fields");
@@ -112,12 +131,19 @@ export default function Subscribe() {
if (!currentUser) { if (!currentUser) {
await signup(data.email, data.password); await signup(data.email, data.password);
} }
const status = await getSubscriptionStatus();
if (status.subscription_status === "active" || status.subscription_status === "trialing") { if (data.selectedTier === "free") {
setStep(2); await activateFreeTier();
setStep(3);
return; 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; window.location.href = url;
} catch (err) { } catch (err) {
if (err.code === "auth/email-already-in-use") { if (err.code === "auth/email-already-in-use") {
@@ -133,7 +159,7 @@ export default function Subscribe() {
} }
} }
async function handleStep2() { async function handleStep3() {
setError(""); setError("");
if (!data.walletAddress) { if (!data.walletAddress) {
setError("Please enter a wallet address"); setError("Please enter a wallet address");
@@ -150,7 +176,7 @@ export default function Subscribe() {
label: data.walletLabel || undefined, label: data.walletLabel || undefined,
}); });
set("createdAddressId", created.id); set("createdAddressId", created.id);
setStep(3); setStep(4);
} catch (err) { } catch (err) {
setError(err.message); setError(err.message);
} finally { } finally {
@@ -158,7 +184,7 @@ export default function Subscribe() {
} }
} }
async function handleStep3() { async function handleStep4() {
setError(""); setError("");
const rules = []; const rules = [];
if (data.alertIncomingTx) rules.push({ type: "incoming_tx" }); if (data.alertIncomingTx) rules.push({ type: "incoming_tx" });
@@ -179,7 +205,7 @@ export default function Subscribe() {
} }
if (rules.length === 0) { if (rules.length === 0) {
setStep(4); setStep(5);
return; return;
} }
@@ -191,7 +217,7 @@ export default function Subscribe() {
created.push(result); created.push(result);
} }
set("alertsCreated", created); set("alertsCreated", created);
setStep(4); setStep(5);
} catch (err) { } catch (err) {
setError(err.message); setError(err.message);
} finally { } finally {
@@ -199,12 +225,12 @@ export default function Subscribe() {
} }
} }
async function handleStep4() { async function handleStep5() {
setError(""); setError("");
const hasAny = const hasAny =
data.discordWebhookUrl || data.slackWebhookUrl || data.notificationEmail; data.discordWebhookUrl || data.slackWebhookUrl || data.notificationEmail;
if (!hasAny) { if (!hasAny) {
setStep(5); setStep(6);
return; return;
} }
try { try {
@@ -216,7 +242,7 @@ export default function Subscribe() {
email: data.notificationEmail || undefined, email: data.notificationEmail || undefined,
}); });
set("notificationConfigured", true); set("notificationConfigured", true);
setStep(5); setStep(6);
} catch (err) { } catch (err) {
setError(err.message); setError(err.message);
} finally { } 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 ────────────────────────────────────────────────────────── // ── Progress bar ──────────────────────────────────────────────────────────
function ProgressBar() { function ProgressBar() {
@@ -284,7 +325,22 @@ export default function Subscribe() {
// ── Step content ────────────────────────────────────────────────────────── // ── Step content ──────────────────────────────────────────────────────────
function Step1() { function StepChoosePlan() {
return (
<>
<h2 className="mb-sm">Choose your plan</h2>
<p className="subscribe__subtitle">
Select the plan that works best for you. You can upgrade anytime.
</p>
<TierPicker
onSelect={(tier) => set("selectedTier", tier)}
selectedTier={data.selectedTier}
/>
</>
);
}
function StepCreateAccount() {
return ( return (
<> <>
<h2 className="mb-lg">Create your account</h2> <h2 className="mb-lg">Create your account</h2>
@@ -317,7 +373,7 @@ export default function Subscribe() {
); );
} }
function Step2() { function StepAddWallet() {
return ( return (
<> <>
<h2 className="mb-sm">Add a wallet address</h2> <h2 className="mb-sm">Add a wallet address</h2>
@@ -343,28 +399,35 @@ export default function Subscribe() {
); );
} }
function Step3() { function StepAlertRules() {
const atLimit = !canSelectMoreAlerts();
const maxTypes = tierLimits.maxAlertTypes;
return ( return (
<> <>
<h2 className="mb-sm">Configure alert rules</h2> <h2 className="mb-sm">Configure alert rules</h2>
<p className="subscribe__subtitle"> <p className="subscribe__subtitle">
Choose which events trigger notifications. You can change these later. Choose which events trigger notifications ({countSelectedAlerts()}/{maxTypes} selected).
You can change these later.
</p> </p>
<CheckboxRow <CheckboxRow
checked={data.alertIncomingTx} checked={data.alertIncomingTx}
onChange={(v) => set("alertIncomingTx", v)} onChange={(v) => set("alertIncomingTx", v)}
label="Incoming transaction" label="Incoming transaction"
disabled={!data.alertIncomingTx && atLimit}
/> />
<CheckboxRow <CheckboxRow
checked={data.alertOutgoingTx} checked={data.alertOutgoingTx}
onChange={(v) => set("alertOutgoingTx", v)} onChange={(v) => set("alertOutgoingTx", v)}
label="Outgoing transaction" label="Outgoing transaction"
disabled={!data.alertOutgoingTx && atLimit}
/> />
<CheckboxRow <CheckboxRow
checked={data.alertLargeTransfer} checked={data.alertLargeTransfer}
onChange={(v) => set("alertLargeTransfer", v)} onChange={(v) => set("alertLargeTransfer", v)}
label="Large transfer" label="Large transfer"
disabled={!data.alertLargeTransfer && atLimit}
> >
{data.alertLargeTransfer && ( {data.alertLargeTransfer && (
<div className="checkbox-row__nested"> <div className="checkbox-row__nested">
@@ -384,6 +447,7 @@ export default function Subscribe() {
checked={data.alertBalanceBelow} checked={data.alertBalanceBelow}
onChange={(v) => set("alertBalanceBelow", v)} onChange={(v) => set("alertBalanceBelow", v)}
label="Balance below" label="Balance below"
disabled={!data.alertBalanceBelow && atLimit}
> >
{data.alertBalanceBelow && ( {data.alertBalanceBelow && (
<div className="checkbox-row__nested"> <div className="checkbox-row__nested">
@@ -399,11 +463,21 @@ export default function Subscribe() {
</div> </div>
)} )}
</CheckboxRow> </CheckboxRow>
{atLimit && data.selectedTier !== "pro" && (
<p className="subscribe__tier-hint">
Your {data.selectedTier} plan allows {maxTypes} alert type{maxTypes !== 1 ? "s" : ""} per address. Upgrade for more.
</p>
)}
</> </>
); );
} }
function Step4() { function StepNotifications() {
const channels = tierLimits.channels;
const canDiscord = channels.includes("discord");
const canSlack = channels.includes("slack");
return ( return (
<> <>
<h2 className="mb-sm">Set up notifications</h2> <h2 className="mb-sm">Set up notifications</h2>
@@ -411,7 +485,17 @@ export default function Subscribe() {
Add at least one channel so you receive alerts. All fields are optional. Add at least one channel so you receive alerts. All fields are optional.
</p> </p>
<div className="mb-md"> <Input
label="Email address for alerts"
type="email"
value={data.notificationEmail}
onChange={(v) => set("notificationEmail", v)}
disabled={loading}
placeholder="you@example.com"
className="form-field--last"
/>
<div className={`mb-md${!canDiscord ? " subscribe__channel-disabled" : ""}`}>
<label className="form-label"> <label className="form-label">
Discord Webhook URL{" "} Discord Webhook URL{" "}
<a <a
@@ -428,12 +512,15 @@ export default function Subscribe() {
type="url" type="url"
value={data.discordWebhookUrl} value={data.discordWebhookUrl}
onChange={(v) => set("discordWebhookUrl", v)} onChange={(v) => set("discordWebhookUrl", v)}
disabled={loading} disabled={loading || !canDiscord}
placeholder="https://discord.com/api/webhooks/..." placeholder="https://discord.com/api/webhooks/..."
/> />
{!canDiscord && (
<p className="subscribe__tier-hint">Upgrade to Premium to enable Discord alerts</p>
)}
</div> </div>
<div className="mb-md"> <div className={`mb-md${!canSlack ? " subscribe__channel-disabled" : ""}`}>
<label className="form-label"> <label className="form-label">
Slack Webhook URL{" "} Slack Webhook URL{" "}
<a <a
@@ -450,25 +537,18 @@ export default function Subscribe() {
type="url" type="url"
value={data.slackWebhookUrl} value={data.slackWebhookUrl}
onChange={(v) => set("slackWebhookUrl", v)} onChange={(v) => set("slackWebhookUrl", v)}
disabled={loading} disabled={loading || !canSlack}
placeholder="https://hooks.slack.com/services/..." placeholder="https://hooks.slack.com/services/..."
/> />
{!canSlack && (
<p className="subscribe__tier-hint">Upgrade to Pro to enable Slack alerts</p>
)}
</div> </div>
<Input
label="Email address for alerts"
type="email"
value={data.notificationEmail}
onChange={(v) => set("notificationEmail", v)}
disabled={loading}
placeholder="you@example.com"
className="form-field--last"
/>
</> </>
); );
} }
function Step5() { function StepDone() {
const alertCount = data.alertsCreated.length; const alertCount = data.alertsCreated.length;
const hasNotif = data.notificationConfigured; const hasNotif = data.notificationConfigured;
@@ -479,6 +559,12 @@ export default function Subscribe() {
<div className="subscribe__summary"> <div className="subscribe__summary">
<p className="subscribe__summary-title">Summary</p> <p className="subscribe__summary-title">Summary</p>
<ul className="subscribe__summary-list"> <ul className="subscribe__summary-list">
<li>
Plan:{" "}
<span className="text-white">
{data.selectedTier === "pro" ? "Pro" : data.selectedTier === "premium" ? "Premium" : "Free Trial"}
</span>
</li>
<li> <li>
Wallet address added:{" "} Wallet address added:{" "}
<span className="text-mono text-white-sm"> <span className="text-mono text-white-sm">
@@ -542,15 +628,16 @@ export default function Subscribe() {
// ── Shared helpers ──────────────────────────────────────────────────────── // ── Shared helpers ────────────────────────────────────────────────────────
function CheckboxRow({ checked, onChange, label, children }) { function CheckboxRow({ checked, onChange, label, children, disabled }) {
return ( return (
<div className="checkbox-row"> <div className={`checkbox-row${disabled ? " checkbox-row--disabled" : ""}`}>
<label className="checkbox-row__label"> <label className="checkbox-row__label">
<input <input
type="checkbox" type="checkbox"
checked={checked} checked={checked}
onChange={(e) => onChange(e.target.checked)} onChange={(e) => onChange(e.target.checked)}
className="checkbox-row__input" className="checkbox-row__input"
disabled={disabled}
/> />
{label} {label}
</label> </label>
@@ -562,17 +649,18 @@ export default function Subscribe() {
// ── Footer navigation ───────────────────────────────────────────────────── // ── Footer navigation ─────────────────────────────────────────────────────
function Footer() { function Footer() {
if (step === 5) return null; if (step === 6) return null;
const canSkip = step === 3 || step === 4; const canSkip = step === 4 || step === 5;
const canBack = step > 2; const canBack = step > 1 && step <= 5;
async function handleNext() { async function handleNext() {
setSkipWarning(""); setSkipWarning("");
if (step === 1) await handleStep1(); if (step === 1) handleStep1();
else if (step === 2) await handleStep2(); else if (step === 2) await handleStep2();
else if (step === 3) await handleStep3(); else if (step === 3) await handleStep3();
else if (step === 4) await handleStep4(); else if (step === 4) await handleStep4();
else if (step === 5) await handleStep5();
} }
function handleSkip() { function handleSkip() {
@@ -588,10 +676,14 @@ export default function Subscribe() {
} }
const nextLabel = step === 1 const nextLabel = step === 1
? "Create Account & Subscribe" ? "Continue"
: step === 4 : step === 2
? "Finish" ? data.selectedTier === "free"
: "Next →"; ? "Create Account"
: "Create Account & Subscribe"
: step === 5
? "Finish"
: "Next →";
return ( return (
<div className="subscribe__footer"> <div className="subscribe__footer">
@@ -619,7 +711,7 @@ export default function Subscribe() {
)} )}
<Button <Button
onClick={handleNext} onClick={handleNext}
disabled={loading} disabled={loading || (step === 1 && !data.selectedTier)}
className="text-bold" className="text-bold"
> >
{loading ? "Please wait..." : nextLabel} {loading ? "Please wait..." : nextLabel}
@@ -632,16 +724,17 @@ export default function Subscribe() {
// ── Render ──────────────────────────────────────────────────────────────── // ── Render ────────────────────────────────────────────────────────────────
const stepContent = { const stepContent = {
1: Step1(), 1: StepChoosePlan(),
2: Step2(), 2: StepCreateAccount(),
3: Step3(), 3: StepAddWallet(),
4: Step4(), 4: StepAlertRules(),
5: Step5(), 5: StepNotifications(),
6: StepDone(),
}; };
return ( return (
<div className="subscribe"> <div className="subscribe">
<div className="subscribe__container"> <div className={`subscribe__container${step === 1 ? " subscribe__container--wide" : ""}`}>
<h1 className="subscribe__title">Koin Ping</h1> <h1 className="subscribe__title">Koin Ping</h1>
{ProgressBar()} {ProgressBar()}
@@ -659,7 +752,7 @@ export default function Subscribe() {
{Footer()} {Footer()}
</div> </div>
{step === 1 && ( {step === 2 && (
<p className="subscribe__login-link"> <p className="subscribe__login-link">
Already have an account?{" "} Already have an account?{" "}
<a href="/login">Log in here</a> <a href="/login">Log in here</a>

View File

@@ -1,10 +1,18 @@
import { useState, useEffect } from "react"; import { useState, useEffect } from "react";
import { useNavigate } from "react-router-dom";
import { updatePassword } from "firebase/auth"; import { updatePassword } from "firebase/auth";
import { auth } from "../../firebase/config"; import { auth } from "../../firebase/config";
import { getAccount, createPortalSession } from "../../api/account"; import { getAccount, createPortalSession } from "../../api/account";
import "./Account.css"; import "./Account.css";
const TIER_LABELS = {
free: "Free Trial",
premium: "Premium",
pro: "Pro",
};
export default function Account() { export default function Account() {
const navigate = useNavigate();
const [account, setAccount] = useState(null); const [account, setAccount] = useState(null);
const [loading, setLoading] = useState(true); const [loading, setLoading] = useState(true);
const [error, setError] = useState(null); const [error, setError] = useState(null);
@@ -78,6 +86,7 @@ export default function Account() {
return <div className="page text-error">Error: {error}</div>; return <div className="page text-error">Error: {error}</div>;
} }
const tier = account.subscription_tier || "free";
const isCanceling = account.cancel_at_period_end; const isCanceling = account.cancel_at_period_end;
const statusLabel = isCanceling const statusLabel = isCanceling
? "Canceling" ? "Canceling"
@@ -86,6 +95,9 @@ export default function Account() {
: account.subscription_status.charAt(0).toUpperCase() + : account.subscription_status.charAt(0).toUpperCase() +
account.subscription_status.slice(1); account.subscription_status.slice(1);
const canUpgrade = tier === "free" || tier === "premium";
const hasPaidSub = tier !== "free";
return ( return (
<div className="page account-page"> <div className="page account-page">
<h1 className="mb-lg">Account</h1> <h1 className="mb-lg">Account</h1>
@@ -114,7 +126,7 @@ export default function Account() {
<h2 className="account__section-title">Subscription</h2> <h2 className="account__section-title">Subscription</h2>
<div className="account__row"> <div className="account__row">
<span className="account__label">Plan</span> <span className="account__label">Plan</span>
<span className="account__value">{account.subscription_plan}</span> <span className="account__value">{TIER_LABELS[tier] || account.subscription_plan}</span>
</div> </div>
<div className="account__row"> <div className="account__row">
<span className="account__label">Status</span> <span className="account__label">Status</span>
@@ -149,15 +161,27 @@ export default function Account() {
)} )}
<div className="account__portal-section"> <div className="account__portal-section">
<button {canUpgrade && (
onClick={handleManageSubscription} <button
disabled={portalLoading} onClick={() => navigate("/subscribe")}
className="btn btn--ghost" className="btn btn--primary"
> >
{portalLoading ? "Redirecting..." : "Manage Subscription"} Upgrade Plan
</button> </button>
)}
{hasPaidSub && (
<button
onClick={handleManageSubscription}
disabled={portalLoading}
className="btn btn--ghost"
>
{portalLoading ? "Redirecting..." : "Manage Subscription"}
</button>
)}
<p className="text-dimmed text-sm account__portal-hint"> <p className="text-dimmed text-sm account__portal-hint">
Cancel subscription, update payment method, or view invoices via Stripe. {hasPaidSub
? "Cancel subscription, update payment method, or view invoices via Stripe."
: "Upgrade to unlock more addresses, alert types, and notification channels."}
</p> </p>
</div> </div>
</div> </div>