add subscription tiers

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

View File

@@ -6,29 +6,40 @@ import (
"time"
"github.com/kjannette/koin-ping/backend/internal/config"
"github.com/kjannette/koin-ping/backend/internal/domain"
"github.com/kjannette/koin-ping/backend/internal/middleware"
"github.com/kjannette/koin-ping/backend/internal/models"
)
type AccountHandler struct {
users *models.UserModel
cfg *config.Config
users *models.UserModel
addresses *models.AddressModel
cfg *config.Config
}
func NewAccountHandler(users *models.UserModel, cfg *config.Config) *AccountHandler {
return &AccountHandler{users: users, cfg: cfg}
func NewAccountHandler(users *models.UserModel, addresses *models.AddressModel, cfg *config.Config) *AccountHandler {
return &AccountHandler{users: users, addresses: addresses, cfg: cfg}
}
type accountResponse struct {
UserID string `json:"user_id"`
Email string `json:"email"`
UserName string `json:"user_name"`
SubscriptionStatus string `json:"subscription_status"`
SubscriptionPlan string `json:"subscription_plan"`
MemberSince *string `json:"member_since,omitempty"`
NextBillingDate *string `json:"next_billing_date,omitempty"`
CancelAtPeriodEnd bool `json:"cancel_at_period_end"`
PeriodEndDate *string `json:"period_end_date,omitempty"`
UserID string `json:"user_id"`
Email string `json:"email"`
UserName string `json:"user_name"`
SubscriptionStatus string `json:"subscription_status"`
SubscriptionTier string `json:"subscription_tier"`
SubscriptionPlan string `json:"subscription_plan"`
TierLimits domain.TierLimits `json:"tier_limits"`
AddressCount int `json:"address_count"`
MemberSince *string `json:"member_since,omitempty"`
NextBillingDate *string `json:"next_billing_date,omitempty"`
CancelAtPeriodEnd bool `json:"cancel_at_period_end"`
PeriodEndDate *string `json:"period_end_date,omitempty"`
}
var tierPlanLabels = map[domain.SubscriptionTier]string{ //nolint:gochecknoglobals
domain.TierFree: "Free Trial",
domain.TierPremium: "Premium / $1.99 mo",
domain.TierPro: "Pro / $11.99 mo",
}
func (h *AccountHandler) GetAccount(w http.ResponseWriter, r *http.Request) {
@@ -42,12 +53,26 @@ func (h *AccountHandler) GetAccount(w http.ResponseWriter, r *http.Request) {
return
}
addrCount, err := h.addresses.CountByUser(r.Context(), userID)
if err != nil {
log.Printf("Account: failed to count addresses for %s: %v", userID, err)
addrCount = 0
}
planLabel := tierPlanLabels[user.SubscriptionTier]
if planLabel == "" {
planLabel = "Free Trial"
}
resp := accountResponse{
UserID: user.ID,
Email: email,
UserName: email,
SubscriptionStatus: user.SubscriptionStatus,
SubscriptionPlan: "monthly/$1.99",
SubscriptionTier: string(user.SubscriptionTier),
SubscriptionPlan: planLabel,
TierLimits: domain.GetTierLimits(user.SubscriptionTier),
AddressCount: addrCount,
}
if user.SubscriptionCreatedAt != nil {
@@ -55,9 +80,6 @@ func (h *AccountHandler) GetAccount(w http.ResponseWriter, r *http.Request) {
resp.MemberSince = &t
}
// MOCKED: NextBillingDate, CancelAtPeriodEnd, PeriodEndDate.
// TODO: stripe-go v82 removed Subscription.CurrentPeriodEnd / CancelAtPeriodEnd.
// Research SubscriptionItem.CurrentPeriodEnd or use Stripe REST API directly.
resp.NextBillingDate = nil
resp.CancelAtPeriodEnd = false
resp.PeriodEndDate = nil

View File

@@ -3,6 +3,7 @@ package handlers
import (
"encoding/json"
"fmt"
"log"
"net/http"
"regexp"
@@ -17,10 +18,11 @@ var ethAddressRe = regexp.MustCompile(`^0x[a-fA-F0-9]{40}$`)
type AddressHandler struct {
addresses *models.AddressModel
users *models.UserModel
}
func NewAddressHandler(addresses *models.AddressModel) *AddressHandler {
return &AddressHandler{addresses: addresses}
func NewAddressHandler(addresses *models.AddressModel, users *models.UserModel) *AddressHandler {
return &AddressHandler{addresses: addresses, users: users}
}
func (h *AddressHandler) Create(w http.ResponseWriter, r *http.Request) {
@@ -49,6 +51,28 @@ func (h *AddressHandler) Create(w http.ResponseWriter, r *http.Request) {
return
}
user, err := h.users.GetByID(r.Context(), userID)
if err != nil || user == nil {
log.Printf("Failed to get user %s for tier check: %v", userID, err)
writeError(w, http.StatusInternalServerError, "INTERNAL_ERROR", "Failed to verify account")
return
}
limits := domain.GetTierLimits(user.SubscriptionTier)
if !limits.IsUnlimitedAddresses() {
count, err := h.addresses.CountByUser(r.Context(), userID)
if err != nil {
log.Printf("Failed to count addresses for user %s: %v", userID, err)
writeError(w, http.StatusInternalServerError, "INTERNAL_ERROR", "Failed to create address")
return
}
if count >= limits.MaxAddresses {
writeError(w, http.StatusForbidden, "TIER_LIMIT_REACHED",
fmt.Sprintf("Your %s plan allows %d address(es). Upgrade to track more.", user.SubscriptionTier, limits.MaxAddresses))
return
}
}
log.Printf("User %s creating address: %s", userID, body.Address)
addr, err := h.addresses.Create(r.Context(), userID, body.Address, body.Label)

View File

@@ -19,10 +19,11 @@ var errThresholdFormat = errors.New("unsupported threshold format")
type AlertRuleHandler struct {
alertRules *models.AlertRuleModel
addresses *models.AddressModel
users *models.UserModel
}
func NewAlertRuleHandler(alertRules *models.AlertRuleModel, addresses *models.AddressModel) *AlertRuleHandler {
return &AlertRuleHandler{alertRules: alertRules, addresses: addresses}
func NewAlertRuleHandler(alertRules *models.AlertRuleModel, addresses *models.AddressModel, users *models.UserModel) *AlertRuleHandler {
return &AlertRuleHandler{alertRules: alertRules, addresses: addresses, users: users}
}
func (h *AlertRuleHandler) Create(w http.ResponseWriter, r *http.Request) {
@@ -132,6 +133,28 @@ func (h *AlertRuleHandler) Create(w http.ResponseWriter, r *http.Request) {
return
}
user, userErr := h.users.GetByID(r.Context(), userID)
if userErr != nil || user == nil {
log.Printf("Failed to get user %s for tier check: %v", userID, userErr)
writeError(w, http.StatusInternalServerError, "INTERNAL_ERROR", "Failed to verify account")
return
}
limits := domain.GetTierLimits(user.SubscriptionTier)
if !limits.IsUnlimitedAlertTypes() {
typeCount, countErr := h.alertRules.CountDistinctTypesByAddress(r.Context(), addressID)
if countErr != nil {
log.Printf("Failed to count alert types for address %d: %v", addressID, countErr)
writeError(w, http.StatusInternalServerError, "INTERNAL_ERROR", "Failed to create alert rule")
return
}
if typeCount >= limits.MaxAlertTypes {
writeError(w, http.StatusForbidden, "TIER_LIMIT_REACHED",
fmt.Sprintf("Your %s plan allows %d alert type(s) per address. Upgrade for more.", user.SubscriptionTier, limits.MaxAlertTypes))
return
}
}
newAlert, err := h.alertRules.Create(r.Context(), addressID, alertType, threshold, minimum, maximum)
if err != nil {
log.Printf("Error creating alert rule: %v", err)

View File

@@ -18,11 +18,12 @@ var emailRe = regexp.MustCompile(`^[^\s@]+@[^\s@]+\.[^\s@]+$`)
type NotificationConfigHandler struct {
configs *models.NotificationConfigModel
users *models.UserModel
cfg *config.Config
}
func NewNotificationConfigHandler(configs *models.NotificationConfigModel, cfg *config.Config) *NotificationConfigHandler {
return &NotificationConfigHandler{configs: configs, cfg: cfg}
func NewNotificationConfigHandler(configs *models.NotificationConfigModel, users *models.UserModel, cfg *config.Config) *NotificationConfigHandler {
return &NotificationConfigHandler{configs: configs, users: users, cfg: cfg}
}
func (h *NotificationConfigHandler) GetConfig(w http.ResponseWriter, r *http.Request) {
@@ -95,6 +96,26 @@ func (h *NotificationConfigHandler) UpdateConfig(w http.ResponseWriter, r *http.
return
}
user, userErr := h.users.GetByID(r.Context(), userID)
if userErr != nil || user == nil {
log.Printf("Failed to get user %s for tier check: %v", userID, userErr)
writeError(w, http.StatusInternalServerError, "INTERNAL_ERROR", "Failed to verify account")
return
}
limits := domain.GetTierLimits(user.SubscriptionTier)
if !limits.ChannelAllowed("discord") {
body.DiscordWebhookURL = nil
}
if !limits.ChannelAllowed("telegram") {
body.TelegramBotToken = nil
body.TelegramChatID = nil
}
if !limits.ChannelAllowed("slack") {
body.SlackWebhookURL = nil
}
enabled := true
if body.NotificationEnabled != nil {
enabled = *body.NotificationEnabled

View File

@@ -2,6 +2,7 @@ package handlers
import (
"encoding/json"
"fmt"
"io"
"log"
"net/http"
@@ -12,6 +13,7 @@ import (
"github.com/stripe/stripe-go/v82/webhook"
"github.com/kjannette/koin-ping/backend/internal/config"
"github.com/kjannette/koin-ping/backend/internal/domain"
"github.com/kjannette/koin-ping/backend/internal/middleware"
"github.com/kjannette/koin-ping/backend/internal/models"
)
@@ -28,10 +30,45 @@ func NewStripeHandler(users *models.UserModel, cfg *config.Config) *StripeHandle
return &StripeHandler{users: users, cfg: cfg}
}
// CreateCheckoutSession creates a Stripe Checkout session for the monthly subscription.
func (h *StripeHandler) priceIDForTier(tier domain.SubscriptionTier) (string, error) {
switch tier {
case domain.TierPremium:
return h.cfg.StripePriceIDPremium, nil
case domain.TierPro:
return h.cfg.StripePriceIDPro, nil
default:
return "", fmt.Errorf("no Stripe price for tier %q", tier) //nolint:err113
}
}
// CreateCheckoutSession creates a Stripe Checkout session for the selected tier.
func (h *StripeHandler) CreateCheckoutSession(w http.ResponseWriter, r *http.Request) {
userID := middleware.GetUserID(r.Context())
var body struct {
Tier string `json:"tier"`
}
if err := json.NewDecoder(r.Body).Decode(&body); err != nil {
writeError(w, http.StatusBadRequest, "BAD_REQUEST", "Invalid request body")
return
}
if body.Tier == "" {
body.Tier = "premium"
}
tier := domain.SubscriptionTier(body.Tier)
if tier != domain.TierPremium && tier != domain.TierPro {
writeError(w, http.StatusBadRequest, "VALIDATION_ERROR", "Tier must be 'premium' or 'pro'")
return
}
priceID, err := h.priceIDForTier(tier)
if err != nil {
writeError(w, http.StatusBadRequest, "VALIDATION_ERROR", err.Error())
return
}
user, err := h.users.GetByID(r.Context(), userID)
if err != nil || user == nil {
log.Printf("Failed to get user %s: %v", userID, err)
@@ -43,7 +80,7 @@ func (h *StripeHandler) CreateCheckoutSession(w http.ResponseWriter, r *http.Req
Mode: stripe.String(string(stripe.CheckoutSessionModeSubscription)),
LineItems: []*stripe.CheckoutSessionLineItemParams{
{
Price: stripe.String(h.cfg.StripePriceID),
Price: stripe.String(priceID),
Quantity: stripe.Int64(1),
},
},
@@ -53,6 +90,8 @@ func (h *StripeHandler) CreateCheckoutSession(w http.ResponseWriter, r *http.Req
CustomerEmail: stripe.String(user.Email),
}
params.AddMetadata("tier", string(tier))
if user.StripeCustomerID != nil && *user.StripeCustomerID != "" {
params.Customer = user.StripeCustomerID
params.CustomerEmail = nil
@@ -80,14 +119,14 @@ func (h *StripeHandler) GetSubscriptionStatus(w http.ResponseWriter, r *http.Req
}
writeJSON(w, http.StatusOK, map[string]any{
"subscription_status": user.SubscriptionStatus,
"subscription_status": user.SubscriptionStatus,
"subscription_tier": user.SubscriptionTier,
"subscription_created_at": user.SubscriptionCreatedAt,
})
}
// VerifyCheckoutSession retrieves a completed checkout session from Stripe,
// confirms payment, and activates the user's subscription in the database.
// This is the primary activation path; webhooks serve as a backup.
func (h *StripeHandler) VerifyCheckoutSession(w http.ResponseWriter, r *http.Request) {
userID := middleware.GetUserID(r.Context())
@@ -116,6 +155,11 @@ func (h *StripeHandler) VerifyCheckoutSession(w http.ResponseWriter, r *http.Req
return
}
tier := domain.TierPremium
if t, ok := s.Metadata["tier"]; ok && domain.IsValidTier(t) {
tier = domain.SubscriptionTier(t)
}
customerID := ""
if s.Customer != nil {
customerID = s.Customer.ID
@@ -131,13 +175,33 @@ func (h *StripeHandler) VerifyCheckoutSession(w http.ResponseWriter, r *http.Req
}
}
if subscriptionID != "" && customerID != "" {
if err := h.users.ActivateSubscription(r.Context(), customerID, subscriptionID, "active"); err != nil {
if err := h.users.ActivateSubscription(r.Context(), customerID, subscriptionID, "active", tier); err != nil {
log.Printf("VerifyCheckout: failed to activate subscription: %v", err)
}
}
log.Printf("Checkout verified for user %s, customer %s, subscription %s", userID, customerID, subscriptionID)
writeJSON(w, http.StatusOK, map[string]string{"subscription_status": "active"})
log.Printf("Checkout verified for user %s, customer %s, subscription %s, tier %s", userID, customerID, subscriptionID, tier)
writeJSON(w, http.StatusOK, map[string]string{
"subscription_status": "active",
"subscription_tier": string(tier),
})
}
// ActivateFreeTier sets the user to the free tier without Stripe involvement.
func (h *StripeHandler) ActivateFreeTier(w http.ResponseWriter, r *http.Request) {
userID := middleware.GetUserID(r.Context())
if err := h.users.ActivateFreeTier(r.Context(), userID); err != nil {
log.Printf("ActivateFreeTier: failed for user %s: %v", userID, err)
writeError(w, http.StatusInternalServerError, "INTERNAL_ERROR", "Failed to activate free tier")
return
}
log.Printf("Free tier activated for user %s", userID)
writeJSON(w, http.StatusOK, map[string]string{
"subscription_status": "active",
"subscription_tier": "free",
})
}
// CreatePortalSession creates a Stripe Billing Portal session so the user can
@@ -173,7 +237,6 @@ func (h *StripeHandler) CreatePortalSession(w http.ResponseWriter, r *http.Reque
}
// HandleWebhook processes incoming Stripe webhook events.
// This endpoint must NOT require authentication (Stripe calls it directly).
func (h *StripeHandler) HandleWebhook(w http.ResponseWriter, r *http.Request) {
payload, err := io.ReadAll(io.LimitReader(r.Body, webhookMaxBodyBytes))
if err != nil {
@@ -217,6 +280,11 @@ func (h *StripeHandler) handleCheckoutCompleted(r *http.Request, event stripe.Ev
return
}
tier := domain.TierPremium
if t, ok := session.Metadata["tier"]; ok && domain.IsValidTier(t) {
tier = domain.SubscriptionTier(t)
}
customerID := ""
if session.Customer != nil {
customerID = session.Customer.ID
@@ -233,12 +301,12 @@ func (h *StripeHandler) handleCheckoutCompleted(r *http.Request, event stripe.Ev
}
if subscriptionID != "" && customerID != "" {
if err := h.users.ActivateSubscription(r.Context(), customerID, subscriptionID, "active"); err != nil {
if err := h.users.ActivateSubscription(r.Context(), customerID, subscriptionID, "active", tier); err != nil {
log.Printf("Failed to activate subscription: %v", err)
}
}
log.Printf("Checkout completed for user %s, customer %s, subscription %s", userID, customerID, subscriptionID)
log.Printf("Checkout completed for user %s, customer %s, subscription %s, tier %s", userID, customerID, subscriptionID, tier)
}
func (h *StripeHandler) handleSubscriptionUpdated(r *http.Request, event stripe.Event) {
@@ -257,7 +325,7 @@ func (h *StripeHandler) handleSubscriptionUpdated(r *http.Request, event stripe.
}
status := string(sub.Status)
if err := h.users.ActivateSubscription(r.Context(), customerID, sub.ID, status); err != nil {
if err := h.users.UpdateSubscriptionStatus(r.Context(), customerID, status); err != nil {
log.Printf("Failed to update subscription status: %v", err)
}