Files
Crussell/backend/main.go
T
popertotsandSisyphus e4b9003439 refactor(handlers): migrate remaining backend handlers to clock.Now() and transaction patterns
Apply clock.Now() migration, transaction wrapping, and minor refactors across admin, scheduling, today, user, auth handler, notifications, webhooks, services, portfolio, ratelimit, testutils, and main.go.

Ultraworked with [Sisyphus](https://github.com/code-yeongyu/oh-my-openagent)

Co-authored-by: Sisyphus <clio-agent@sisyphuslabs.ai>
2026-06-24 23:43:50 +01:00

450 lines
16 KiB
Go

package main
import (
"context"
"crussell/auth"
"crussell/internal/dav"
"crussell/internal/s3"
"crussell/internal/square"
"encoding/json"
"fmt"
"log"
"net/http"
"os"
"os/signal"
"strings"
"syscall"
"time"
"github.com/go-chi/chi/v5"
"github.com/go-chi/chi/v5/middleware"
"crussell/db"
"crussell/mw"
authHandlers "crussell/handlers/auth"
"crussell/handlers/admin"
"crussell/handlers/bookings"
"crussell/handlers/notifications"
"crussell/handlers/payments"
"crussell/handlers/portfolio"
"crussell/handlers/scheduling"
"crussell/handlers/services"
"crussell/handlers/today"
"crussell/handlers/user"
"crussell/handlers/webhooks"
)
func init() {
// The testing framework passes -test.* to the binary; skip JWT init
// because test packages handle it via testutils/jwt. This is more
// precise than checking GO_TESTING env var, which can leak from the
// test runner into the environment during seeding.
for _, arg := range os.Args {
if strings.HasPrefix(arg, "-test.") {
return
}
}
jwtSecret := os.Getenv("JWT_SECRET_KEY")
if jwtSecret == "" {
log.Fatal("FATAL: JWT_SECRET_KEY environment variable not set. Application cannot start.")
}
auth.InitJWT(jwtSecret)
}
func limitBody(limit int64) func(http.Handler) http.Handler {
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.Body != nil {
r.Body = http.MaxBytesReader(w, r.Body, limit)
}
next.ServeHTTP(w, r)
})
}
}
const (
defaultBodyLimit int64 = 1 * 1024 * 1024 // 1MB
uploadBodyLimit int64 = 15 * 1024 * 1024 // 15MB
portfolioBodyLimit int64 = 40 * 1024 * 1024 // 40MB (7 variants from 20MB source)
)
func initDB() {
if err := db.Connect(); err != nil {
log.Fatal("Failed to connect to DB:", err)
}
fmt.Println("Connected to DB successfully")
}
func initDav() {
if dav.Service == nil {
log.Fatal("Failed to initialize DAV service")
}
fmt.Println("DAV Service connected successfully")
}
func initS3() {
if err := s3.Connect(); err != nil {
log.Printf("WARNING: Failed to connect to S3: %v", err)
} else {
fmt.Println("S3 client initialized")
}
}
func initSquare() {
payments.SquareClient = square.NewClient()
fmt.Println("Square client initialized (dev mock)")
}
func healthCheckHandler(w http.ResponseWriter, r *http.Request) {
status := "ok"
services := map[string]string{
"backend": "ok",
"database": "ok",
"s3_storage": "ok",
"square_payments": "not_implemented",
"frontend": "unknown",
}
if db.Conn != nil {
if err := db.Conn.Ping(r.Context()); err != nil {
services["database"] = "error"
status = "degraded"
}
} else {
services["database"] = "error"
status = "degraded"
}
if s3.Client == nil {
services["s3_storage"] = "not_configured"
}
if status == "degraded" {
w.WriteHeader(http.StatusServiceUnavailable)
} else {
w.WriteHeader(http.StatusOK)
}
json.NewEncoder(w).Encode(map[string]interface{}{
"status": status,
"services": services,
})
}
func main() {
initDB()
initDav()
initS3()
initSquare()
auth.StartJTICleanup()
r := chi.NewRouter()
// --- Global Middleware ---
r.Use(middleware.RequestID)
r.Use(middleware.ClientIPFromHeader("X-Real-IP"))
r.Use(middleware.Logger)
r.Use(middleware.Recoverer)
r.Use(middleware.Timeout(15 * time.Second))
r.Use(func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("X-Content-Type-Options", "nosniff")
w.Header().Set("X-Frame-Options", "DENY")
w.Header().Set("X-XSS-Protection", "1; mode=block")
// TODO: Enable HSTS in production
w.Header().Set("Strict-Transport-Security", "max-age=31536000; includeSubDomains")
// TODO: Enable Referrer-Policy in production
w.Header().Set("Referrer-Policy", "strict-origin-when-cross-origin")
w.Header().Set("Content-Security-Policy", "default-src 'none'; frame-ancestors 'none'")
// Reflect origin (not wildcard '*') so credentialed cross-origin
// requests with Authorization: Bearer work in browsers.
origin := r.Header.Get("Origin")
if origin != "" {
w.Header().Set("Access-Control-Allow-Origin", origin)
w.Header().Set("Vary", "Origin")
}
w.Header().Set("Access-Control-Allow-Methods", "GET, POST, PUT, PATCH, DELETE, OPTIONS")
w.Header().Set("Access-Control-Allow-Headers", "Authorization, Content-Type, Idempotency-Key")
if r.Method == http.MethodOptions {
w.WriteHeader(http.StatusNoContent)
return
}
next.ServeHTTP(w, r)
})
})
// All API routes grouped under /api for clarity
r.Route("/api", func(r chi.Router) {
r.Use(mw.JsonContentType)
// Public read-only (but check auth context if present for eligibility)
r.Group(func(r chi.Router) {
r.Use(mw.RateLimit(120, time.Minute))
r.Get("/services", services.ServicesHandler)
r.Get("/services/eligible-for/{user_id}", services.ServicesEligibleForUserHandler)
})
// Registration: 10/min to prevent spam + progressive per-IP backoff
r.With(mw.ProgressiveRateLimit, mw.RateLimit(10, time.Minute), limitBody(defaultBodyLimit)).Post("/register", authHandlers.RegisterHandler)
// Login: Has its own internal rate limiting + progressive per-IP backoff
r.With(mw.ProgressiveRateLimit, mw.RateLimit(10, time.Minute), limitBody(defaultBodyLimit)).Post("/login", authHandlers.LoginHandler)
// Logout: requires valid token
r.With(mw.RequireAuth).Post("/logout", authHandlers.LogoutHandler)
// Email verification
r.With(mw.RateLimit(10, time.Minute), limitBody(defaultBodyLimit)).Post("/verify/generate", authHandlers.GenerateVerificationCodeHandler)
r.With(mw.RateLimit(20, time.Minute), limitBody(defaultBodyLimit)).Post("/verify/check", authHandlers.VerifyCodeHandler)
// Health check
r.Get("/health", healthCheckHandler)
// Public contact info
r.Get("/contact", user.GetContactInfoHandler)
// Public business info (limited, safe for non-admin users)
r.Get("/business-info", admin.GetPublicBusinessInfo)
// Portfolio
r.Route("/portfolio", func(r chi.Router) {
r.Get("/images", portfolio.ListImages)
r.Get("/tags", portfolio.ListTags)
r.With(mw.RateLimit(60, time.Minute)).Get("/filters", portfolio.ListFilters)
r.With(mw.RateLimit(120, time.Minute)).Get("/images/{id}", portfolio.GetImage)
r.Group(func(r chi.Router) {
r.Use(mw.RequireAuth)
r.Use(mw.RequireAdmin)
r.Use(mw.RateLimit(60, time.Minute))
r.With(limitBody(portfolioBodyLimit)).Post("/images", portfolio.UploadImage)
r.Delete("/images/{id}", portfolio.DeleteImage)
})
})
// Scheduling
r.Route("/scheduling", func(r chi.Router) {
r.Use(mw.RateLimit(120, time.Minute))
r.Use(mw.OptionalAuth)
r.Get("/default-hours", scheduling.GetDefaultHours)
r.Get("/exceptional-groups", scheduling.ListExceptionalGroups)
r.Get("/working-hours", scheduling.GetWorkingHours)
r.Get("/available-hours", scheduling.GetAvailableHours)
r.Group(func(r chi.Router) {
r.Use(mw.RequireAuth)
r.Use(mw.RequireAdmin)
r.Use(mw.RateLimit(60, time.Minute))
r.Put("/default-hours", scheduling.UpdateDefaultHours)
r.Post("/exceptional-groups", scheduling.CreateExceptionalGroup)
r.Delete("/exceptional-groups", scheduling.DeleteExceptionalGroup)
r.Put("/exceptional-applications", scheduling.UpdateExceptionalApplications)
})
})
// Public booking endpoints (optional auth for slot reservation and guest bookings)
r.Group(func(r chi.Router) {
r.Use(mw.RateLimit(30, time.Minute), mw.OptionalAuth)
r.Use(limitBody(defaultBodyLimit))
r.Post("/bookings/reserve", bookings.ReserveSlotHandler)
r.Post("/bookings", bookings.CreateBookingHandler)
})
// Guest user creation (public, no auth required)
r.With(mw.RateLimit(10, time.Minute), limitBody(defaultBodyLimit)).Post("/users/guest", user.CreateGuestUserHandler)
// Public email check (used by BookingFlow for proactive registered-email detection)
r.With(mw.RateLimit(60, time.Minute)).Get("/check-email", user.CheckEmailHandler)
// Authenticated users
r.Group(func(r chi.Router) {
r.Use(mw.RequireAuth)
r.Use(mw.RateLimit(120, time.Minute))
r.Use(limitBody(defaultBodyLimit))
r.Post("/refresh-token", authHandlers.RefreshTokenHandler)
r.Get("/user/profile", user.GetProfileHandler)
r.Put("/user/profile", user.UpdateProfileHandler)
r.Put("/user/change-password", user.ChangePasswordHandler)
r.Get("/user/notification-preferences", user.GetNotificationPreferencesHandler)
r.Put("/user/notification-preferences", user.UpdateNotificationPreferencesHandler)
r.Delete("/user/account", user.DeleteAccountHandler)
r.Get("/user/gdpr-export", user.GetGDPRExportHandler)
r.Get("/user/loyalty", user.GetLoyaltyHandler)
r.Get("/bookings", bookings.GetAllUserBookingsHandler)
r.Get("/bookings/{id}", bookings.GetBookingHandler)
r.Get("/bookings/{id}/calendar", bookings.GetBookingCalendarHandler)
r.Put("/bookings/{id}", bookings.EditBookingHandler)
r.Delete("/bookings/{id}", bookings.DeleteBookingHandler)
r.Post("/bookings/{id}/edit-request", bookings.RequestEditHandler)
r.Delete("/bookings/{id}/edit-request", bookings.DeleteEditRequestHandler)
r.Get("/bookings/{id}/edit-request", bookings.GetMyEditRequestHandler)
r.Get("/bookings/edit-requests", bookings.GetMyEditRequestsHandler)
// User payment routes
r.Post("/bookings/{id}/payment", payments.CreateBookingPayment)
r.Post("/bookings/{id}/apply-redemption", payments.ApplyLoyaltyRedemption)
r.Post("/bookings/{id}/payment-lock", payments.AcquirePaymentLock)
r.Delete("/bookings/{id}/payment-lock", payments.ReleasePaymentLock)
r.Get("/user/payment-methods", payments.GetUserPaymentMethods)
r.Post("/user/payment-methods", payments.CreatePaymentMethod)
r.Delete("/user/payment-methods/{id}", payments.DeletePaymentMethod)
r.Post("/bookings/{id}/tip", payments.CreateTipPayment)
r.Get("/bookings/{id}/payment-summary", payments.GetBookingPaymentSummary)
r.With(mw.RateLimit(10, time.Minute), limitBody(defaultBodyLimit)).
Get("/bookings/{id}/discount-preview", payments.GetDiscountPreviewHandler)
// User gift card routes
r.Post("/user/giftcards/redeem", payments.RedeemGiftCard)
r.Get("/user/giftcards/balance", payments.GetGiftCardBalance)
r.Post("/user/giftcards/buy", payments.BuyGiftCard)
})
r.With(mw.RequireAuth, mw.RequireVerified, limitBody(uploadBodyLimit)).Post("/user/profile-picture", user.UploadProfilePictureHandler)
// Admin-only (no rate limit - trusted users with authenticated sessions)
r.Group(func(r chi.Router) {
r.Use(mw.RequireAuth)
r.Use(mw.RequireAdmin)
r.Use(limitBody(defaultBodyLimit))
r.Route("/admin/services", func(r chi.Router) {
r.Post("/", services.CreateServiceHandler)
r.Delete("/{id}", services.DeleteServiceHandler)
r.Get("/", services.AllServicesHandler)
r.Put("/{id}/toggle", services.ToggleService)
})
r.Route("/admin/patch-tests", func(r chi.Router) {
r.Get("/", admin.GetPatchTests)
r.Post("/", admin.CreatePatchTest)
r.Put("/{id}", admin.UpdatePatchTest)
r.Delete("/{id}", admin.DeletePatchTest)
})
r.Route("/admin/custom-services", func(r chi.Router) {
r.Get("/", admin.GetCustomServices)
r.Post("/", admin.CreateCustomService)
r.Get("/{id}", admin.GetCustomService)
r.Put("/{id}", admin.UpdateCustomService)
r.Post("/{id}/promote", admin.PromoteCustomService)
r.Delete("/{id}", admin.DeleteCustomService)
})
r.Route("/admin/bookings", func(r chi.Router) {
r.Get("/", bookings.GetAllAdminBookingsHandler)
r.Post("/", bookings.AdminCreateBookingForUserHandler)
r.With(mw.RateLimit(60, time.Minute)).Get("/search", bookings.SearchAdminBookingsHandler)
r.Get("/user/{user_id}", bookings.GetAllBookingsByUserHandler)
r.Get("/{id}", bookings.GetAdminBookingHandler)
r.Put("/{id}", bookings.UpdateBookingServicesHandler)
r.Get("/{id}/overlapping", bookings.GetOverlappingBookingsHandler)
r.Get("/overlapping", bookings.GetOverlappingBookingsByTimeHandler)
r.Get("/by-date-range", bookings.GetBookingsByDateRangeHandler)
r.Get("/by-created-range", bookings.GetBookingsByCreatedRangeHandler)
r.Put("/{id}/reschedule", bookings.AdminRescheduleBookingHandler)
r.Put("/{id}/progress", bookings.ProgressBookingHandler)
r.Post("/{id}/confirm", bookings.ConfirmBookingHandler)
r.Post("/{id}/cancel", bookings.AdminCancelBookingHandler)
r.Post("/reserve", bookings.AdminReserveSlotHandler)
// Edit request endpoints
r.Get("/edit-requests", bookings.AdminListAllEditRequestsHandler)
r.Get("/{id}/edit-request", bookings.AdminGetBookingEditRequestHandler)
r.Post("/{id}/edit-requests/{request_id}/approve", bookings.AdminApproveEditRequestHandler)
r.Post("/{id}/edit-requests/{request_id}/deny", bookings.AdminRejectEditRequestHandler)
})
r.Route("/admin/users", func(r chi.Router) {
r.Get("/", user.ListAdminUsersHandler)
r.Get("/{id}", user.GetAdminUserHandler)
r.Get("/{id}/relationship", user.GetCustomerRelationshipHandler)
r.Get("/{id}/patch-tests/eligible", user.GetEligiblePatchTestServicesHandler)
r.Post("/{id}/patch-tests", user.AddPatchTestHandler)
r.Get("/{id}/giftcard-balance", payments.GetUserGiftCardBalanceAdmin)
r.Get("/{id}/payment-methods", payments.AdminGetUserPaymentMethods)
})
r.Route("/admin/today", func(r chi.Router) {
r.Get("/current-next", today.GetCurrentAndNextHandler)
r.Get("/appointments", today.GetTodayAppointmentsHandler)
r.Get("/pending-approvals", today.GetPendingApprovalsHandler)
})
r.Route("/admin/notifications", func(r chi.Router) {
r.Get("/", notifications.GetNotifications)
r.Get("/unread-count", notifications.GetUnreadCount)
r.Post("/{id}/acknowledge", notifications.AcknowledgeNotification)
})
r.Route("/admin/time-blockers", func(r chi.Router) {
r.Get("/", scheduling.ListTimeBlockers)
r.Post("/", scheduling.CreateTimeBlocker)
r.Delete("/{id}", scheduling.DeleteTimeBlocker)
})
r.Route("/admin/discount-campaigns", func(r chi.Router) {
r.Get("/", admin.GetDiscountCampaigns)
r.Post("/", admin.CreateDiscountCampaign)
r.Put("/{id}", admin.UpdateDiscountCampaign)
r.Delete("/{id}", admin.DeleteDiscountCampaign)
r.Get("/{id}/stats", admin.GetCampaignStats)
})
// Admin payment routes
r.Post("/admin/bookings/{id}/payment", payments.CreateTerminalPayment)
r.Get("/admin/payments/{checkout_id}/status", payments.GetCheckoutStatus)
r.Post("/admin/payments/{payment_id}/refund", payments.RefundPayment)
// Admin gift card routes
r.Get("/admin/gift-cards", payments.GetGiftCards)
r.Post("/admin/gift-cards", payments.CreateGiftCard)
r.Put("/admin/gift-cards/{id}/topup", payments.TopUpGiftCard)
r.Post("/admin/gift-cards/{from}/transfer", payments.TransferGiftCard)
r.Get("/admin/gift-cards/expired-balances", payments.GetExpiredBalances)
r.Post("/admin/gift-cards/expired-balances/claim", payments.ClaimExpiredBalance)
// Admin till sale routes (POS transactions not linked to bookings)
r.Post("/admin/till/sale", payments.CreateTillSale)
r.Get("/admin/till/sale/checkout/{checkout_id}/status", payments.GetTillCheckoutStatus)
r.Route("/admin/settings", func(r chi.Router) {
r.Get("/", admin.GetBusinessSettings)
r.Put("/", admin.UpdateBusinessSettings)
})
})
})
// Webhooks (no auth - Square sends to base path)
r.Post("/webhooks/square", webhooks.HandleSquareWebhook)
srv := &http.Server{
Addr: ":8080",
Handler: r,
}
quit := make(chan os.Signal, 1)
signal.Notify(quit, syscall.SIGTERM, syscall.SIGINT)
go func() {
<-quit
log.Println("Shutting down server...")
ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second)
defer cancel()
if err := srv.Shutdown(ctx); err != nil {
log.Printf("Server forced to shutdown: %v", err)
}
}()
fmt.Println("Server is listening on :8080")
if err := srv.ListenAndServe(); err != nil && err != http.ErrServerClosed {
log.Fatalf("Server failed to start: %v", err)
}
log.Println("Server exited")
}