Compare commits

..

9 Commits

Author SHA1 Message Date
KS Jannette
40ed4f6afd m 2026-03-28 22:34:23 -04:00
KS Jannette
26324150d2 removed F up 2026-03-28 11:58:29 -04:00
KS Jannette
b0572451d3 but 2026-03-28 11:55:51 -04:00
S Jannette
d9c3bd1db5 Merge pull request #27 from kjannette/setup-free-tier
Setup free tier
2026-03-28 11:27:18 -04:00
KS Jannette
615dd1dddc Chockfull 2026-03-28 11:10:25 -04:00
KS Jannette
288092e4b4 More 2026-03-28 10:31:14 -04:00
KS Jannette
69a2112df9 update readme 2026-03-28 08:27:06 -04:00
KS Jannette
4ba91c7d9b Stripe integration tweaks 2026-03-10 17:18:54 -04:00
KS Jannette
f9fa7def2b add subscription tiers 2026-03-10 15:22:48 -04:00
40 changed files with 1384 additions and 223 deletions

2
.gitignore vendored
View File

@@ -22,6 +22,8 @@ node_modules/
# Go build artifacts # Go build artifacts
backend/bin/ backend/bin/
backend/api
backend/poller
*.exe *.exe
*.exe~ *.exe~
*.dll *.dll

View File

@@ -37,7 +37,7 @@ tidy:
# Database setup # Database setup
db-setup: db-setup:
psql -d koin_ping_dev -f infra/schema.sql psql -d koin_ping -f infra/schema.sql
vet: vet:
go vet ./... go vet ./...

View File

