diff --git a/backend/handlers/bookings/bookings.go b/backend/handlers/bookings/bookings.go index 3980604..1ee7a92 100644 --- a/backend/handlers/bookings/bookings.go +++ b/backend/handlers/bookings/bookings.go @@ -912,6 +912,10 @@ func UpdateBookingServicesHandler(w http.ResponseWriter, r *http.Request) { http.Error(w, "Invalid request", http.StatusBadRequest) return } + if err := validators.Validate.Struct(&req); err != nil { + http.Error(w, err.Error(), http.StatusBadRequest) + return + } if len(req.ServiceIDs) == 0 { http.Error(w, "At least one service is required", http.StatusBadRequest) @@ -1216,6 +1220,10 @@ func SearchAdminBookingsHandler(w http.ResponseWriter, r *http.Request) { http.Error(w, "Search query 'q' is required", http.StatusBadRequest) return } + if len(query) > 200 { + http.Error(w, "search query too long", http.StatusBadRequest) + return + } page, perPage := 1, 10 if pageStr := r.URL.Query().Get("page"); pageStr != "" { @@ -1492,6 +1500,11 @@ func CreateBookingHandler(w http.ResponseWriter, r *http.Request) { isGuest = true } + if err := validators.Validate.Struct(&req); err != nil { + http.Error(w, err.Error(), http.StatusBadRequest) + return + } + if req.StartTime.IsZero() { http.Error(w, "Start time is required", http.StatusBadRequest) return @@ -1566,15 +1579,16 @@ func CreateBookingHandler(w http.ResponseWriter, r *http.Request) { } eligibleFrom := testedAt.Add(time.Duration(noticeHours) * time.Hour) - if time.Now().Before(eligibleFrom) { - hoursLeft := time.Until(eligibleFrom).Hours() - http.Error(w, fmt.Sprintf("You must wait %.0f hours after your patch test before booking this service.", hoursLeft), http.StatusBadRequest) + if req.StartTime.Before(eligibleFrom) { + hoursLeft := eligibleFrom.Sub(req.StartTime).Hours() + http.Error(w, fmt.Sprintf("Booking time is before the %.0f hour notice period after patch test. Earliest booking: %s", hoursLeft, eligibleFrom.Format("2006-01-02 15:04")), http.StatusBadRequest) return } var expiryMonths int if err := db.DB.QueryRow(r.Context(), `SELECT expiry_months FROM patch_tests WHERE id = $1`, patchTestID).Scan(&expiryMonths); err == nil { - if time.Now().After(testedAt.AddDate(0, expiryMonths, 0)) { + expiresAt := testedAt.AddDate(0, expiryMonths, 0) + if req.StartTime.After(expiresAt) { http.Error(w, "Your patch test has expired. Please complete a new patch test.", http.StatusBadRequest) return } @@ -1777,6 +1791,10 @@ func EditBookingHandler(w http.ResponseWriter, r *http.Request) { http.Error(w, "Invalid request", http.StatusBadRequest) return } + if err := validators.Validate.Struct(&req); err != nil { + http.Error(w, err.Error(), http.StatusBadRequest) + return + } if req.StartTime.IsZero() { http.Error(w, "Start time is required", http.StatusBadRequest) @@ -2220,6 +2238,10 @@ func ConfirmBookingHandler(w http.ResponseWriter, r *http.Request) { http.Error(w, "Invalid request", http.StatusBadRequest) return } + if err := validators.Validate.Struct(&req); err != nil { + http.Error(w, err.Error(), http.StatusBadRequest) + return + } for _, override := range req.ServiceOverrides { if override.OverridePrice != nil && *override.OverridePrice < 0 { @@ -2434,6 +2456,10 @@ func DeleteBookingHandler(w http.ResponseWriter, r *http.Request) { http.Error(w, "Invalid request", http.StatusBadRequest) return } + if err := validators.Validate.Struct(&req); err != nil { + http.Error(w, err.Error(), http.StatusBadRequest) + return + } allowed := map[string]bool{ "client_cancelled": true, "we_cancelled": true, "re-schedule": true, "no_show": true, } @@ -3291,6 +3317,10 @@ func AdminRescheduleBookingHandler(w http.ResponseWriter, r *http.Request) { http.Error(w, "Invalid request", http.StatusBadRequest) return } + if err := validators.Validate.Struct(&req); err != nil { + http.Error(w, err.Error(), http.StatusBadRequest) + return + } if req.StartTime.IsZero() { http.Error(w, "Start time is required", http.StatusBadRequest) diff --git a/backend/handlers/bookings/bookings_test.go b/backend/handlers/bookings/bookings_test.go index 239fd9b..b9c5fa9 100644 --- a/backend/handlers/bookings/bookings_test.go +++ b/backend/handlers/bookings/bookings_test.go @@ -3608,9 +3608,12 @@ func TestBookings_Create_PatchTestRequired_WithinNoticePeriod(t *testing.T) { token := jwt.GenerateUserToken(userID) - // Try to book within notice period (24h required, but only 1h passed) - futureTime := time.Now().Add(72 * time.Hour).Truncate(time.Second) - futureTime = time.Date(futureTime.Year(), futureTime.Month(), futureTime.Day(), 10, 0, 0, 0, futureTime.Location()) + // Try to book within notice period (24h required, but only 1h passed). + // Booking start_time must be before testedAt + noticeHours to trigger this. + // Use a booking time in the future (passes 1-hour advance check) but before + // eligibleFrom (testedAt + 24h = now + 23h). + futureTime := time.Now().Add(2 * time.Hour).Truncate(time.Second) + futureTime = time.Date(futureTime.Year(), futureTime.Month(), futureTime.Day(), futureTime.Hour(), 0, 0, 0, futureTime.Location()) req := CreateBookingRequest{ StartTime: futureTime, ServiceIDs: []string{serviceID}, @@ -3623,8 +3626,8 @@ func TestBookings_Create_PatchTestRequired_WithinNoticePeriod(t *testing.T) { t.Errorf("expected status 400 for within notice period, got %d. body: %s", w.Code, w.Body.String()) } - if !bytes.Contains(w.Body.Bytes(), []byte("wait")) { - t.Errorf("expected error message about waiting, got: %s", w.Body.String()) + if !bytes.Contains(w.Body.Bytes(), []byte("notice period")) { + t.Errorf("expected error message about notice period, got: %s", w.Body.String()) } } diff --git a/backend/handlers/bookings/manage.go b/backend/handlers/bookings/manage.go index f139446..ebe511e 100644 --- a/backend/handlers/bookings/manage.go +++ b/backend/handlers/bookings/manage.go @@ -440,6 +440,10 @@ func AdminCreateBookingForUserHandler(w http.ResponseWriter, r *http.Request) { http.Error(w, "Invalid request", http.StatusBadRequest) return } + if err := validators.Validate.Struct(&req); err != nil { + http.Error(w, err.Error(), http.StatusBadRequest) + return + } // Extract idempotency key from header idempotencyKey := r.Header.Get("Idempotency-Key")