518 lines
14 KiB
Go
518 lines
14 KiB
Go
package auth
|
|
|
|
import (
|
|
"context"
|
|
"crussell/auth"
|
|
"crussell/db"
|
|
"crussell/internal/dav"
|
|
"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)
|
|
)
|
|
|
|
// Login state management
|
|
var (
|
|
loginStateMu sync.Mutex
|
|
loginInProgress = make(map[string]bool)
|
|
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)
|
|
}
|
|
}
|
|
loginStateMu.Unlock()
|
|
}
|
|
}()
|
|
}
|
|
|
|
type RegisterRequest struct {
|
|
FirstName string `json:"firstName"`
|
|
LastName string `json:"lastName"`
|
|
Email string `json:"email"`
|
|
Password string `json:"password"`
|
|
Phone string `json:"phone"`
|
|
DateOfBirth string `json:"dateOfBirth"`
|
|
AgreedToPolicy bool `json:"agreedToPolicy"`
|
|
}
|
|
|
|
type LoginRequest struct {
|
|
Email string `json:"email"`
|
|
Password string `json:"password"`
|
|
}
|
|
|
|
// 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
|
|
}
|
|
|
|
// 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
|
|
}
|
|
|
|
// 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
|
|
}
|
|
|
|
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
|
|
}
|
|
|
|
// 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 loginInProgress[userID] {
|
|
loginStateMu.Unlock()
|
|
http.Error(w, "login already in progress", http.StatusConflict) // 409
|
|
return
|
|
}
|
|
loginInProgress[userID] = true
|
|
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, 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})
|
|
}
|
|
|
|
// 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())
|
|
|
|
// 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
|
|
}
|
|
|
|
// Generate new token
|
|
newToken, 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})
|
|
}
|
|
|
|
type VerificationCodeRequest struct {
|
|
Email string `json:"email"`
|
|
}
|
|
|
|
type VerifyCodeRequest struct {
|
|
Code string `json:"code"`
|
|
}
|
|
|
|
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
|
|
}
|
|
|
|
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
|
|
}
|
|
|
|
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))
|
|
}
|