@@ -1,20 +1,32 @@
Start DB: Backend startup quickstarat:
-----------------------------> BEST
## 1. brew services start postgresql@15
brew services start postgresql@15 OR
From the backend directory, you have a few options: /opt/homebrew/opt/postgresql@15/bin/pg_ctl -D /opt/homebrew/var/postgresql@15 start
Option 1: Single command (both API + poller) ## 2. ALLIN ONE:
cd /Users/kjannette/workspace/koin_ping_0.2.0/backendmake dev-all Make dev-all — Runs both the API and poller concurrently.
-----------------------------> BEST
Option 2: Two separate terminals - OR -
Terminal 1 (API server):
cd /Users/kjannette/workspace/koin_ping_0.2.0/backend go run ./cmd/api
Terminal 2 (Poller): ## 3. Option 1: Single command (both API + poller)
cd /Users/kjannette/workspace/koin_ping_0.2.0/backend go run ./cmd/poller cd /Users/kjannette/workspace/koin_ping_0.2.0/backendmake dev-all
## 4. Option 2: Two separate terminals
Terminal 1 (API server):
cd /Users/kjannette/workspace/koin_ping_0.2.0/backend go run ./cmd/api
make run — Builds and runs the API server.
make dev — Runs the API server with auto-reload via air (falls back to go run if air isn't installed).
## 5. Terminal 2 (Poller):
cd /Users/kjannette/workspace/koin_ping_0.2.0/backend
go run ./cmd/poller
make poller — Builds and runs the poller.
make poller-dev — Runs the poller with auto-reload.
make run — Builds and runs the API server.
make dev — Runs the API server with auto-reload via air (falls back to go run if air isn't installed).
make poller — Builds and runs the poller.
make poller-dev — Runs the poller with auto-reload.
make dev-all — Runs both the API and poller concurrently.

Binary file not shown.

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)
@@ -82,6 +82,9 @@ func main() {
// Stripe webhook (public — called by Stripe, not authenticated) // Stripe webhook (public — called by Stripe, not authenticated)
mux.HandleFunc("POST "+b+"/stripe/webhook", stripeHandler.HandleWebhook) mux.HandleFunc("POST "+b+"/stripe/webhook", stripeHandler.HandleWebhook)
// Onboarding checkout (public — account doesn't exist yet)
mux.HandleFunc("POST "+b+"/stripe/create-onboarding-checkout", stripeHandler.CreateOnboardingCheckout)
// Stripe routes (auth required, NO subscription required) // Stripe routes (auth required, NO subscription required)
mux.Handle("POST "+b+"/stripe/create-checkout-session", mux.Handle("POST "+b+"/stripe/create-checkout-session",
authenticate(http.HandlerFunc(stripeHandler.CreateCheckoutSession))) authenticate(http.HandlerFunc(stripeHandler.CreateCheckoutSession)))
@@ -91,6 +94,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

@@ -51,6 +51,7 @@ func main() {
defer database.Close() defer database.Close()
userModel := models.NewUserModel(pool)
addressModel := models.NewAddressModel(pool) addressModel := models.NewAddressModel(pool)
alertRuleModel := models.NewAlertRuleModel(pool) alertRuleModel := models.NewAlertRuleModel(pool)
alertEventModel := models.NewAlertEventModel(pool) alertEventModel := models.NewAlertEventModel(pool)
@@ -59,8 +60,7 @@ func main() {
observer := services.NewObserverService(eth, addressModel, checkpointModel) observer := services.NewObserverService(eth, addressModel, checkpointModel)
evaluator := services.NewEvaluatorService( evaluator := services.NewEvaluatorService(
eth, alertRuleModel, alertEventModel, addressModel, notifConfigModel, eth, alertRuleModel, alertEventModel, addressModel, userModel, notifConfigModel,
cfg.ResendAPIKey, cfg.EmailFrom,
) )
digestSvc := services.NewEmailDigestService(cfg.ResendAPIKey, cfg.EmailFrom, alertEventModel, notifConfigModel) digestSvc := services.NewEmailDigestService(cfg.ResendAPIKey, cfg.EmailFrom, alertEventModel, notifConfigModel)

View File

@@ -0,0 +1,8 @@
ALTER TABLE users ADD COLUMN IF NOT EXISTS subscription_tier VARCHAR(20) DEFAULT 'free';
-- Existing active/trialing subscribers were on the single paid plan,
-- which is now the "premium" tier. Backfill them so they aren't downgraded.
UPDATE users
SET subscription_tier = 'premium'
WHERE subscription_status IN ('active', 'trialing')
AND stripe_subscription_id IS NOT NULL;

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"),
} }
@@ -92,6 +94,19 @@ func (c *Config) DSN() string {
) )
} }
// TierForPriceID maps a Stripe price ID back to the corresponding
// subscription tier. Returns empty string if the price is unrecognised.
func (c *Config) TierForPriceID(priceID string) string {
switch priceID {
case c.StripePriceIDPremium:
return "premium"
case c.StripePriceIDPro:
return "pro"
default:
return ""
}
}
func getEnv(key, fallback string) string { func getEnv(key, fallback string) string {
if v := os.Getenv(key); v != "" { if v := os.Getenv(key); v != "" {
return v return v

View File

@@ -0,0 +1,28 @@
package config
import "testing"
func TestTierForPriceID(t *testing.T) {
t.Parallel()
cfg := &Config{
StripePriceIDPremium: "price_premium_123",
StripePriceIDPro: "price_pro_456",
}
tests := []struct {
priceID string
want string
}{
{"price_premium_123", "premium"},
{"price_pro_456", "pro"},
{"price_unknown", ""},
{"", ""},
}
for _, tt := range tests {
if got := cfg.TierForPriceID(tt.priceID); got != tt.want {
t.Errorf("TierForPriceID(%q) = %q, want %q", tt.priceID, got, tt.want)
}
}
}

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

@@ -0,0 +1,104 @@
package domain
import "testing"
func TestIsValidTier(t *testing.T) {
t.Parallel()
valid := []string{"free", "premium", "pro"}
for _, tier := range valid {
if !IsValidTier(tier) {
t.Errorf("expected %q to be valid", tier)
}
}
invalid := []string{"", "basic", "enterprise", "FREE", "Pro"}
for _, tier := range invalid {
if IsValidTier(tier) {
t.Errorf("expected %q to be invalid", tier)
}
}
}
func TestGetTierLimits_Free(t *testing.T) {
t.Parallel()
limits := GetTierLimits(TierFree)
if limits.MaxAddresses != 1 {
t.Errorf("free MaxAddresses = %d, want 1", limits.MaxAddresses)
}
if limits.MaxAlertTypes != 1 {
t.Errorf("free MaxAlertTypes = %d, want 1", limits.MaxAlertTypes)
}
if len(limits.AllowedChannels) != 1 || limits.AllowedChannels[0] != "email" {
t.Errorf("free AllowedChannels = %v, want [email]", limits.AllowedChannels)
}
if limits.IsUnlimitedAddresses() {
t.Error("free should not have unlimited addresses")
}
if limits.IsUnlimitedAlertTypes() {
t.Error("free should not have unlimited alert types")
}
}
func TestGetTierLimits_Premium(t *testing.T) {
t.Parallel()
limits := GetTierLimits(TierPremium)
if limits.MaxAddresses != 3 {
t.Errorf("premium MaxAddresses = %d, want 3", limits.MaxAddresses)
}
if limits.MaxAlertTypes != 2 {
t.Errorf("premium MaxAlertTypes = %d, want 2", limits.MaxAlertTypes)
}
if !limits.ChannelAllowed("email") {
t.Error("premium should allow email")
}
if !limits.ChannelAllowed("discord") {
t.Error("premium should allow discord")
}
if !limits.ChannelAllowed("telegram") {
t.Error("premium should allow telegram")
}
if limits.ChannelAllowed("slack") {
t.Error("premium should NOT allow slack")
}
}
func TestGetTierLimits_Pro(t *testing.T) {
t.Parallel()
limits := GetTierLimits(TierPro)
if !limits.IsUnlimitedAddresses() {
t.Error("pro should have unlimited addresses")
}
if !limits.IsUnlimitedAlertTypes() {
t.Error("pro should have unlimited alert types")
}
for _, ch := range []string{"email", "discord", "telegram", "slack"} {
if !limits.ChannelAllowed(ch) {
t.Errorf("pro should allow %s", ch)
}
}
}
func TestGetTierLimits_Unknown(t *testing.T) {
t.Parallel()
limits := GetTierLimits(SubscriptionTier("unknown"))
if limits.MaxAddresses != 1 {
t.Errorf("unknown tier should default to free limits, got MaxAddresses=%d", limits.MaxAddresses)
}
}
func TestChannelAllowed_NotInList(t *testing.T) {
t.Parallel()
limits := GetTierLimits(TierFree)
if limits.ChannelAllowed("discord") {
t.Error("free should not allow discord")
}
if limits.ChannelAllowed("nonexistent") {
t.Error("nonexistent channel should not be allowed")
}
}

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())
@@ -106,7 +145,7 @@ func (h *StripeHandler) VerifyCheckoutSession(w http.ResponseWriter, r *http.Req
return return
} }
if s.ClientReferenceID != userID { if s.ClientReferenceID != "" && s.ClientReferenceID != userID {
writeError(w, http.StatusForbidden, "FORBIDDEN", "Session does not belong to this user") writeError(w, http.StatusForbidden, "FORBIDDEN", "Session does not belong to this user")
return return
} }
@@ -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,92 @@ 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",
})
}
// CreateOnboardingCheckout creates a Stripe Checkout session for a user who
// has not yet created an account. This is a public endpoint (no auth required).
// The Firebase account is created on the frontend only after payment succeeds.
func (h *StripeHandler) CreateOnboardingCheckout(w http.ResponseWriter, r *http.Request) {
var body struct {
Email string `json:"email"`
Tier string `json:"tier"`
}
if err := json.NewDecoder(r.Body).Decode(&body); err != nil {
writeError(w, http.StatusBadRequest, "BAD_REQUEST", "Invalid request body")
return
}
if body.Email == "" {
writeError(w, http.StatusBadRequest, "VALIDATION_ERROR", "Email is required")
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
}
params := &stripe.CheckoutSessionParams{
Mode: stripe.String(string(stripe.CheckoutSessionModeSubscription)),
LineItems: []*stripe.CheckoutSessionLineItemParams{
{
Price: stripe.String(priceID),
Quantity: stripe.Int64(1),
},
},
SuccessURL: stripe.String(h.cfg.FrontendURL + "/subscribe?payment=success&session_id={CHECKOUT_SESSION_ID}"),
CancelURL: stripe.String(h.cfg.FrontendURL + "/subscribe?payment=cancelled"),
CustomerEmail: stripe.String(body.Email),
}
params.AddMetadata("tier", string(tier))
params.AddMetadata("onboarding", "true")
s, err := checkoutsession.New(params)
if err != nil {
log.Printf("Failed to create onboarding checkout session: %v", err)
writeError(w, http.StatusInternalServerError, "STRIPE_ERROR", "Failed to create checkout session")
return
}
writeJSON(w, http.StatusOK, map[string]string{"url": s.URL})
} }
// CreatePortalSession creates a Stripe Billing Portal session so the user can // CreatePortalSession creates a Stripe Billing Portal session so the user can
@@ -173,7 +296,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 +339,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 +360,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,11 +384,26 @@ 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 {
log.Printf("Failed to update subscription status: %v", err) // Detect tier from the subscription's current price so that
// upgrades/downgrades via the Stripe portal are reflected.
tier := domain.TierPremium
if sub.Items != nil {
for _, item := range sub.Items.Data {
if item.Price != nil {
if t := h.cfg.TierForPriceID(item.Price.ID); t != "" {
tier = domain.SubscriptionTier(t)
break
}
}
}
} }
log.Printf("Subscription %s updated to %s for customer %s", sub.ID, status, customerID) if err := h.users.ActivateSubscription(r.Context(), customerID, sub.ID, status, tier); err != nil {
log.Printf("Failed to update subscription: %v", err)
}
log.Printf("Subscription %s updated to %s (tier %s) for customer %s", sub.ID, status, tier, customerID)
} }
func (h *StripeHandler) handleSubscriptionDeleted(r *http.Request, event stripe.Event) { func (h *StripeHandler) handleSubscriptionDeleted(r *http.Request, event stripe.Event) {

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

@@ -0,0 +1,60 @@
package middleware
import (
"context"
"testing"
)
func TestGetUserID_Empty(t *testing.T) {
t.Parallel()
ctx := context.Background()
if id := GetUserID(ctx); id != "" {
t.Errorf("expected empty user ID, got %q", id)
}
}
func TestGetUserID_Set(t *testing.T) {
t.Parallel()
ctx := context.WithValue(context.Background(), UserIDKey, "abc-123")
if id := GetUserID(ctx); id != "abc-123" {
t.Errorf("expected abc-123, got %q", id)
}
}
func TestGetUserEmail_Empty(t *testing.T) {
t.Parallel()
ctx := context.Background()
if email := GetUserEmail(ctx); email != "" {
t.Errorf("expected empty email, got %q", email)
}
}
func TestGetUserEmail_Set(t *testing.T) {
t.Parallel()
ctx := context.WithValue(context.Background(), UserEmailKey, "test@example.com")
if email := GetUserEmail(ctx); email != "test@example.com" {
t.Errorf("expected test@example.com, got %q", email)
}
}
func TestGetUserTier_Default(t *testing.T) {
t.Parallel()
ctx := context.Background()
if tier := GetUserTier(ctx); tier != "free" {
t.Errorf("expected default tier 'free', got %q", tier)
}
}
func TestGetUserTier_Set(t *testing.T) {
t.Parallel()
ctx := context.WithValue(context.Background(), UserTierKey, "pro")
if tier := GetUserTier(ctx); tier != "pro" {
t.Errorf("expected 'pro', got %q", tier)
}
}

View File

@@ -55,12 +55,14 @@ func (m *AddressModel) ListByUser(ctx context.Context, userID string) ([]domain.
return addresses, rows.Err() return addresses, rows.Err()
} }
// ListAll returns all addresses system-wide (used by the poller). // ListAll returns addresses for users with an active subscription (used by the poller).
func (m *AddressModel) ListAll(ctx context.Context) ([]domain.Address, error) { func (m *AddressModel) ListAll(ctx context.Context) ([]domain.Address, error) {
rows, err := m.pool.Query(ctx, rows, err := m.pool.Query(ctx,
`SELECT id, user_id, address, label, created_at `SELECT a.id, a.user_id, a.address, a.label, a.created_at
FROM addresses FROM addresses a
ORDER BY created_at DESC`, JOIN users u ON u.id = a.user_id
WHERE u.subscription_status IN ('active', 'trialing')
ORDER BY a.created_at DESC`,
) )
if err != nil { if err != nil {
return nil, err return nil, err
@@ -125,6 +127,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

@@ -77,10 +77,12 @@ func (m *NotificationConfigModel) Remove(ctx context.Context, userID string) (bo
func (m *NotificationConfigModel) ListEnabled(ctx context.Context) ([]domain.NotificationConfig, error) { func (m *NotificationConfigModel) ListEnabled(ctx context.Context) ([]domain.NotificationConfig, error) {
rows, err := m.pool.Query(ctx, rows, err := m.pool.Query(ctx,
`SELECT user_id, discord_webhook_url, telegram_chat_id, telegram_bot_token, `SELECT nc.user_id, nc.discord_webhook_url, nc.telegram_chat_id, nc.telegram_bot_token,
email, slack_webhook_url nc.email, nc.slack_webhook_url
FROM user_notification_configs FROM user_notification_configs nc
WHERE notification_enabled = TRUE`, JOIN users u ON u.id = nc.user_id
WHERE nc.notification_enabled = TRUE
AND u.subscription_status IN ('active', 'trialing')`,
) )
if err != nil { if err != nil {
return nil, err return nil, err

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

@@ -28,9 +28,8 @@ type EvaluatorService struct {
alertRules *models.AlertRuleModel alertRules *models.AlertRuleModel
alertEvents *models.AlertEventModel alertEvents *models.AlertEventModel
addresses *models.AddressModel addresses *models.AddressModel
users *models.UserModel
notifConfigs *models.NotificationConfigModel notifConfigs *models.NotificationConfigModel
resendAPIKey string
emailFrom string
notifSem *semaphore.Weighted notifSem *semaphore.Weighted
notifWg sync.WaitGroup notifWg sync.WaitGroup
} }
@@ -40,18 +39,16 @@ func NewEvaluatorService(
alertRules *models.AlertRuleModel, alertRules *models.AlertRuleModel,
alertEvents *models.AlertEventModel, alertEvents *models.AlertEventModel,
addresses *models.AddressModel, addresses *models.AddressModel,
users *models.UserModel,
notifConfigs *models.NotificationConfigModel, notifConfigs *models.NotificationConfigModel,
resendAPIKey string,
emailFrom string,
) *EvaluatorService { ) *EvaluatorService {
return &EvaluatorService{ return &EvaluatorService{
eth: eth, eth: eth,
alertRules: alertRules, alertRules: alertRules,
alertEvents: alertEvents, alertEvents: alertEvents,
addresses: addresses, addresses: addresses,
users: users,
notifConfigs: notifConfigs, notifConfigs: notifConfigs,
resendAPIKey: resendAPIKey,
emailFrom: emailFrom,
notifSem: semaphore.NewWeighted(maxConcurrentNotifications), notifSem: semaphore.NewWeighted(maxConcurrentNotifications),
} }
} }
@@ -254,14 +251,16 @@ func (s *EvaluatorService) WaitForNotifications() {
s.notifWg.Wait() s.notifWg.Wait()
} }
func (s *EvaluatorService) buildNotifiers(cfg *domain.NotificationConfig) []notifications.Notifier { func (s *EvaluatorService) buildNotifiers(cfg *domain.NotificationConfig, limits domain.TierLimits) []notifications.Notifier {
var notifiers []notifications.Notifier var notifiers []notifications.Notifier
if cfg.DiscordWebhookURL != nil && *cfg.DiscordWebhookURL != "" { if limits.ChannelAllowed("discord") &&
cfg.DiscordWebhookURL != nil && *cfg.DiscordWebhookURL != "" {
notifiers = append(notifiers, &notifications.DiscordNotifier{WebhookURL: *cfg.DiscordWebhookURL}) notifiers = append(notifiers, &notifications.DiscordNotifier{WebhookURL: *cfg.DiscordWebhookURL})
} }
if cfg.TelegramBotToken != nil && *cfg.TelegramBotToken != "" && if limits.ChannelAllowed("telegram") &&
cfg.TelegramBotToken != nil && *cfg.TelegramBotToken != "" &&
cfg.TelegramChatID != nil && *cfg.TelegramChatID != "" { cfg.TelegramChatID != nil && *cfg.TelegramChatID != "" {
notifiers = append(notifiers, &notifications.TelegramNotifier{ notifiers = append(notifiers, &notifications.TelegramNotifier{
BotToken: *cfg.TelegramBotToken, BotToken: *cfg.TelegramBotToken,
@@ -269,17 +268,12 @@ func (s *EvaluatorService) buildNotifiers(cfg *domain.NotificationConfig) []noti
}) })
} }
if cfg.SlackWebhookURL != nil && *cfg.SlackWebhookURL != "" { if limits.ChannelAllowed("slack") &&
cfg.SlackWebhookURL != nil && *cfg.SlackWebhookURL != "" {
notifiers = append(notifiers, &notifications.SlackNotifier{WebhookURL: *cfg.SlackWebhookURL}) notifiers = append(notifiers, &notifications.SlackNotifier{WebhookURL: *cfg.SlackWebhookURL})
} }
if cfg.Email != nil && *cfg.Email != "" { // Email is delivered exclusively via the daily digest service, never real-time.
notifiers = append(notifiers, &notifications.EmailNotifier{
APIKey: s.resendAPIKey,
From: s.emailFrom,
To: *cfg.Email,
})
}
return notifiers return notifiers
} }
@@ -325,6 +319,13 @@ func (s *EvaluatorService) sendNotification(ctx context.Context, userID, message
return return
} }
user, err := s.users.GetByID(ctx, userID)
if err != nil || user == nil {
log.Printf("Failed to load user %s for tier check in notification: %v", userID, err)
return
}
limits := domain.GetTierLimits(user.SubscriptionTier)
meta := notifications.AlertMetadata{ meta := notifications.AlertMetadata{
TxHash: obs.Hash, TxHash: obs.Hash,
AddressLabel: addressLabel, AddressLabel: addressLabel,
@@ -332,7 +333,7 @@ func (s *EvaluatorService) sendNotification(ctx context.Context, userID, message
Address: address, Address: address,
} }
for _, n := range s.buildNotifiers(notifConfig) { for _, n := range s.buildNotifiers(notifConfig, limits) {
if err := sendWithRetry(ctx, n, message, meta); err != nil { if err := sendWithRetry(ctx, n, message, meta); err != nil {
log.Printf("Notification channel failed for user %s after retries: %v", userID, err) log.Printf("Notification channel failed for user %s after retries: %v", userID, err)
} else { } else {

Binary file not shown.

View File

@@ -10,7 +10,7 @@ import AlertHistory from "./pages/alertHistory/AlertHistory";
import Account from "./pages/user_account/Account"; import Account from "./pages/user_account/Account";
export default function App() { export default function App() {
const { currentUser } = useAuth(); const { currentUser, isSubscribed } = useAuth();
if (!currentUser) { if (!currentUser) {
return ( return (
@@ -23,6 +23,16 @@ export default function App() {
); );
} }
if (!isSubscribed) {
return (
<Routes>
<Route path="/subscribe" element={<Subscribe />} />
<Route path="/account" element={<><Navbar /><Account /></>} />
<Route path="*" element={<Navigate to="/subscribe" />} />
</Routes>
);
}
return ( return (
<div> <div>
<Navbar /> <Navbar />

View File

@@ -1,15 +1,42 @@
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 createOnboardingCheckout(email, tier = "premium") {
const res = await fetch(`${API_BASE}/stripe/create-onboarding-checkout`, {
method: "POST",
headers: { "Content-Type": "application/json" },
body: JSON.stringify({ email, 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,96 @@
import "./TierPicker.css";
const TIERS = [
{
id: "free",
name: "Trial Monitoring",
price: "$0",
period: "",
features: [
"Monitor 1 blockchain address 24/7",
"Configure alerts to fire on trigger events",
"1 transaction alert type per trigger (email digest)",
],
disabledFeatures: [
"Real-time Discord alerts",
"Real-time Telegram alerts",
"Real-time Slack alerts",
],
},
{
id: "premium",
name: "Premium Monitoring",
price: "$1.99",
period: "/month",
features: [
"Monitor 3 blockchain addresses",
"Configure alerts to fire on trigger events",
"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: "Professional Monitoring",
price: "$11.99",
period: "/month",
features: [
"Monitor unlimited blockchain addresses",
"Configure alerts to fire on trigger events",
"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 = "/subscribe?upgrade=true" }) {
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,30 @@ 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 [subscriptionStatus, setSubscriptionStatus] = useState("none");
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);
setSubscriptionStatus(data.subscription_status || "none");
} 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 +74,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 +89,13 @@ 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);
setSubscriptionStatus("none");
await signOut(auth); await signOut(auth);
} catch (err) { } catch (err) {
setError(err.message); setError(err.message);
@@ -84,18 +103,19 @@ 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 isSubscribed =
subscriptionStatus === "active" || subscriptionStatus === "trialing";
const value = { const value = {
currentUser, currentUser,
@@ -104,6 +124,12 @@ export function AuthProvider({ children }) {
logout, logout,
error, error,
loading, loading,
userTier,
tierLimits,
addressCount,
subscriptionStatus,
isSubscribed,
refreshAccount,
}; };
return ( return (

View File

@@ -397,6 +397,14 @@ button {
} }
} }
/* Tier-locked: greyed out, non-interactive overlay for features above the user's plan */
.tier-locked {
opacity: 0.45;
pointer-events: none;
filter: grayscale(40%);
user-select: none;
}
/* Mobile */ /* Mobile */
@media (max-width: 480px) { @media (max-width: 480px) {
html { html {

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,7 +89,13 @@ 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 && (
<UpgradeBanner
message={`Your plan allows ${maxAddresses} address${maxAddresses !== 1 ? "es" : ""}. Upgrade to track more.`}
/>
)}
<div className={`mb-xl${atLimit ? " tier-locked" : ""}`}>
<AddressForm onSubmit={handleAddressSubmit} /> <AddressForm onSubmit={handleAddressSubmit} />
</div> </div>

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,7 +376,13 @@ export default function Alerts() {
</div> </div>
</div> </div>
<div className="mb-xl"> {atAlertLimit && userTier !== "pro" && (
<UpgradeBanner
message={`Your ${userTier} plan allows ${maxAlertTypes} alert type${maxAlertTypes !== 1 ? "s" : ""} per address. Upgrade for more.`}
/>
)}
<div className={`mb-xl${atAlertLimit ? " tier-locked" : ""}`}>
<h3>Create New Alert</h3> <h3>Create New Alert</h3>
<AlertForm onSubmit={handleAlertSubmit} /> <AlertForm onSubmit={handleAlertSubmit} />
</div> </div>
@@ -540,19 +565,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 +627,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 +650,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,13 +7,21 @@ import {
updateNotificationConfig, updateNotificationConfig,
testNotificationChannels, testNotificationChannels,
} from "../../api/notificationConfig"; } from "../../api/notificationConfig";
import { createCheckoutSession, getSubscriptionStatus, verifyCheckoutSession } from "../../api/stripe"; import { createOnboardingCheckout, 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 = [
"Create Account", "Create Account",
"Choose Plan",
"Add Wallet", "Add Wallet",
"Alert Rules", "Alert Rules",
"Notifications", "Notifications",
@@ -25,14 +33,17 @@ export default function Subscribe() {
const navigate = useNavigate(); const navigate = useNavigate();
const [searchParams, setSearchParams] = useSearchParams(); const [searchParams, setSearchParams] = useSearchParams();
const [step, setStep] = useState(1); const hasPaymentReturn = searchParams.get("payment") === "success";
const [loading, setLoading] = useState(false); const isUpgrade = searchParams.get("upgrade") === "true";
const [step, setStep] = useState(hasPaymentReturn || currentUser ? 2 : 1);
const [loading, setLoading] = useState(hasPaymentReturn);
const [error, setError] = useState(""); const [error, setError] = useState("");
const [skipWarning, setSkipWarning] = useState(""); const [skipWarning, setSkipWarning] = useState("");
const [testResults, setTestResults] = useState(null); const [testResults, setTestResults] = useState(null);
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,8 +67,10 @@ 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 || isUpgrade) return;
getAddresses() getAddresses()
.then((addresses) => { .then((addresses) => {
if (addresses.length > 0) { if (addresses.length > 0) {
@@ -65,35 +78,64 @@ export default function Subscribe() {
} }
}) })
.catch(() => { }); .catch(() => { });
}, [currentUser, navigate]); }, [currentUser, navigate, isUpgrade]);
// Handle Stripe redirect back from checkout
useEffect(() => { useEffect(() => {
if (!currentUser) return;
const payment = searchParams.get("payment"); const payment = searchParams.get("payment");
const sessionId = searchParams.get("session_id"); const sessionId = searchParams.get("session_id");
if (payment === "success" && sessionId) {
if (payment === "cancelled") {
setSearchParams({}, { replace: true }); setSearchParams({}, { replace: true });
setLoading(true); setStep(2);
verifyCheckoutSession(sessionId)
.then(() => {
setStep(2);
})
.catch((err) => {
setError("Payment verification failed: " + err.message);
setStep(1);
})
.finally(() => setLoading(false));
} else if (payment === "cancelled") {
setSearchParams({}, { replace: true });
setStep(1);
setError("Payment was cancelled. Please try again."); setError("Payment was cancelled. Please try again.");
return;
} }
}, [currentUser, searchParams, setSearchParams]);
if (payment !== "success" || !sessionId) return;
// Phase 1: no account yet — create it. This triggers an auth state change
// which remounts the component (App.jsx swaps route trees). The URL params
// are preserved so Phase 2 runs on the next mount.
if (!currentUser) {
const savedEmail = sessionStorage.getItem("kp_onboard_email");
const savedPw = sessionStorage.getItem("kp_onboard_pw");
if (!savedEmail || !savedPw) {
setSearchParams({}, { replace: true });
setError("Session expired. Please start the signup process again.");
setStep(1);
setLoading(false);
return;
}
setLoading(true);
signup(savedEmail, savedPw).catch((err) => {
setSearchParams({}, { replace: true });
setError("Account creation failed: " + err.message);
setStep(1);
setLoading(false);
});
return;
}
// Phase 2: authenticated — verify the checkout and activate subscription.
setSearchParams({}, { replace: true });
setLoading(true);
sessionStorage.removeItem("kp_onboard_email");
sessionStorage.removeItem("kp_onboard_pw");
verifyCheckoutSession(sessionId)
.then(() => {
setStep(3);
})
.catch((err) => {
setError("Payment verification failed: " + err.message);
setStep(2);
})
.finally(() => setLoading(false));
}, [currentUser, searchParams, setSearchParams, signup]);
// ── Step handlers ───────────────────────────────────────────────────────── // ── Step handlers ─────────────────────────────────────────────────────────
async function handleStep1() { function handleStep1() {
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");
@@ -107,17 +149,30 @@ export default function Subscribe() {
setError("Password must be at least 6 characters"); setError("Password must be at least 6 characters");
return; return;
} }
setStep(2);
}
async function handleStep2() {
setError("");
if (!data.selectedTier) {
setError("Please select a plan to continue");
return;
}
try { try {
setLoading(true); setLoading(true);
if (!currentUser) {
await signup(data.email, data.password); if (data.selectedTier === "free") {
} if (!currentUser) {
const status = await getSubscriptionStatus(); await signup(data.email, data.password);
if (status.subscription_status === "active" || status.subscription_status === "trialing") { }
setStep(2); await activateFreeTier();
setStep(3);
return; return;
} }
const { url } = await createCheckoutSession();
sessionStorage.setItem("kp_onboard_email", data.email);
sessionStorage.setItem("kp_onboard_pw", data.password);
const { url } = await createOnboardingCheckout(data.email, 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") {
@@ -127,13 +182,14 @@ export default function Subscribe() {
} else if (err.code === "auth/weak-password") { } else if (err.code === "auth/weak-password") {
setError("Password is too weak"); setError("Password is too weak");
} else { } else {
setError("Failed to create account: " + err.message); setError("Failed to process plan selection: " + err.message);
} }
} finally {
setLoading(false); setLoading(false);
} }
} }
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 +206,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 +214,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 +235,7 @@ export default function Subscribe() {
} }
if (rules.length === 0) { if (rules.length === 0) {
setStep(4); setStep(5);
return; return;
} }
@@ -191,7 +247,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 +255,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 +272,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 +293,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 +355,7 @@ export default function Subscribe() {
// ── Step content ────────────────────────────────────────────────────────── // ── Step content ──────────────────────────────────────────────────────────
function Step1() { function StepCreateAccount() {
return ( return (
<> <>
<h2 className="mb-lg">Create your account</h2> <h2 className="mb-lg">Create your account</h2>
@@ -317,7 +388,22 @@ export default function Subscribe() {
); );
} }
function Step2() { function StepChoosePlan() {
return (
<>
<h2 className="mb-sm">Choose your monitoring 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 StepAddWallet() {
return ( return (
<> <>
<h2 className="mb-sm">Add a wallet address</h2> <h2 className="mb-sm">Add a wallet address</h2>
@@ -343,28 +429,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 +477,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 +493,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 +515,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 +542,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 +567,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 +589,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 +658,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 +679,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 +706,16 @@ export default function Subscribe() {
} }
const nextLabel = step === 1 const nextLabel = step === 1
? "Create Account & Subscribe" ? "Create Account"
: step === 4 : step === 2
? "Finish" ? !data.selectedTier
: "Next →"; ? "Continue"
: data.selectedTier === "free"
? "Start Free Trial"
: "Subscribe & Continue"
: step === 5
? "Finish"
: "Next →";
return ( return (
<div className="subscribe__footer"> <div className="subscribe__footer">
@@ -619,7 +743,7 @@ export default function Subscribe() {
)} )}
<Button <Button
onClick={handleNext} onClick={handleNext}
disabled={loading} disabled={loading || (step === 2 && !data.selectedTier)}
className="text-bold" className="text-bold"
> >
{loading ? "Please wait..." : nextLabel} {loading ? "Please wait..." : nextLabel}
@@ -632,16 +756,17 @@ export default function Subscribe() {
// ── Render ──────────────────────────────────────────────────────────────── // ── Render ────────────────────────────────────────────────────────────────
const stepContent = { const stepContent = {
1: Step1(), 1: StepCreateAccount(),
2: Step2(), 2: StepChoosePlan(),
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 === 2 ? " subscribe__container--wide" : ""}`}>
<h1 className="subscribe__title">Koin Ping</h1> <h1 className="subscribe__title">Koin Ping</h1>
{ProgressBar()} {ProgressBar()}

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>

BIN
memberships.pdf Normal file

Binary file not shown.