package auth import ( "context" "crussell/auth" "crussell/db" "crussell/internal/dav" "crussell/internal/validators" "crussell/mw" "crypto/rand" "database/sql" "encoding/json" "errors" "fmt" "log" "net/http" "net/mail" "regexp" "strings" "sync" "time" "github.com/nyaruka/phonenumbers" "golang.org/x/crypto/bcrypt" "golang.org/x/text/cases" "golang.org/x/text/language" ) var ( titleCaser = cases.Title(language.English) ) const maxLoginInProgress = 20 // Login state management var ( loginStateMu sync.Mutex loginInProgress = make(map[string]time.Time) loginAttempts = make(map[string]time.Time) ) func init() { go func() { ticker := time.NewTicker(1 * time.Hour) defer ticker.Stop() for range ticker.C { loginStateMu.Lock() now := time.Now() for userID, lastAttempt := range loginAttempts { // Remove attempts older than 1 hour if now.Sub(lastAttempt) > 1*time.Hour { delete(loginAttempts, userID) } } // Clean up stuck loginInProgress entries (older than 30s) for userID, startedAt := range loginInProgress { if now.Sub(startedAt) > 30*time.Second { delete(loginInProgress, userID) } } loginStateMu.Unlock() } }() } type RegisterRequest struct { FirstName string `json:"firstName" validate:"required,min=1,max=50"` LastName string `json:"lastName" validate:"required,min=1,max=50"` Email string `json:"email" validate:"required,email,max=254"` Password string `json:"password" validate:"required,max=72"` Phone string `json:"phone" validate:"required"` DateOfBirth string `json:"dateOfBirth" validate:"required"` AgreedToPolicy bool `json:"agreedToPolicy"` ReferralCode string `json:"referralCode,omitempty" validate:"omitempty,max=12"` } type LoginRequest struct { Email string `json:"email" validate:"required,email,max=254"` Password string `json:"password" validate:"required,max=72"` } // POST /api/register func RegisterHandler(w http.ResponseWriter, r *http.Request) { var req RegisterRequest if err := json.NewDecoder(r.Body).Decode(&req); err != nil { http.Error(w, "invalid request", http.StatusBadRequest) return } if err := validators.Validate.Struct(&req); err != nil { http.Error(w, err.Error(), http.StatusBadRequest) return } // Must accept terms if !req.AgreedToPolicy { http.Error(w, "must agree to terms", http.StatusBadRequest) return } // Normalize input req.Email = strings.ToLower(strings.TrimSpace(req.Email)) req.FirstName = strings.TrimSpace(req.FirstName) req.LastName = strings.TrimSpace(req.LastName) req.Phone = strings.TrimSpace(req.Phone) req.DateOfBirth = strings.TrimSpace(req.DateOfBirth) // Check required fields if req.FirstName == "" || req.LastName == "" || req.Email == "" || req.Phone == "" || req.DateOfBirth == "" || req.Password == "" { http.Error(w, "all fields are required", http.StatusBadRequest) return } // Password must not exceed bcrypt's 72-byte limit if len(req.Password) > 72 { http.Error(w, "password must be 72 characters or less", http.StatusBadRequest) return } // Validate name (unicode letters, spaces, hyphen, apostrophe, dot) nameRegex := regexp.MustCompile(`^[\p{L}\p{M}\s\-'\.]+$`) if !nameRegex.MatchString(req.FirstName) { http.Error(w, "invalid characters in name", http.StatusBadRequest) return } // Validate length if len(req.FirstName) > 50 || len(req.FirstName) < 1 { http.Error(w, "first name must be 1-50 characters", http.StatusBadRequest) return } if len(req.LastName) > 50 || len(req.LastName) < 1 { http.Error(w, "last name must be 1-50 characters", http.StatusBadRequest) return } // Validate email format _, err := mail.ParseAddress(req.Email) if err != nil { http.Error(w, "invalid email format", http.StatusBadRequest) return } // Normalize phone (remove spaces, hyphens, parentheses) req.Phone = strings.Map(func(r rune) rune { if r >= '0' && r <= '9' || r == '+' { return r } return -1 }, req.Phone) // Validate UK phone number phone, err := ValidateUKPhoneNumber(req.Phone) if err != nil { http.Error(w, "invalid phone number format", http.StatusBadRequest) return } req.Phone = strings.TrimSpace(phone) // Convert names to title case req.FirstName = titleCaser.String(strings.ToLower(req.FirstName)) req.LastName = titleCaser.String(strings.ToLower(req.LastName)) // Parse date of birth dob, err := time.Parse("2006-01-02", req.DateOfBirth) if err != nil { http.Error(w, "invalid date format", http.StatusBadRequest) return } // Reject if younger than 16 if !dob.Before(time.Now().AddDate(-16, 0, 0)) { http.Error(w, "account creation prohibited for users under 16. Please call to book an appointment.", http.StatusBadRequest) return } // Validate referral code if provided var referrerID *string req.ReferralCode = strings.TrimSpace(req.ReferralCode) if req.ReferralCode != "" { if len(req.ReferralCode) != 12 { http.Error(w, "referral code must be exactly 12 characters", http.StatusBadRequest) return } referralCodeRegex := regexp.MustCompile(`^[a-zA-Z0-9]{12}$`) if !referralCodeRegex.MatchString(req.ReferralCode) { http.Error(w, "referral code must be alphanumeric", http.StatusBadRequest) return } // Look up referrer by referral code err := db.DB.QueryRow(r.Context(), "SELECT id FROM users WHERE referral_code = $1", req.ReferralCode).Scan(&referrerID) if err != nil { http.Error(w, "invalid referral code", http.StatusBadRequest) return } } // Hash password hash, err := bcrypt.GenerateFromPassword([]byte(req.Password), bcrypt.DefaultCost) if err != nil { http.Error(w, "server error", http.StatusInternalServerError) return } tx, err := db.DB.Begin(r.Context()) if err != nil { http.Error(w, "server error", http.StatusInternalServerError) return } defer tx.Rollback(r.Context()) now := time.Now() // Insert and return the generated ID var userID string err = tx.QueryRow(r.Context(), ` INSERT INTO users (n_first_name, n_last_name, phone, date_of_birth, email, password_hash, account_type, privacy_policy_and_terms_consent, policy_consent_updated_at, created_at, updated_at) VALUES ($1, $2, $3, $4, $5, $6, 'email', $7, $8, $8, $8) RETURNING id `, req.FirstName, req.LastName, req.Phone, dob, req.Email, string(hash), req.AgreedToPolicy, now).Scan(&userID) if err != nil { if strings.Contains(err.Error(), "duplicate key") { http.Error(w, "an account with this email already exists", http.StatusConflict) } else { http.Error(w, "could not create user", http.StatusInternalServerError) } return } // Record referral relationship if referral code was provided if referrerID != nil { _, err = tx.Exec(r.Context(), ` INSERT INTO user_referrals (referrer_id, referred_id) VALUES ($1, $2) ON CONFLICT DO NOTHING `, *referrerID, userID) if err != nil { http.Error(w, "could not process referral", http.StatusInternalServerError) return } } if err := tx.Commit(r.Context()); err != nil { http.Error(w, "server error", http.StatusInternalServerError) return } go func() { input := dav.ContactInput{ UserID: userID, FirstName: req.FirstName, LastName: req.LastName, Email: req.Email, Phone: req.Phone, DOB: req.DateOfBirth, } if err := dav.Service.CreateContact(1, userID, input); err != nil { log.Printf("Warning: Failed to create contact in DAV for user %s: %v", userID, err) } }() w.WriteHeader(http.StatusCreated) } func ValidateUKPhoneNumber(phone string) (string, error) { num, err := phonenumbers.Parse(phone, "GB") if err != nil { return "", err } if !phonenumbers.IsValidNumber(num) { return "", fmt.Errorf("invalid phone number") } // Check if it's actually a UK number if phonenumbers.GetRegionCodeForNumber(num) != "GB" { return "", fmt.Errorf("only UK numbers allowed") } // Format in E.164 format (+44...) return phonenumbers.Format(num, phonenumbers.E164), nil } // POST /api/login func LoginHandler(w http.ResponseWriter, r *http.Request) { var req LoginRequest if err := json.NewDecoder(r.Body).Decode(&req); err != nil { http.Error(w, "invalid request", http.StatusBadRequest) return } if err := validators.Validate.Struct(&req); err != nil { http.Error(w, err.Error(), http.StatusBadRequest) return } // Normalize email req.Email = strings.ToLower(strings.TrimSpace(req.Email)) var userID, passwordHash, role string ctx := context.Background() err := db.DB.QueryRow(ctx, ` SELECT id, password_hash, account_role FROM users WHERE email = $1 AND account_type = 'email' `, req.Email).Scan(&userID, &passwordHash, &role) if err != nil { http.Error(w, "invalid credentials", http.StatusUnauthorized) return } // Check if user is already logging in loginStateMu.Lock() if t, ok := loginInProgress[userID]; ok && time.Since(t) < 30*time.Second { loginStateMu.Unlock() http.Error(w, "login already in progress", http.StatusConflict) // 409 return } // Cap the map size - drop new request if at capacity if len(loginInProgress) >= maxLoginInProgress { loginStateMu.Unlock() http.Error(w, "server busy, try again later", http.StatusTooManyRequests) return } loginInProgress[userID] = time.Now() loginStateMu.Unlock() // Always clear flag when done defer func() { loginStateMu.Lock() delete(loginInProgress, userID) loginStateMu.Unlock() }() // Enforce 1 attempt per 5s loginStateMu.Lock() if last, ok := loginAttempts[userID]; ok { since := time.Since(last) if since < 5*time.Second { wait := 5*time.Second - since loginStateMu.Unlock() time.Sleep(wait) } else { loginStateMu.Unlock() } } else { loginStateMu.Unlock() } // Verify password if err := bcrypt.CompareHashAndPassword([]byte(passwordHash), []byte(req.Password)); err != nil { loginStateMu.Lock() loginAttempts[userID] = time.Now() loginStateMu.Unlock() http.Error(w, "invalid credentials", http.StatusUnauthorized) return } // On success, clear attempts loginStateMu.Lock() delete(loginAttempts, userID) loginStateMu.Unlock() // Update last login _, err = db.DB.Exec(ctx, `UPDATE users SET last_login_at = NOW() WHERE id = $1`, userID) if err != nil { fmt.Println("Failed to update last_login_at:", err) } // Generate JWT tokenString, jti, err := auth.GenerateToken(userID, role) if err != nil { http.Error(w, "could not generate token", http.StatusInternalServerError) return } w.Header().Set("Content-Type", "application/json") json.NewEncoder(w).Encode(auth.AuthResponse{Token: tokenString, JTI: jti}) } // POST /api/refresh-token (requires auth middleware) func RefreshTokenHandler(w http.ResponseWriter, r *http.Request) { userID, _ := mw.GetUserID(r.Context()) role, _ := mw.GetUserRole(r.Context()) oldJTI, _ := mw.GetJTI(r.Context()) // Verify user still exists and role hasn't changed var currentRole string err := db.DB.QueryRow(r.Context(), ` SELECT account_role FROM users WHERE id = $1 `, userID).Scan(¤tRole) if err != nil { http.Error(w, "user not found", http.StatusUnauthorized) return } // If role changed, force re-login if currentRole != role { http.Error(w, "role changed, please log in again", http.StatusUnauthorized) return } // Revoke the old token's JTI before issuing a new one if oldJTI != "" { // Use a 30-day expiry from now for the revoked JTI (matching token lifetime) auth.RevokeJTI(oldJTI, time.Now().Add(30*24*time.Hour)) } // Generate new token newToken, jti, err := auth.GenerateToken(userID, currentRole) if err != nil { http.Error(w, "could not generate token", http.StatusInternalServerError) return } w.Header().Set("Content-Type", "application/json") json.NewEncoder(w).Encode(auth.AuthResponse{Token: newToken, JTI: jti}) } // POST /api/logout (requires auth middleware) func LogoutHandler(w http.ResponseWriter, r *http.Request) { jti, ok := mw.GetJTI(r.Context()) if !ok || jti == "" { http.Error(w, "invalid token", http.StatusUnauthorized) return } // Revoke the JTI — it will be kept until the token's natural expiry (30 days) auth.RevokeJTI(jti, time.Now().Add(30*24*time.Hour)) w.Header().Set("Content-Type", "application/json") json.NewEncoder(w).Encode(map[string]bool{"success": true}) } type VerificationCodeRequest struct { Email string `json:"email" validate:"required,email,max=254"` } type VerifyCodeRequest struct { Code string `json:"code" validate:"required,max=64"` } type VerificationResponse struct { Success bool `json:"success"` Message string `json:"message,omitempty"` } func GenerateVerificationCodeHandler(w http.ResponseWriter, r *http.Request) { var req VerificationCodeRequest if err := json.NewDecoder(r.Body).Decode(&req); err != nil { http.Error(w, "invalid request", http.StatusBadRequest) return } if err := validators.Validate.Struct(&req); err != nil { http.Error(w, err.Error(), http.StatusBadRequest) return } email := strings.TrimSpace(strings.ToLower(req.Email)) if email == "" { http.Error(w, "email is required", http.StatusBadRequest) return } var userID string err := db.DB.QueryRow(r.Context(), "SELECT id FROM users WHERE LOWER(email) = $1", email, ).Scan(&userID) if err != nil { if errors.Is(err, sql.ErrNoRows) { w.Header().Set("Content-Type", "application/json") json.NewEncoder(w).Encode(VerificationResponse{Success: true, Message: "If the email exists, a verification code will be sent"}) return } log.Printf("Failed to look up user: %v", err) http.Error(w, "internal error", http.StatusInternalServerError) return } expiresAt := time.Now().Add(24 * time.Hour) var code string err = db.DB.QueryRow(r.Context(), `INSERT INTO verification_codes (user_id, purpose, expires_at) VALUES ($1, 'email_verify', $2) RETURNING code`, userID, expiresAt, ).Scan(&code) if err != nil { log.Printf("Failed to insert verification code: %v", err) http.Error(w, "internal error", http.StatusInternalServerError) return } log.Printf("DEBUG: Verification code for %s: %s (expires at %s)", email, code, expiresAt.Format(time.RFC3339)) w.Header().Set("Content-Type", "application/json") json.NewEncoder(w).Encode(VerificationResponse{Success: true, Message: "Verification code generated"}) } func VerifyCodeHandler(w http.ResponseWriter, r *http.Request) { var req VerifyCodeRequest if err := json.NewDecoder(r.Body).Decode(&req); err != nil { http.Error(w, "invalid request", http.StatusBadRequest) return } if err := validators.Validate.Struct(&req); err != nil { http.Error(w, err.Error(), http.StatusBadRequest) return } code := strings.TrimSpace(req.Code) if code == "" { http.Error(w, "code is required", http.StatusBadRequest) return } var userID string var purpose string var expiresAt time.Time err := db.DB.QueryRow(r.Context(), `SELECT user_id, purpose, expires_at FROM verification_codes WHERE code = $1 AND used_at IS NULL AND expires_at > NOW()`, code, ).Scan(&userID, &purpose, &expiresAt) if err != nil { if errors.Is(err, sql.ErrNoRows) { // Check if code exists but was already used or expired var checkUsedAt *time.Time checkErr := db.DB.QueryRow(r.Context(), `SELECT used_at FROM verification_codes WHERE code = $1`, code, ).Scan(&checkUsedAt) if checkErr != nil { // Code doesn't exist at all http.Error(w, "invalid or expired code", http.StatusBadRequest) return } // Code exists but was already used if checkUsedAt != nil { http.Error(w, "code already used", http.StatusForbidden) return } // Code exists but expired http.Error(w, "invalid or expired code", http.StatusBadRequest) return } log.Printf("Failed to verify code: %v", err) http.Error(w, "internal error", http.StatusInternalServerError) return } tx, err := db.DB.Begin(r.Context()) if err != nil { log.Printf("Failed to start transaction: %v", err) http.Error(w, "internal error", http.StatusInternalServerError) return } defer tx.Rollback(r.Context()) _, err = tx.Exec(r.Context(), `UPDATE verification_codes SET used_at = NOW() WHERE code = $1`, code, ) if err != nil { log.Printf("Failed to mark code as used: %v", err) http.Error(w, "internal error", http.StatusInternalServerError) return } if purpose == "email_verify" { _, err = tx.Exec(r.Context(), `UPDATE users SET account_role = 'verified_email' WHERE id = $1 AND account_role = 'unverified_email'`, userID, ) if err != nil { log.Printf("Failed to update user role: %v", err) http.Error(w, "internal error", http.StatusInternalServerError) return } } if err := tx.Commit(r.Context()); err != nil { log.Printf("Failed to commit verification: %v", err) http.Error(w, "internal error", http.StatusInternalServerError) return } w.Header().Set("Content-Type", "application/json") json.NewEncoder(w).Encode(VerificationResponse{Success: true, Message: "Email verified successfully"}) } func generateSecureCode(length int) string { bytes := make([]byte, length) if _, err := rand.Read(bytes); err != nil { log.Printf("Failed to generate random code: %v", err) return strings.ToLower(fmt.Sprintf("%x", time.Now().UnixNano())) } return strings.ToLower(fmt.Sprintf("%x", bytes)) }