Add CreateTestDatabase function for parallel isolated test databases per package. - Add CreateTestDatabase() for isolated test DBs (parallel-safe) - Move all TestMain functions to per-package testmain_test.go files - Remove old TestMain from handlers_test.go, jwt_test.go, main_test.go - Add JWT init guard in main.go to skip when -test.* flags detected - Update testdb.go with admin DSN and proper cleanup - Rename test database to crussell_test_db for consistency - Replace testdb.NewPool + testdb.Migrate pattern with CreateTestDatabase Ultraworked with [Sisyphus](https://github.com/code-yeongyu/oh-my-openagent) Co-authored-by: Sisyphus <clio-agent@sisyphuslabs.ai>
446 lines
16 KiB
Go
446 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.DB != nil {
|
|
if err := db.DB.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"
|
|
}
|
|
|
|
w.Header().Set("Content-Type", "application/json")
|
|
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) {
|
|
|
|
// 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)
|
|
|
|
// 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.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")
|
|
}
|