Fix test setup and middleware chain - Handler tests now passing

- Fix TestRequireRoleMiddleware by chaining RequireAuth before RequireRole (role context requirement)
- Remove unused 'strings' import from testdb.go
- Create crussell_test database in Docker setup
- Tests now properly initialize authentication context for role-based tests

Result: handlers test suite passes (13/13 tests)
Remaining failures in admin/auth/bookings/portfolio/scheduling/services/user packages need further investigation (environment setup, database constraints, endpoint initialization)
This commit is contained in:
2026-02-21 23:50:17 +00:00
parent e858c782a4
commit 44cac94f64
20 changed files with 6725 additions and 7 deletions
File diff suppressed because it is too large Load Diff
+255
View File
@@ -0,0 +1,255 @@
//go:build test
// +build test
package admin
import (
"context"
"encoding/json"
"net/http"
"testing"
"crussell/db"
"crussell/handlers/services"
"crussell/mw"
)
func TestAdminServices_Create(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
// Create admin user in DB first
_, err := db.DB.Exec(context.Background(), `
INSERT INTO users (n_first_name, n_last_name, email, password_hash, account_role, account_type)
VALUES ('Admin', 'User', 'admin@test.com', 'hash', 'admin', 'email')
`)
if err != nil {
t.Fatalf("failed to create admin user: %v", err)
}
handler := http.HandlerFunc(services.CreateServiceHandler)
createReq := services.CreateServiceRequest{
Name: "Test Manicure",
Description: stringPtr("A test manicure service"),
Price: 35.00,
DurationMinutes: 45,
PatchTestDurationHours: 0,
MinimumAgeRequired: 16,
}
w := makeAdminRequest(handler, "POST", "/api/admin/services", createReq)
if w.Code != http.StatusCreated {
t.Errorf("expected status 201, got %d. body: %s", w.Code, w.Body.String())
}
var response services.Service
if err := json.Unmarshal(w.Body.Bytes(), &response); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
if response.Name != "Test Manicure" {
t.Errorf("expected name 'Test Manicure', got %s", response.Name)
}
if response.Price != 35.00 {
t.Errorf("expected price 35.00, got %f", response.Price)
}
if !response.IsActive {
t.Error("expected new service to be active by default")
}
}
func TestAdminServices_List(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
// Insert test services
_, err := db.DB.Exec(context.Background(), `
INSERT INTO services (name, description, price, duration_minutes, is_active, patch_test_duration_hours, minimum_age_required)
VALUES
('Manicure', 'Basic manicure', 25.00, 30, true, 0, 0),
('Pedicure', 'Basic pedicure', 30.00, 45, false, 0, 0),
('Gel Polish', 'Gel polish service', 40.00, 60, true, 48, 16)
`)
if err != nil {
t.Fatalf("failed to create services: %v", err)
}
handler := http.HandlerFunc(services.AllServicesHandler)
w := makeAdminRequest(handler, "GET", "/api/admin/services", nil)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
var response []services.Service
if err := json.Unmarshal(w.Body.Bytes(), &response); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
if len(response) != 3 {
t.Errorf("expected 3 services, got %d", len(response))
}
// Verify all services including inactive are returned
found := map[string]bool{}
for _, s := range response {
found[s.Name] = true
}
if !found["Manicure"] {
t.Error("expected Manicure in response")
}
if !found["Pedicure"] {
t.Error("expected Pedicure in response (including inactive)")
}
if !found["Gel Polish"] {
t.Error("expected Gel Polish in response")
}
}
func TestAdminServices_Toggle(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
// Create a service
var serviceID string
err := db.DB.QueryRow(context.Background(), `
INSERT INTO services (name, description, price, duration_minutes, is_active, patch_test_duration_hours, minimum_age_required)
VALUES ('Test Service', 'A test service', 50.00, 60, true, 0, 16)
RETURNING id
`).Scan(&serviceID)
if err != nil {
t.Fatalf("failed to create service: %v", err)
}
handler := http.HandlerFunc(services.ToggleService)
w := makeAdminRequest(handler, "PUT", "/api/admin/services/"+serviceID+"/toggle", nil)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
// Verify service is now inactive
var isActive bool
err = db.DB.QueryRow(context.Background(), "SELECT is_active FROM services WHERE id = $1", serviceID).Scan(&isActive)
if err != nil {
t.Fatalf("failed to check service: %v", err)
}
if isActive {
t.Error("expected service to be inactive after toggle")
}
// Toggle again
w = makeAdminRequest(handler, "PUT", "/api/admin/services/"+serviceID+"/toggle", nil)
if w.Code != http.StatusOK {
t.Errorf("expected status 200 on second toggle, got %d", w.Code)
}
// Verify service is active again
err = db.DB.QueryRow(context.Background(), "SELECT is_active FROM services WHERE id = $1", serviceID).Scan(&isActive)
if err != nil {
t.Fatalf("failed to check service: %v", err)
}
if !isActive {
t.Error("expected service to be active after second toggle")
}
}
func TestAdminServices_Delete(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
// Create a service
var serviceID string
err := db.DB.QueryRow(context.Background(), `
INSERT INTO services (name, description, price, duration_minutes, is_active, patch_test_duration_hours, minimum_age_required)
VALUES ('Test Service', 'A test service', 50.00, 60, true, 0, 16)
RETURNING id
`).Scan(&serviceID)
if err != nil {
t.Fatalf("failed to create service: %v", err)
}
handler := http.HandlerFunc(services.DeleteServiceHandler)
w := makeAdminRequest(handler, "DELETE", "/api/admin/services/"+serviceID, nil)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
// Verify service is deleted
var count int
err = db.DB.QueryRow(context.Background(), "SELECT COUNT(*) FROM services WHERE id = $1", serviceID).Scan(&count)
if err != nil {
t.Fatalf("failed to check service: %v", err)
}
if count != 0 {
t.Error("expected service to be deleted")
}
}
func TestAdminServices_NonAdmin(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
// Create regular user in DB
_, err := db.DB.Exec(context.Background(), `
INSERT INTO users (n_first_name, n_last_name, email, password_hash, account_role, account_type)
VALUES ('Regular', 'User', 'user@test.com', 'hash', 'verified_email', 'email')
`)
if err != nil {
t.Fatalf("failed to create user: %v", err)
}
// Test CREATE - should get 403 when using middleware
createHandler := mw.RequireAdmin(http.HandlerFunc(services.CreateServiceHandler))
createReq := services.CreateServiceRequest{
Name: "Test Service",
Description: stringPtr("Test"),
Price: 50.00,
DurationMinutes: 60,
PatchTestDurationHours: 0,
MinimumAgeRequired: 16,
}
w := makeUserRequest(createHandler, "POST", "/api/admin/services", createReq)
if w.Code != http.StatusForbidden {
t.Errorf("CREATE: expected status 403, got %d", w.Code)
}
// Test LIST - should get 403 when using middleware
listHandler := mw.RequireAdmin(http.HandlerFunc(services.AllServicesHandler))
w = makeUserRequest(listHandler, "GET", "/api/admin/services", nil)
if w.Code != http.StatusForbidden {
t.Errorf("LIST: expected status 403, got %d", w.Code)
}
// Test TOGGLE - should get 403 when using middleware
var serviceID string
err = db.DB.QueryRow(context.Background(), `
INSERT INTO services (name, description, price, duration_minutes, is_active, patch_test_duration_hours, minimum_age_required)
VALUES ('Test Service', 'A test service', 50.00, 60, true, 0, 16)
RETURNING id
`).Scan(&serviceID)
if err != nil {
t.Fatalf("failed to create service: %v", err)
}
toggleHandler := mw.RequireAdmin(http.HandlerFunc(services.ToggleService))
w = makeUserRequest(toggleHandler, "PUT", "/api/admin/services/"+serviceID+"/toggle", nil)
if w.Code != http.StatusForbidden {
t.Errorf("TOGGLE: expected status 403, got %d", w.Code)
}
// Test DELETE - should get 403 when using middleware
deleteHandler := mw.RequireAdmin(http.HandlerFunc(services.DeleteServiceHandler))
w = makeUserRequest(deleteHandler, "DELETE", "/api/admin/services/"+serviceID, nil)
if w.Code != http.StatusForbidden {
t.Errorf("DELETE: expected status 403, got %d", w.Code)
}
}
func stringPtr(s string) *string {
return &s
}
+71
View File
@@ -0,0 +1,71 @@
//go:build test
// +build test
package admin
import (
"bytes"
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
"crussell/db"
"crussell/mw"
"crussell/testutils/jwt"
"crussell/testutils/testdb"
)
// setupTestDB replaces the global db.DB with a test pool and returns a cleanup function
func setupTestDB(t *testing.T) func() {
t.Helper()
pool := testdb.Pool(t)
testdb.Migrate(t, pool)
originalDB := db.DB
db.DB = pool
jwt.Init()
return func() {
db.DB = originalDB
pool.Close()
}
}
// makeAdminRequest creates a request with admin context
func makeAdminRequest(handler http.Handler, method, path string, body interface{}) *httptest.ResponseRecorder {
return makeRequestWithContext(handler, method, path, body, "admin-test-001", "admin")
}
// makeUserRequest creates a request with regular user context
func makeUserRequest(handler http.Handler, method, path string, body interface{}) *httptest.ResponseRecorder {
return makeRequestWithContext(handler, method, path, body, "user-test-001", "verified_email")
}
// makeRequestWithContext creates a request with specific user context
func makeRequestWithContext(handler http.Handler, method, path string, body interface{}, userID, role string) *httptest.ResponseRecorder {
var req *http.Request
if body != nil {
bodyBytes, _ := json.Marshal(body)
req = httptest.NewRequest(method, path, bytes.NewReader(bodyBytes))
req.Header.Set("Content-Type", "application/json")
} else {
req = httptest.NewRequest(method, path, nil)
}
// Set up context with user ID and role (simulating middleware)
ctx := context.WithValue(req.Context(), mw.UserIDKey, userID)
ctx = context.WithValue(ctx, mw.UserRoleKey, role)
req = req.WithContext(ctx)
w := httptest.NewRecorder()
handler.ServeHTTP(w, req)
return w
}
func parseResponseBody(w *httptest.ResponseRecorder, dest interface{}) error {
return json.Unmarshal(w.Body.Bytes(), dest)
}
+362
View File
@@ -0,0 +1,362 @@
//go:build test
// +build test
package admin
import (
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"strconv"
"testing"
"crussell/db"
"crussell/handlers/notifications"
"crussell/handlers/today"
"crussell/mw"
)
func TestAdminToday_CurrentNext(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
// Create test user
var userID string
err := db.DB.QueryRow(context.Background(), `
INSERT INTO users (n_first_name, n_last_name, email, password_hash, account_role, account_type)
VALUES ('Test', 'User', 'testuser@test.com', 'hash', 'verified_email', 'standard')
RETURNING id
`).Scan(&userID)
if err != nil {
t.Fatalf("failed to create user: %v", err)
}
// Create service
var serviceID string
err = db.DB.QueryRow(context.Background(), `
INSERT INTO services (name, description, price, duration_minutes, is_active)
VALUES ('Manicure', 'Basic manicure', 25.00, 30, true)
RETURNING id
`).Scan(&serviceID)
if err != nil {
t.Fatalf("failed to create service: %v", err)
}
// Create booking for today (in_progress)
_, err = db.DB.Exec(context.Background(), `
INSERT INTO bookings (user_id, start_time, status, created_at)
VALUES ($1, NOW(), 'in_progress', NOW())
`, userID)
if err != nil {
t.Fatalf("failed to create booking: %v", err)
}
// Get the booking ID
var bookingID string
err = db.DB.QueryRow(context.Background(), `
SELECT id FROM bookings WHERE user_id = $1 ORDER BY created_at DESC LIMIT 1
`, userID).Scan(&bookingID)
if err != nil {
t.Fatalf("failed to get booking ID: %v", err)
}
// Add service to booking
_, err = db.DB.Exec(context.Background(), `
INSERT INTO booking_services (booking_id, service_id)
VALUES ($1, $2)
`, bookingID, serviceID)
if err != nil {
t.Fatalf("failed to add service to booking: %v", err)
}
handler := http.HandlerFunc(today.GetCurrentAndNextHandler)
w := makeAdminRequest(handler, "GET", "/api/admin/today/current-next", nil)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
var response today.CurrentNextResponse
if err := json.Unmarshal(w.Body.Bytes(), &response); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
if response.Current == nil {
t.Errorf("expected current appointment, got nil")
}
if response.Current != nil && response.Current.ID != bookingID {
t.Errorf("expected booking ID %s, got %s", bookingID, response.Current.ID)
}
}
func TestAdminToday_Appointments(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
// Create test user
var userID string
err := db.DB.QueryRow(context.Background(), `
INSERT INTO users (n_first_name, n_last_name, email, password_hash, account_role, account_type)
VALUES ('Test', 'User', 'testuser@test.com', 'hash', 'verified_email', 'standard')
RETURNING id
`).Scan(&userID)
if err != nil {
t.Fatalf("failed to create user: %v", err)
}
// Create service
var serviceID string
err = db.DB.QueryRow(context.Background(), `
INSERT INTO services (name, description, price, duration_minutes, is_active)
VALUES ('Manicure', 'Basic manicure', 25.00, 30, true)
RETURNING id
`).Scan(&serviceID)
if err != nil {
t.Fatalf("failed to create service: %v", err)
}
// Create booking for today
_, err = db.DB.Exec(context.Background(), `
INSERT INTO bookings (user_id, start_time, status, created_at)
VALUES ($1, NOW(), 'confirmed', NOW())
`, userID)
if err != nil {
t.Fatalf("failed to create booking: %v", err)
}
// Get the booking ID
var bookingID string
err = db.DB.QueryRow(context.Background(), `
SELECT id FROM bookings WHERE user_id = $1 ORDER BY created_at DESC LIMIT 1
`, userID).Scan(&bookingID)
if err != nil {
t.Fatalf("failed to get booking ID: %v", err)
}
// Add service to booking
_, err = db.DB.Exec(context.Background(), `
INSERT INTO booking_services (booking_id, service_id)
VALUES ($1, $2)
`, bookingID, serviceID)
if err != nil {
t.Fatalf("failed to add service to booking: %v", err)
}
handler := http.HandlerFunc(today.GetTodayAppointmentsHandler)
w := makeAdminRequest(handler, "GET", "/api/admin/today/appointments", nil)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
var response today.TodayAppointmentsResponse
if err := json.Unmarshal(w.Body.Bytes(), &response); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
if len(response.Appointments) != 1 {
t.Errorf("expected 1 appointment, got %d", len(response.Appointments))
}
if len(response.Appointments) > 0 && response.Appointments[0].ID != bookingID {
t.Errorf("expected booking ID %s, got %s", bookingID, response.Appointments[0].ID)
}
}
func TestAdminToday_PendingApprovals(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
// Create test user
var userID string
err := db.DB.QueryRow(context.Background(), `
INSERT INTO users (n_first_name, n_last_name, email, password_hash, account_role, account_type)
VALUES ('Test', 'User', 'testuser@test.com', 'hash', 'verified_email', 'standard')
RETURNING id
`).Scan(&userID)
if err != nil {
t.Fatalf("failed to create user: %v", err)
}
// Create service
var serviceID string
err = db.DB.QueryRow(context.Background(), `
INSERT INTO services (name, description, price, duration_minutes, is_active)
VALUES ('Manicure', 'Basic manicure', 25.00, 30, true)
RETURNING id
`).Scan(&serviceID)
if err != nil {
t.Fatalf("failed to create service: %v", err)
}
// Create pending booking
_, err = db.DB.Exec(context.Background(), `
INSERT INTO bookings (user_id, start_time, status, created_at)
VALUES ($1, NOW() + INTERVAL '1 day', 'pending', NOW())
`, userID)
if err != nil {
t.Fatalf("failed to create booking: %v", err)
}
// Get the booking ID
var bookingID string
err = db.DB.QueryRow(context.Background(), `
SELECT id FROM bookings WHERE user_id = $1 ORDER BY created_at DESC LIMIT 1
`, userID).Scan(&bookingID)
if err != nil {
t.Fatalf("failed to get booking ID: %v", err)
}
// Add service to booking
_, err = db.DB.Exec(context.Background(), `
INSERT INTO booking_services (booking_id, service_id)
VALUES ($1, $2)
`, bookingID, serviceID)
if err != nil {
t.Fatalf("failed to add service to booking: %v", err)
}
handler := http.HandlerFunc(today.GetPendingApprovalsHandler)
w := makeAdminRequest(handler, "GET", "/api/admin/today/pending-approvals", nil)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
var response today.PendingApprovalsResponse
if err := json.Unmarshal(w.Body.Bytes(), &response); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
if len(response.Approvals) != 1 {
t.Errorf("expected 1 pending approval, got %d", len(response.Approvals))
}
if len(response.Approvals) > 0 && response.Approvals[0].ID != bookingID {
t.Errorf("expected booking ID %s, got %s", bookingID, response.Approvals[0].ID)
}
}
func TestAdminNotifications_List(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
// Create a notification
_, err := db.DB.Exec(context.Background(), `
INSERT INTO admin_notifications (reason, booking_id, user_id, created_at)
VALUES ('pending_booking', 1, 1, NOW())
`)
if err != nil {
t.Fatalf("failed to create notification: %v", err)
}
handler := http.HandlerFunc(notifications.GetNotifications)
w := makeAdminRequest(handler, "GET", "/api/admin/notifications", nil)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
var response notifications.AdminNotificationListResponse
if err := json.Unmarshal(w.Body.Bytes(), &response); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
if response.Total != 1 {
t.Errorf("expected 1 notification, got %d", response.Total)
}
if len(response.Notifications) != 1 {
t.Errorf("expected 1 notification in list, got %d", len(response.Notifications))
}
if len(response.Notifications) > 0 && response.Notifications[0].Reason != "pending_booking" {
t.Errorf("expected reason 'pending_booking', got %s", response.Notifications[0].Reason)
}
}
func TestAdminNotifications_Acknowledge(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
// Create a notification
var notificationID int
err := db.DB.QueryRow(context.Background(), `
INSERT INTO admin_notifications (reason, booking_id, user_id, created_at)
VALUES ('pending_booking', 1, 1, NOW())
RETURNING id
`).Scan(&notificationID)
if err != nil {
t.Fatalf("failed to create notification: %v", err)
}
// Create request to acknowledge
req := httptest.NewRequest("POST", "/api/admin/notifications/"+strconv.Itoa(notificationID)+"/acknowledge", nil)
ctx := req.Context()
ctx = context.WithValue(ctx, mw.UserIDKey, "admin-test-001")
ctx = context.WithValue(ctx, mw.UserRoleKey, "admin")
req = req.WithContext(ctx)
w := httptest.NewRecorder()
handler := http.HandlerFunc(notifications.AcknowledgeNotification)
handler.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
// Verify it's acknowledged
var acknowledged bool
err = db.DB.QueryRow(context.Background(), `
SELECT acknowledged_at IS NOT NULL FROM admin_notifications WHERE id = $1
`, notificationID).Scan(&acknowledged)
if err != nil {
t.Fatalf("failed to check acknowledgment: %v", err)
}
if !acknowledged {
t.Errorf("expected notification to be acknowledged")
}
}
func TestAdminToday_NonAdmin(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
// Test current-next endpoint
currentNextHandler := mw.RequireAdmin(http.HandlerFunc(today.GetCurrentAndNextHandler))
w := makeUserRequest(currentNextHandler, "GET", "/api/admin/today/current-next", nil)
if w.Code != http.StatusForbidden {
t.Errorf("CurrentNext: expected status 403, got %d", w.Code)
}
// Test appointments endpoint
appointmentsHandler := mw.RequireAdmin(http.HandlerFunc(today.GetTodayAppointmentsHandler))
w = makeUserRequest(appointmentsHandler, "GET", "/api/admin/today/appointments", nil)
if w.Code != http.StatusForbidden {
t.Errorf("Appointments: expected status 403, got %d", w.Code)
}
// Test pending-approvals endpoint
pendingApprovalsHandler := mw.RequireAdmin(http.HandlerFunc(today.GetPendingApprovalsHandler))
w = makeUserRequest(pendingApprovalsHandler, "GET", "/api/admin/today/pending-approvals", nil)
if w.Code != http.StatusForbidden {
t.Errorf("PendingApprovals: expected status 403, got %d", w.Code)
}
// Test notifications list endpoint
notificationsHandler := mw.RequireAdmin(http.HandlerFunc(notifications.GetNotifications))
w = makeUserRequest(notificationsHandler, "GET", "/api/admin/notifications", nil)
if w.Code != http.StatusForbidden {
t.Errorf("Notifications List: expected status 403, got %d", w.Code)
}
// Test notifications acknowledge endpoint
ackHandler := mw.RequireAdmin(http.HandlerFunc(notifications.AcknowledgeNotification))
w = makeUserRequest(ackHandler, "POST", "/api/admin/notifications/1/acknowledge", nil)
if w.Code != http.StatusForbidden {
t.Errorf("Notifications Acknowledge: expected status 403, got %d", w.Code)
}
}
+350
View File
@@ -0,0 +1,350 @@
//go:build test
// +build test
package admin
import (
"context"
"encoding/json"
"net/http"
"testing"
"crussell/db"
"crussell/handlers/user"
"crussell/mw"
)
func TestAdminUsers_List(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
// Create test users
_, err := db.DB.Exec(context.Background(), `
INSERT INTO users (n_first_name, n_last_name, email, password_hash, account_role, account_type)
VALUES
('Alice', 'Smith', 'alice@test.com', 'hash1', 'admin', 'standard'),
('Bob', 'Jones', 'bob@test.com', 'hash2', 'verified_email', 'standard'),
('Charlie', 'Brown', 'charlie@test.com', 'hash3', 'verified_email', 'vip')
`)
if err != nil {
t.Fatalf("failed to create users: %v", err)
}
handler := http.HandlerFunc(user.ListAdminUsersHandler)
w := makeAdminRequest(handler, "GET", "/api/admin/users", nil)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
var response user.UserListResponse
if err := json.Unmarshal(w.Body.Bytes(), &response); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
if response.Total != 3 {
t.Errorf("expected 3 users, got %d", response.Total)
}
if len(response.Users) != 3 {
t.Errorf("expected 3 users in list, got %d", len(response.Users))
}
}
func TestAdminUsers_Get(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
// Create test user
var userID string
err := db.DB.QueryRow(context.Background(), `
INSERT INTO users (n_first_name, n_last_name, email, password_hash, account_role, account_type)
VALUES ('Test', 'User', 'testuser@test.com', 'hash', 'verified_email', 'vip')
RETURNING id
`).Scan(&userID)
if err != nil {
t.Fatalf("failed to create user: %v", err)
}
handler := http.HandlerFunc(user.GetAdminUserHandler)
w := makeAdminRequest(handler, "GET", "/api/admin/users/"+userID, nil)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
var response user.AdminUserDetail
if err := json.Unmarshal(w.Body.Bytes(), &response); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
if response.ID != userID {
t.Errorf("expected user ID %s, got %s", userID, response.ID)
}
if response.AccountType != "vip" {
t.Errorf("expected account type 'vip', got %s", response.AccountType)
}
}
func TestAdminUsers_Get_NotFound(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
handler := http.HandlerFunc(user.GetAdminUserHandler)
w := makeAdminRequest(handler, "GET", "/api/admin/users/nonexistent-id", nil)
if w.Code != http.StatusNotFound {
t.Errorf("expected status 404, got %d", w.Code)
}
}
func TestAdminUsers_PatchTests_Eligible(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
// Create test user
var userID string
err := db.DB.QueryRow(context.Background(), `
INSERT INTO users (n_first_name, n_last_name, email, password_hash, account_role, account_type)
VALUES ('Test', 'User', 'testuser@test.com', 'hash', 'verified_email', 'standard')
RETURNING id
`).Scan(&userID)
if err != nil {
t.Fatalf("failed to create user: %v", err)
}
// Create services - some with patch test, some without
_, err = db.DB.Exec(context.Background(), `
INSERT INTO services (name, description, price, duration_minutes, is_active, patch_test_duration_hours, minimum_age_required)
VALUES
('Basic Manicure', 'Basic manicure', 25.00, 30, true, 0, 0),
('Gel Polish Full Set', 'Gel polish service', 45.00, 60, true, 48, 16),
('Luxury Gel Manicure', 'Luxury gel', 55.00, 75, true, 48, 16),
('Inactive Service', 'Inactive', 30.00, 30, false, 48, 16)
`)
if err != nil {
t.Fatalf("failed to create services: %v", err)
}
handler := http.HandlerFunc(user.GetEligiblePatchTestServicesHandler)
w := makeAdminRequest(handler, "GET", "/api/admin/users/"+userID+"/patch-tests/eligible", nil)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
var response []user.ServiceForPatchTest
if err := json.Unmarshal(w.Body.Bytes(), &response); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
// Should return 2 services (the two with patch test duration > 0 that are active)
if len(response) != 2 {
t.Errorf("expected 2 eligible services, got %d. body: %s", len(response), w.Body.String())
}
}
func TestAdminUsers_PatchTests_Eligible_WithExisting(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
// Create test user
var userID string
err := db.DB.QueryRow(context.Background(), `
INSERT INTO users (n_first_name, n_last_name, email, password_hash, account_role, account_type)
VALUES ('Test', 'User', 'testuser@test.com', 'hash', 'verified_email', 'standard')
RETURNING id
`).Scan(&userID)
if err != nil {
t.Fatalf("failed to create user: %v", err)
}
// Create services with patch test
var serviceID1, serviceID2 string
err = db.DB.QueryRow(context.Background(), `
INSERT INTO services (name, description, price, duration_minutes, is_active, patch_test_duration_hours, minimum_age_required)
VALUES ('Gel Polish Full Set', 'Gel polish service', 45.00, 60, true, 48, 16)
RETURNING id
`).Scan(&serviceID1)
if err != nil {
t.Fatalf("failed to create service 1: %v", err)
}
err = db.DB.QueryRow(context.Background(), `
INSERT INTO services (name, description, price, duration_minutes, is_active, patch_test_duration_hours, minimum_age_required)
VALUES ('Luxury Gel Manicure', 'Luxury gel', 55.00, 75, true, 48, 16)
RETURNING id
`).Scan(&serviceID2)
if err != nil {
t.Fatalf("failed to create service 2: %v", err)
}
// Add one patch test for the user
_, err = db.DB.Exec(context.Background(), `
INSERT INTO user_service_patch_tests (user_id, service_id, last_time)
VALUES ($1, $2, NOW())
`, userID, serviceID1)
if err != nil {
t.Fatalf("failed to add patch test: %v", err)
}
handler := http.HandlerFunc(user.GetEligiblePatchTestServicesHandler)
w := makeAdminRequest(handler, "GET", "/api/admin/users/"+userID+"/patch-tests/eligible", nil)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
var response []user.ServiceForPatchTest
if err := json.Unmarshal(w.Body.Bytes(), &response); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
// Should return only 1 service (the one not already added)
if len(response) != 1 {
t.Errorf("expected 1 eligible service, got %d. body: %s", len(response), w.Body.String())
}
if len(response) > 0 && response[0].ID != serviceID2 {
t.Errorf("expected service %s, got %s", serviceID2, response[0].ID)
}
}
func TestAdminUsers_AddPatchTest(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
// Create test user
var userID string
err := db.DB.QueryRow(context.Background(), `
INSERT INTO users (n_first_name, n_last_name, email, password_hash, account_role, account_type)
VALUES ('Test', 'User', 'testuser@test.com', 'hash', 'verified_email', 'standard')
RETURNING id
`).Scan(&userID)
if err != nil {
t.Fatalf("failed to create user: %v", err)
}
// Create service with patch test
var serviceID string
err = db.DB.QueryRow(context.Background(), `
INSERT INTO services (name, description, price, duration_minutes, is_active, patch_test_duration_hours, minimum_age_required)
VALUES ('Gel Polish Full Set', 'Gel polish service', 45.00, 60, true, 48, 16)
RETURNING id
`).Scan(&serviceID)
if err != nil {
t.Fatalf("failed to create service: %v", err)
}
handler := http.HandlerFunc(user.AddPatchTestHandler)
reqBody := user.AddPatchTestRequest{ServiceID: serviceID}
w := makeAdminRequest(handler, "POST", "/api/admin/users/"+userID+"/patch-tests", reqBody)
if w.Code != http.StatusCreated {
t.Errorf("expected status 201, got %d. body: %s", w.Code, w.Body.String())
}
// Verify patch test was added
var count int
err = db.DB.QueryRow(context.Background(), `
SELECT COUNT(*) FROM user_service_patch_tests WHERE user_id = $1 AND service_id = $2
`, userID, serviceID).Scan(&count)
if err != nil {
t.Fatalf("failed to check patch test: %v", err)
}
if count != 1 {
t.Errorf("expected 1 patch test record, got %d", count)
}
}
func TestAdminUsers_AddPatchTest_InvalidService(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
// Create test user
var userID string
err := db.DB.QueryRow(context.Background(), `
INSERT INTO users (n_first_name, n_last_name, email, password_hash, account_role, account_type)
VALUES ('Test', 'User', 'testuser@test.com', 'hash', 'verified_email', 'standard')
RETURNING id
`).Scan(&userID)
if err != nil {
t.Fatalf("failed to create user: %v", err)
}
// Create service without patch test requirement
var serviceID string
err = db.DB.QueryRow(context.Background(), `
INSERT INTO services (name, description, price, duration_minutes, is_active, patch_test_duration_hours, minimum_age_required)
VALUES ('Basic Manicure', 'Basic manicure', 25.00, 30, true, 0, 0)
RETURNING id
`).Scan(&serviceID)
if err != nil {
t.Fatalf("failed to create service: %v", err)
}
handler := http.HandlerFunc(user.AddPatchTestHandler)
reqBody := user.AddPatchTestRequest{ServiceID: serviceID}
w := makeAdminRequest(handler, "POST", "/api/admin/users/"+userID+"/patch-tests", reqBody)
if w.Code != http.StatusBadRequest {
t.Errorf("expected status 400, got %d. body: %s", w.Code, w.Body.String())
}
}
func TestAdminUsers_NonAdmin(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
// Create regular user in DB
_, err := db.DB.Exec(context.Background(), `
INSERT INTO users (n_first_name, n_last_name, email, password_hash, account_role, account_type)
VALUES ('Regular', 'User', 'user@test.com', 'hash', 'verified_email', 'standard')
`)
if err != nil {
t.Fatalf("failed to create user: %v", err)
}
// Create test user for GET
var targetUserID string
err = db.DB.QueryRow(context.Background(), `
INSERT INTO users (n_first_name, n_last_name, email, password_hash, account_role, account_type)
VALUES ('Target', 'User', 'target@test.com', 'hash', 'verified_email', 'standard')
RETURNING id
`).Scan(&targetUserID)
if err != nil {
t.Fatalf("failed to create target user: %v", err)
}
// Test LIST - should get 403 when using middleware
listHandler := mw.RequireAdmin(http.HandlerFunc(user.ListAdminUsersHandler))
w := makeUserRequest(listHandler, "GET", "/api/admin/users", nil)
if w.Code != http.StatusForbidden {
t.Errorf("LIST: expected status 403, got %d", w.Code)
}
// Test GET - should get 403 when using middleware
getHandler := mw.RequireAdmin(http.HandlerFunc(user.GetAdminUserHandler))
w = makeUserRequest(getHandler, "GET", "/api/admin/users/"+targetUserID, nil)
if w.Code != http.StatusForbidden {
t.Errorf("GET: expected status 403, got %d", w.Code)
}
// Test eligible patch tests - should get 403 when using middleware
eligibleHandler := mw.RequireAdmin(http.HandlerFunc(user.GetEligiblePatchTestServicesHandler))
w = makeUserRequest(eligibleHandler, "GET", "/api/admin/users/"+targetUserID+"/patch-tests/eligible", nil)
if w.Code != http.StatusForbidden {
t.Errorf("ELIGIBLE: expected status 403, got %d", w.Code)
}
// Test add patch test - should get 403 when using middleware
addHandler := mw.RequireAdmin(http.HandlerFunc(user.AddPatchTestHandler))
w = makeUserRequest(addHandler, "POST", "/api/admin/users/"+targetUserID+"/patch-tests", map[string]string{"service_id": "some-service-id"})
if w.Code != http.StatusForbidden {
t.Errorf("ADD: expected status 403, got %d", w.Code)
}
}
+754
View File
@@ -0,0 +1,754 @@
//go:build test
// +build test
package auth
import (
"bytes"
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
"time"
"crussell/db"
"crussell/internal/dav"
"crussell/testutils/fixtures"
"crussell/testutils/jwt"
"crussell/testutils/testdb"
"github.com/jackc/pgx/v5/pgxpool"
)
// setupTestDB replaces the global db.DB with a test pool and returns a cleanup function
func setupTestDB(t *testing.T) func() {
t.Helper()
pool := testdb.Pool(t)
testdb.Migrate(t, pool)
// Replace global db.DB with test pool
originalDB := db.DB
db.DB = pool
// Initialize JWT for tests
jwt.Init()
// Set up a minimal dav.Service to avoid nil pointer panic
// The real service is only used in a goroutine in RegisterHandler
dav.Service = &dav.BaseService{}
return func() {
db.DB = originalDB
pool.Close()
}
}
// helper function to make JSON request
func makeRequest(handler http.Handler, method, path string, body interface{}) *httptest.ResponseRecorder {
var req *http.Request
if body != nil {
bodyBytes, _ := json.Marshal(body)
req = httptest.NewRequest(method, path, bytes.NewReader(bodyBytes))
req.Header.Set("Content-Type", "application/json")
} else {
req = httptest.NewRequest(method, path, nil)
}
w := httptest.NewRecorder()
handler.ServeHTTP(w, req)
return w
}
// Helper to parse response body
func parseResponseBody(w *httptest.ResponseRecorder, dest interface{}) error {
return json.Unmarshal(w.Body.Bytes(), dest)
}
// =============================================================================
// Register Handler Tests
// =============================================================================
func TestRegister_Success(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
handler := http.HandlerFunc(RegisterHandler)
body := RegisterRequest{
FirstName: "John",
LastName: "Doe",
Email: "john.doe@test.com",
Password: "password123",
Phone: "07700900000",
DateOfBirth: "1990-01-15",
AgreedToPolicy: true,
}
w := makeRequest(handler, "POST", "/api/register", body)
if w.Code != http.StatusCreated {
t.Errorf("expected status 201, got %d. body: %s", w.Code, w.Body.String())
}
// Verify user was created in DB
var userID string
err := db.DB.QueryRow(context.Background(),
"SELECT id FROM users WHERE email = $1", "john.doe@test.com").Scan(&userID)
if err != nil {
t.Errorf("failed to find user in DB: %v", err)
}
// Clean up
db.DB.Exec(context.Background(), "DELETE FROM users WHERE id = $1", userID)
}
func TestRegister_InvalidInput_MissingFields(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
handler := http.HandlerFunc(RegisterHandler)
tests := []struct {
name string
body RegisterRequest
}{
{
name: "missing firstName",
body: RegisterRequest{LastName: "Doe", Email: "test@test.com", Password: "pass", Phone: "07700900000", DateOfBirth: "1990-01-15", AgreedToPolicy: true},
},
{
name: "missing lastName",
body: RegisterRequest{FirstName: "John", Email: "test@test.com", Password: "pass", Phone: "07700900000", DateOfBirth: "1990-01-15", AgreedToPolicy: true},
},
{
name: "missing email",
body: RegisterRequest{FirstName: "John", LastName: "Doe", Password: "pass", Phone: "07700900000", DateOfBirth: "1990-01-15", AgreedToPolicy: true},
},
{
name: "missing phone",
body: RegisterRequest{FirstName: "John", LastName: "Doe", Email: "test@test.com", Password: "pass", DateOfBirth: "1990-01-15", AgreedToPolicy: true},
},
{
name: "missing dateOfBirth",
body: RegisterRequest{FirstName: "John", LastName: "Doe", Email: "test@test.com", Password: "pass", Phone: "07700900000", AgreedToPolicy: true},
},
{
name: "did not agree to policy",
body: RegisterRequest{FirstName: "John", LastName: "Doe", Email: "test@test.com", Password: "pass", Phone: "07700900000", DateOfBirth: "1990-01-15", AgreedToPolicy: false},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
w := makeRequest(handler, "POST", "/api/register", tt.body)
if w.Code != http.StatusBadRequest {
t.Errorf("expected status 400, got %d", w.Code)
}
})
}
}
func TestRegister_InvalidInput_InvalidEmail(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
handler := http.HandlerFunc(RegisterHandler)
body := RegisterRequest{
FirstName: "John",
LastName: "Doe",
Email: "not-an-email",
Password: "password123",
Phone: "07700900000",
DateOfBirth: "1990-01-15",
AgreedToPolicy: true,
}
w := makeRequest(handler, "POST", "/api/register", body)
if w.Code != http.StatusBadRequest {
t.Errorf("expected status 400, got %d. body: %s", w.Code, w.Body.String())
}
}
func TestRegister_InvalidInput_InvalidPhone(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
handler := http.HandlerFunc(RegisterHandler)
body := RegisterRequest{
FirstName: "John",
LastName: "Doe",
Email: "john@test.com",
Password: "password123",
Phone: "12345", // Not a valid UK phone
DateOfBirth: "1990-01-15",
AgreedToPolicy: true,
}
w := makeRequest(handler, "POST", "/api/register", body)
if w.Code != http.StatusBadRequest {
t.Errorf("expected status 400, got %d. body: %s", w.Code, w.Body.String())
}
}
func TestRegister_InvalidInput_Under16(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
handler := http.HandlerFunc(RegisterHandler)
// Calculate a date that makes them under 16
under16DOB := time.Now().AddDate(-15, 0, 0).Format("2006-01-02")
body := RegisterRequest{
FirstName: "Young",
LastName: "User",
Email: "young@test.com",
Password: "password123",
Phone: "07700900000",
DateOfBirth: under16DOB,
AgreedToPolicy: true,
}
w := makeRequest(handler, "POST", "/api/register", body)
if w.Code != http.StatusBadRequest {
t.Errorf("expected status 400, got %d. body: %s", w.Code, w.Body.String())
}
}
func TestRegister_DuplicateEmail(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
handler := http.HandlerFunc(RegisterHandler)
// First create a user
userID, err := fixtures.CreateTestUser(db.DB)
if err != nil {
t.Fatalf("failed to create test user: %v", err)
}
defer fixtures.DeleteUser(db.DB, userID)
// Now try to register with same email
body := RegisterRequest{
FirstName: "John",
LastName: "Doe",
Email: "user@test.com", // Same as fixture
Password: "password123",
Phone: "07700900000",
DateOfBirth: "1990-01-15",
AgreedToPolicy: true,
}
w := makeRequest(handler, "POST", "/api/register", body)
if w.Code != http.StatusConflict {
t.Errorf("expected status 409, got %d. body: %s", w.Code, w.Body.String())
}
}
// =============================================================================
// Login Handler Tests
// =============================================================================
func TestLogin_Success(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
handler := http.HandlerFunc(LoginHandler)
// Create a test user
userID, err := fixtures.CreateTestUser(db.DB)
if err != nil {
t.Fatalf("failed to create test user: %v", err)
}
defer fixtures.DeleteUser(db.DB, userID)
body := LoginRequest{
Email: "user@test.com",
Password: "testpassword123",
}
w := makeRequest(handler, "POST", "/api/login", body)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
var resp struct {
Token string `json:"token"`
}
if err := parseResponseBody(w, &resp); err != nil {
t.Errorf("failed to parse response: %v", err)
}
if resp.Token == "" {
t.Error("expected token in response, got empty string")
}
}
func TestLogin_InvalidCredentials_WrongPassword(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
handler := http.HandlerFunc(LoginHandler)
// Create a test user
userID, err := fixtures.CreateTestUser(db.DB)
if err != nil {
t.Fatalf("failed to create test user: %v", err)
}
defer fixtures.DeleteUser(db.DB, userID)
body := LoginRequest{
Email: "user@test.com",
Password: "wrongpassword",
}
w := makeRequest(handler, "POST", "/api/login", body)
if w.Code != http.StatusUnauthorized {
t.Errorf("expected status 401, got %d. body: %s", w.Code, w.Body.String())
}
}
func TestLogin_InvalidCredentials_NonExistentEmail(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
handler := http.HandlerFunc(LoginHandler)
body := LoginRequest{
Email: "nonexistent@test.com",
Password: "password123",
}
w := makeRequest(handler, "POST", "/api/login", body)
if w.Code != http.StatusUnauthorized {
t.Errorf("expected status 401, got %d. body: %s", w.Code, w.Body.String())
}
}
// =============================================================================
// Refresh Token Handler Tests
// =============================================================================
func TestRefreshToken_Success(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
handler := http.HandlerFunc(RefreshTokenHandler)
// Create a test user
userID, err := fixtures.CreateTestUser(db.DB)
if err != nil {
t.Fatalf("failed to create test user: %v", err)
}
defer fixtures.DeleteUser(db.DB, userID)
// Generate a valid token
token := jwt.GenerateTestToken(userID, "verified_email")
req := httptest.NewRequest("POST", "/api/refresh-token", nil)
req.Header.Set("Authorization", "Bearer "+token)
w := httptest.NewRecorder()
// Use the middleware to set up context
ctx := req.Context()
ctx = context.WithValue(ctx, "user_id", userID)
ctx = context.WithValue(ctx, "user_role", "verified_email")
req = req.WithContext(ctx)
handler.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
var resp struct {
Token string `json:"token"`
}
if err := parseResponseBody(w, &resp); err != nil {
t.Errorf("failed to parse response: %v", err)
}
if resp.Token == "" {
t.Error("expected new token in response, got empty string")
}
}
func TestRefreshToken_Unauthorized_NoToken(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
handler := http.HandlerFunc(RefreshTokenHandler)
req := httptest.NewRequest("POST", "/api/refresh-token", nil)
w := httptest.NewRecorder()
handler.ServeHTTP(w, req)
// Without proper auth middleware, userID/role won't be in context
// The handler tries to query DB with empty userID, which should fail
if w.Code != http.StatusUnauthorized && w.Code != http.StatusInternalServerError {
t.Errorf("expected status 401 or 500, got %d. body: %s", w.Code, w.Body.String())
}
}
// =============================================================================
// Verify Generate Handler Tests
// =============================================================================
func TestVerifyGenerate_ValidEmail(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
handler := http.HandlerFunc(GenerateVerificationCodeHandler)
// Create a test user
userID, err := fixtures.CreateTestUser(db.DB)
if err != nil {
t.Fatalf("failed to create test user: %v", err)
}
defer fixtures.DeleteUser(db.DB, userID)
body := VerificationCodeRequest{
Email: "user@test.com",
}
w := makeRequest(handler, "POST", "/api/verify/generate", body)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
var resp VerificationResponse
if err := parseResponseBody(w, &resp); err != nil {
t.Errorf("failed to parse response: %v", err)
}
if !resp.Success {
t.Error("expected success=true in response")
}
// Verify a code was created in DB
var codeID string
err = db.DB.QueryRow(context.Background(),
"SELECT id FROM verification_codes WHERE user_id = $1", userID).Scan(&codeID)
if err != nil {
t.Errorf("failed to find verification code in DB: %v", err)
}
// Clean up
db.DB.Exec(context.Background(), "DELETE FROM verification_codes WHERE user_id = $1", userID)
}
func TestVerifyGenerate_NonExistentEmail(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
handler := http.HandlerFunc(GenerateVerificationCodeHandler)
// Security: should return success even if email doesn't exist
body := VerificationCodeRequest{
Email: "nonexistent@test.com",
}
w := makeRequest(handler, "POST", "/api/verify/generate", body)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
var resp VerificationResponse
if err := parseResponseBody(w, &resp); err != nil {
t.Errorf("failed to parse response: %v", err)
}
// Should return success for security (don't reveal if email exists)
if !resp.Success {
t.Error("expected success=true in response for non-existent email")
}
}
// =============================================================================
// Verify Check Handler Tests
// =============================================================================
func TestVerifyCheck_ValidCode(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
handler := http.HandlerFunc(VerifyCodeHandler)
// Create a test user
userID, err := fixtures.CreateTestUser(db.DB)
if err != nil {
t.Fatalf("failed to create test user: %v", err)
}
defer fixtures.DeleteUser(db.DB, userID)
// Create a verification code
var code string
expiresAt := time.Now().Add(24 * time.Hour)
err = db.DB.QueryRow(context.Background(),
`INSERT INTO verification_codes (user_id, purpose, expires_at) VALUES ($1, 'email_verify', $2) RETURNING code`,
userID, expiresAt).Scan(&code)
if err != nil {
t.Fatalf("failed to create verification code: %v", err)
}
defer db.DB.Exec(context.Background(), "DELETE FROM verification_codes WHERE user_id = $1", userID)
body := VerifyCodeRequest{
Code: code,
}
w := makeRequest(handler, "POST", "/api/verify/check", body)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
var resp VerificationResponse
if err := parseResponseBody(w, &resp); err != nil {
t.Errorf("failed to parse response: %v", err)
}
if !resp.Success {
t.Error("expected success=true in response")
}
// Verify code is marked as used
var usedAt *time.Time
err = db.DB.QueryRow(context.Background(),
"SELECT used_at FROM verification_codes WHERE code = $1", code).Scan(&usedAt)
if err != nil || usedAt == nil {
t.Error("expected verification code to be marked as used")
}
}
func TestVerifyCheck_InvalidCode(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
handler := http.HandlerFunc(VerifyCodeHandler)
body := VerifyCodeRequest{
Code: "nonexistent-code-12345",
}
w := makeRequest(handler, "POST", "/api/verify/check", body)
if w.Code != http.StatusBadRequest {
t.Errorf("expected status 400, got %d. body: %s", w.Code, w.Body.String())
}
}
func TestVerifyCheck_ExpiredCode(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
handler := http.HandlerFunc(VerifyCodeHandler)
// Create a test user
userID, err := fixtures.CreateTestUser(db.DB)
if err != nil {
t.Fatalf("failed to create test user: %v", err)
}
defer fixtures.DeleteUser(db.DB, userID)
// Create an expired verification code
var code string
expiresAt := time.Now().Add(-1 * time.Hour) // Expired 1 hour ago
err = db.DB.QueryRow(context.Background(),
`INSERT INTO verification_codes (user_id, purpose, expires_at) VALUES ($1, 'email_verify', $2) RETURNING code`,
userID, expiresAt).Scan(&code)
if err != nil {
t.Fatalf("failed to create verification code: %v", err)
}
defer db.DB.Exec(context.Background(), "DELETE FROM verification_codes WHERE user_id = $1", userID)
body := VerifyCodeRequest{
Code: code,
}
w := makeRequest(handler, "POST", "/api/verify/check", body)
if w.Code != http.StatusBadRequest {
t.Errorf("expected status 400, got %d. body: %s", w.Code, w.Body.String())
}
}
// =============================================================================
// Additional Edge Case Tests
// =============================================================================
func TestLogin_InvalidRequest(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
handler := http.HandlerFunc(LoginHandler)
// Send invalid JSON
req := httptest.NewRequest("POST", "/api/login", bytes.NewReader([]byte("not json")))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
handler.ServeHTTP(w, req)
if w.Code != http.StatusBadRequest {
t.Errorf("expected status 400, got %d", w.Code)
}
}
func TestRegister_NameTooLong(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
handler := http.HandlerFunc(RegisterHandler)
// First name > 50 chars
longName := string(bytes.Repeat([]byte("a"), 51))
body := RegisterRequest{
FirstName: longName,
LastName: "Doe",
Email: "john@test.com",
Password: "password123",
Phone: "07700900000",
DateOfBirth: "1990-01-15",
AgreedToPolicy: true,
}
w := makeRequest(handler, "POST", "/api/register", body)
if w.Code != http.StatusBadRequest {
t.Errorf("expected status 400, got %d. body: %s", w.Code, w.Body.String())
}
}
func TestRegister_InvalidNameCharacters(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
handler := http.HandlerFunc(RegisterHandler)
// Name with numbers (invalid)
body := RegisterRequest{
FirstName: "John123",
LastName: "Doe",
Email: "john@test.com",
Password: "password123",
Phone: "07700900000",
DateOfBirth: "1990-01-15",
AgreedToPolicy: true,
}
w := makeRequest(handler, "POST", "/api/register", body)
if w.Code != http.StatusBadRequest {
t.Errorf("expected status 400, got %d. body: %s", w.Code, w.Body.String())
}
}
func TestVerifyCheck_AlreadyUsed(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
handler := http.HandlerFunc(VerifyCodeHandler)
// Create a test user
userID, err := fixtures.CreateTestUser(db.DB)
if err != nil {
t.Fatalf("failed to create test user: %v", err)
}
defer fixtures.DeleteUser(db.DB, userID)
// Create a verification code
var code string
expiresAt := time.Now().Add(24 * time.Hour)
err = db.DB.QueryRow(context.Background(),
`INSERT INTO verification_codes (user_id, purpose, expires_at) VALUES ($1, 'email_verify', $2) RETURNING code`,
userID, expiresAt).Scan(&code)
if err != nil {
t.Fatalf("failed to create verification code: %v", err)
}
defer db.DB.Exec(context.Background(), "DELETE FROM verification_codes WHERE user_id = $1", userID)
// First verification should succeed
body := VerifyCodeRequest{
Code: code,
}
w := makeRequest(handler, "POST", "/api/verify/check", body)
if w.Code != http.StatusOK {
t.Errorf("first verification: expected status 200, got %d", w.Code)
}
// Second verification with same code should return 403 (already used)
w = makeRequest(handler, "POST", "/api/verify/check", body)
if w.Code != http.StatusForbidden {
t.Errorf("second verification: expected status 403, got %d", w.Code)
}
}
func TestVerifyCheck_RoleChangeToVerified(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
handler := http.HandlerFunc(VerifyCodeHandler)
// Create an unverified user
userID, err := fixtures.CreateTestUnverifiedUser(db.DB)
if err != nil {
t.Fatalf("failed to create test user: %v", err)
}
defer fixtures.DeleteUser(db.DB, userID)
// Verify initial role is unverified_email
var initialRole string
err = db.DB.QueryRow(context.Background(),
"SELECT account_role FROM users WHERE id = $1", userID).Scan(&initialRole)
if err != nil {
t.Fatalf("failed to check initial role: %v", err)
}
if initialRole != "unverified_email" {
t.Errorf("expected initial role 'unverified_email', got %s", initialRole)
}
// Create a verification code
var code string
expiresAt := time.Now().Add(24 * time.Hour)
err = db.DB.QueryRow(context.Background(),
`INSERT INTO verification_codes (user_id, purpose, expires_at) VALUES ($1, 'email_verify', $2) RETURNING code`,
userID, expiresAt).Scan(&code)
if err != nil {
t.Fatalf("failed to create verification code: %v", err)
}
defer db.DB.Exec(context.Background(), "DELETE FROM verification_codes WHERE user_id = $1", userID)
// Verify the code
body := VerifyCodeRequest{
Code: code,
}
w := makeRequest(handler, "POST", "/api/verify/check", body)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
// Check that user's role changed to verified_email
var newRole string
err = db.DB.QueryRow(context.Background(),
"SELECT account_role FROM users WHERE id = $1", userID).Scan(&newRole)
if err != nil {
t.Errorf("failed to check new role: %v", err)
}
if newRole != "verified_email" {
t.Errorf("expected role to change to 'verified_email', got %s", newRole)
}
}
// Ensure test compilation - import pgxpool to avoid unused import
var _ = func() *pgxpool.Pool { return nil }
File diff suppressed because it is too large Load Diff
+1 -1
View File
@@ -82,7 +82,7 @@ func UserCancelBookingHandler(w http.ResponseWriter, r *http.Request) {
}
if err := tx.Commit(r.Context()); err != nil {
log.Printf("Failed to commit user cancel: %v", bookingID, err)
log.Printf("Failed to commit user cancel: %v, %v", bookingID, err)
http.Error(w, "Internal server error", http.StatusInternalServerError)
return
}
+174
View File
@@ -0,0 +1,174 @@
//go:build test
// +build test
package handlers
import (
"encoding/json"
"io"
"net/http"
"net/http/httptest"
"testing"
"crussell/mw"
"crussell/testutils/fixtures"
"crussell/testutils/jwt"
"crussell/testutils/testdb"
)
func TestHealthCheck(t *testing.T) {
handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusOK)
w.Write([]byte(`{"status":"ok"}`))
})
server := httptest.NewServer(handler)
defer server.Close()
resp, err := server.Client().Get(server.URL)
if err != nil {
t.Fatalf("failed to make request: %v", err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
t.Errorf("expected status 200, got %d", resp.StatusCode)
}
}
func TestRequireAuthMiddleware(t *testing.T) {
handler := mw.RequireAuth(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
userID, _ := r.Context().Value(mw.UserIDKey).(string)
role, _ := r.Context().Value(mw.UserRoleKey).(string)
w.Write([]byte(`{"user_id":"` + userID + `","role":"` + role + `"}`))
}))
t.Run("no auth header returns 401", func(t *testing.T) {
req := httptest.NewRequest("GET", "/", nil)
w := httptest.NewRecorder()
handler.ServeHTTP(w, req)
if w.Code != http.StatusUnauthorized {
t.Errorf("expected status 401, got %d", w.Code)
}
})
t.Run("valid token passes auth", func(t *testing.T) {
jwt.Init()
token := jwt.GenerateTestToken("test-user-123", "verified_email")
req := httptest.NewRequest("GET", "/", nil)
req.Header.Set("Authorization", "Bearer "+token)
w := httptest.NewRecorder()
handler.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d", w.Code)
}
body, _ := io.ReadAll(w.Body)
var resp map[string]string
json.Unmarshal(body, &resp)
if resp["user_id"] != "test-user-123" {
t.Errorf("expected user_id test-user-123, got %s", resp["user_id"])
}
if resp["role"] != "verified_email" {
t.Errorf("expected role verified_email, got %s", resp["role"])
}
})
t.Run("invalid token returns 401", func(t *testing.T) {
jwt.Init()
req := httptest.NewRequest("GET", "/", nil)
req.Header.Set("Authorization", "Bearer invalid-token")
w := httptest.NewRecorder()
handler.ServeHTTP(w, req)
if w.Code != http.StatusUnauthorized {
t.Errorf("expected status 401, got %d", w.Code)
}
})
}
func TestRequireRoleMiddleware(t *testing.T) {
// Chain RequireAuth before RequireRole to set the role in context
// RequireRole expects role to be in context, but that's only set by RequireAuth
adminOnlyHandler := mw.RequireAuth(mw.RequireRole("admin")(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusOK)
w.Write([]byte(`{"success":true}`))
})))
t.Run("admin role passes", func(t *testing.T) {
jwt.Init()
token := jwt.GenerateAdminToken()
req := httptest.NewRequest("GET", "/", nil)
req.Header.Set("Authorization", "Bearer "+token)
w := httptest.NewRecorder()
adminOnlyHandler.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d", w.Code)
}
})
t.Run("non-admin role returns 403", func(t *testing.T) {
jwt.Init()
token := jwt.GenerateUserToken("test-user")
req := httptest.NewRequest("GET", "/", nil)
req.Header.Set("Authorization", "Bearer "+token)
w := httptest.NewRecorder()
adminOnlyHandler.ServeHTTP(w, req)
if w.Code != http.StatusForbidden {
t.Errorf("expected status 403, got %d", w.Code)
}
})
}
func TestIntegration_UserFlow(t *testing.T) {
if testing.Short() {
t.Skip("skipping integration test in short mode")
}
pool := testdb.Pool(t)
defer pool.Close()
userID, err := fixtures.CreateTestUser(pool)
if err != nil {
t.Fatalf("failed to create test user: %v", err)
}
defer fixtures.DeleteUser(pool, userID)
jwt.Init()
token := jwt.GenerateUserToken(userID)
handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
authHeader := r.Header.Get("Authorization")
if authHeader == "" {
http.Error(w, "no auth", http.StatusUnauthorized)
return
}
userIDCtx, ok := r.Context().Value(mw.UserIDKey).(string)
if !ok {
http.Error(w, "no user id in context", http.StatusInternalServerError)
return
}
w.Write([]byte(`{"user_id":"` + userIDCtx + `"}`))
})
req := httptest.NewRequest("GET", "/", nil)
req.Header.Set("Authorization", "Bearer "+token)
w := httptest.NewRecorder()
mw.RequireAuth(handler).ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d", w.Code)
}
}
@@ -182,9 +182,17 @@ func AcknowledgePendingBookingNotification(tx interface{}, ctx context.Context,
SET acknowledged_at = NOW()
WHERE booking_id = $1 AND reason = 'pending_booking' AND acknowledged_at IS NULL
`
_, err := tx.(interface {
// Use type assertion to get the Exec method - pgx.Tx satisfies this interface
execer, ok := tx.(interface {
Exec(ctx context.Context, sql string, arguments ...interface{}) (pgconn.CommandTag, error)
}).Exec(ctx, query, bookingID)
})
if !ok {
log.Printf("Warning: cannot acknowledge notification - tx does not satisfy Execer interface for booking %s", bookingID)
return nil // Don't fail the main operation if notification ack fails
}
_, err := execer.Exec(ctx, query, bookingID)
if err != nil {
log.Printf("Failed to acknowledge pending booking notification for %s: %v", bookingID, err)
return err
+468
View File
@@ -0,0 +1,468 @@
//go:build test
// +build test
package portfolio
import (
"bytes"
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
"crussell/db"
"crussell/mw"
"crussell/testutils/jwt"
"crussell/testutils/testdb"
)
func setupTestDB(t *testing.T) func() {
t.Helper()
pool := testdb.Pool(t)
testdb.Migrate(t, pool)
originalDB := db.DB
db.DB = pool
jwt.Init()
return func() {
db.DB = originalDB
pool.Close()
}
}
func makeRequest(handler http.HandlerFunc, method, path string, body interface{}) *httptest.ResponseRecorder {
var req *http.Request
if body != nil {
bodyBytes, _ := json.Marshal(body)
req = httptest.NewRequest(method, path, bytes.NewReader(bodyBytes))
req.Header.Set("Content-Type", "application/json")
} else {
req = httptest.NewRequest(method, path, nil)
}
w := httptest.NewRecorder()
handler.ServeHTTP(w, req)
return w
}
func makeRequestWithContext(handler http.HandlerFunc, method, path string, body interface{}, userID, role string) *httptest.ResponseRecorder {
var req *http.Request
if body != nil {
bodyBytes, _ := json.Marshal(body)
req = httptest.NewRequest(method, path, bytes.NewReader(bodyBytes))
req.Header.Set("Content-Type", "application/json")
} else {
req = httptest.NewRequest(method, path, nil)
}
// Add user context
ctx := req.Context()
ctx = context.WithValue(ctx, mw.UserIDKey, userID)
ctx = context.WithValue(ctx, mw.UserRoleKey, role)
req = req.WithContext(ctx)
w := httptest.NewRecorder()
handler.ServeHTTP(w, req)
return w
}
// =============================================================================
// List Images Tests
// =============================================================================
func TestPortfolio_ListImages(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
// Insert test images
_, err := db.DB.Exec(context.Background(), `
INSERT INTO images (url, thumbnail_url, tag_names)
VALUES
('https://example.com/img1.jpg', 'https://example.com/img1_thumb.jpg', ARRAY['nature:forest', 'color:green']),
('https://example.com/img2.jpg', 'https://example.com/img2_thumb.jpg', ARRAY['nature:ocean', 'color:blue'])
`)
if err != nil {
t.Fatalf("failed to create images: %v", err)
}
handler := http.HandlerFunc(ListImages)
w := makeRequest(handler, "GET", "/api/portfolio/images", nil)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
var images []Image
if err := json.Unmarshal(w.Body.Bytes(), &images); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
if len(images) != 2 {
t.Errorf("expected 2 images, got %d", len(images))
}
}
func TestPortfolio_ListImages_WithTagFilter(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
// Insert test images
_, err := db.DB.Exec(context.Background(), `
INSERT INTO images (url, thumbnail_url, tag_names)
VALUES
('https://example.com/img1.jpg', 'https://example.com/img1_thumb.jpg', ARRAY['nature:forest']),
('https://example.com/img2.jpg', 'https://example.com/img2_thumb.jpg', ARRAY['nature:ocean'])
`)
if err != nil {
t.Fatalf("failed to create images: %v", err)
}
handler := http.HandlerFunc(ListImages)
w := makeRequest(handler, "GET", "/api/portfolio/images?tag=forest", nil)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
var images []Image
if err := json.Unmarshal(w.Body.Bytes(), &images); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
if len(images) != 1 {
t.Errorf("expected 1 image, got %d", len(images))
}
}
func TestPortfolio_ListImages_Empty(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
handler := http.HandlerFunc(ListImages)
w := makeRequest(handler, "GET", "/api/portfolio/images", nil)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
var images []Image
if err := json.Unmarshal(w.Body.Bytes(), &images); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
if len(images) != 0 {
t.Errorf("expected 0 images, got %d", len(images))
}
}
// =============================================================================
// List Tags Tests
// =============================================================================
func TestPortfolio_ListTags(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
// Insert test tags
_, err := db.DB.Exec(context.Background(), `
INSERT INTO tags (name) VALUES ('nature:forest'), ('nature:ocean'), ('color:green')
`)
if err != nil {
t.Fatalf("failed to create tags: %v", err)
}
handler := http.HandlerFunc(ListTags)
w := makeRequest(handler, "GET", "/api/portfolio/tags", nil)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
var tags []Tag
if err := json.Unmarshal(w.Body.Bytes(), &tags); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
if len(tags) != 3 {
t.Errorf("expected 3 tags, got %d", len(tags))
}
}
func TestPortfolio_ListTags_WithQuery(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
// Insert test tags
_, err := db.DB.Exec(context.Background(), `
INSERT INTO tags (name) VALUES ('nature:forest'), ('nature:ocean'), ('color:green')
`)
if err != nil {
t.Fatalf("failed to create tags: %v", err)
}
handler := http.HandlerFunc(ListTags)
w := makeRequest(handler, "GET", "/api/portfolio/tags?q=forest", nil)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
var tags []Tag
if err := json.Unmarshal(w.Body.Bytes(), &tags); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
if len(tags) != 1 {
t.Errorf("expected 1 tag, got %d", len(tags))
}
}
func TestPortfolio_ListTags_Empty(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
handler := http.HandlerFunc(ListTags)
w := makeRequest(handler, "GET", "/api/portfolio/tags", nil)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
var tags []Tag
if err := json.Unmarshal(w.Body.Bytes(), &tags); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
if len(tags) != 0 {
t.Errorf("expected 0 tags, got %d", len(tags))
}
}
// =============================================================================
// List Filters Tests
// =============================================================================
func TestPortfolio_ListFilters(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
// Insert test images with tags
_, err := db.DB.Exec(context.Background(), `
INSERT INTO images (url, thumbnail_url, tag_names)
VALUES
('https://example.com/img1.jpg', 'https://example.com/img1_thumb.jpg', ARRAY['nature:forest', 'color:green']),
('https://example.com/img2.jpg', 'https://example.com/img2_thumb.jpg', ARRAY['nature:ocean', 'color:blue'])
`)
if err != nil {
t.Fatalf("failed to create images: %v", err)
}
handler := http.HandlerFunc(ListFilters)
w := makeRequest(handler, "GET", "/api/portfolio/filters", nil)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
var filters []FilterCategory
if err := json.Unmarshal(w.Body.Bytes(), &filters); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
if len(filters) == 0 {
t.Error("expected filters, got empty")
}
}
func TestPortfolio_ListFilters_Empty(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
handler := http.HandlerFunc(ListFilters)
w := makeRequest(handler, "GET", "/api/portfolio/filters", nil)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
var filters []FilterCategory
if err := json.Unmarshal(w.Body.Bytes(), &filters); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
if len(filters) != 0 {
t.Errorf("expected 0 filters, got %d", len(filters))
}
}
// =============================================================================
// Get Image Tests
// =============================================================================
func TestPortfolio_GetImage(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
// Insert test image
var imageID string
err := db.DB.QueryRow(context.Background(), `
INSERT INTO images (url, thumbnail_url, tag_names)
VALUES ('https://example.com/img1.jpg', 'https://example.com/img1_thumb.jpg', ARRAY['nature:forest'])
RETURNING id
`).Scan(&imageID)
if err != nil {
t.Fatalf("failed to create image: %v", err)
}
handler := http.HandlerFunc(GetImage)
w := makeRequest(handler, "GET", "/api/portfolio/images/"+imageID, nil)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
var img Image
if err := json.Unmarshal(w.Body.Bytes(), &img); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
if img.ID != imageID {
t.Errorf("expected image ID %s, got %s", imageID, img.ID)
}
}
func TestPortfolio_GetImage_NotFound(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
handler := http.HandlerFunc(GetImage)
w := makeRequest(handler, "GET", "/api/portfolio/images/nonexistent-id", nil)
if w.Code != http.StatusNotFound {
t.Errorf("expected status 404, got %d. body: %s", w.Code, w.Body.String())
}
}
// =============================================================================
// Upload Image Tests (Admin Only)
// =============================================================================
func TestPortfolio_Upload_Admin(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
// Create a minimal S3 client mock by setting it to nil (handler will check and return error)
// The handler requires S3 client, so we test the auth check first
// Since S3 client setup is complex, we test that admin gets past auth check
handler := http.HandlerFunc(UploadImage)
w := makeRequestWithContext(handler, "POST", "/api/portfolio/images", nil, "admin-001", "admin")
// Should not get 403 (forbidden), will get another error due to missing S3 or file
// The important thing is it's not 403 for admin
if w.Code == http.StatusForbidden {
t.Error("admin should not get 403 - admin access should be granted")
}
}
func TestPortfolio_Upload_NonAdmin(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
handler := http.HandlerFunc(UploadImage)
w := makeRequestWithContext(handler, "POST", "/api/portfolio/images", nil, "user-001", "verified_email")
if w.Code != http.StatusForbidden {
t.Errorf("expected status 403, got %d. body: %s", w.Code, w.Body.String())
}
}
func TestPortfolio_Upload_Unauthenticated(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
handler := http.HandlerFunc(UploadImage)
w := makeRequest(handler, "POST", "/api/portfolio/images", nil)
if w.Code != http.StatusUnauthorized {
t.Errorf("expected status 401, got %d. body: %s", w.Code, w.Body.String())
}
}
// =============================================================================
// Delete Image Tests (Admin Only)
// =============================================================================
func TestPortfolio_Delete_Admin(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
// Insert test image
var imageID string
err := db.DB.QueryRow(context.Background(), `
INSERT INTO images (url, thumbnail_url, tag_names)
VALUES ('https://example.com/img1.jpg', 'https://example.com/img1_thumb.jpg', ARRAY['nature:forest'])
RETURNING id
`).Scan(&imageID)
if err != nil {
t.Fatalf("failed to create image: %v", err)
}
handler := http.HandlerFunc(DeleteImage)
w := makeRequestWithContext(handler, "DELETE", "/api/portfolio/images/"+imageID, nil, "admin-001", "admin")
// Should not get 403 (forbidden) - will get error due to S3 client being nil
// but the important thing is admin auth passed
if w.Code == http.StatusForbidden {
t.Error("admin should not get 403 - admin access should be granted")
}
}
func TestPortfolio_Delete_NonAdmin(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
// Insert test image
var imageID string
err := db.DB.QueryRow(context.Background(), `
INSERT INTO images (url, thumbnail_url, tag_names)
VALUES ('https://example.com/img1.jpg', 'https://example.com/img1_thumb.jpg', ARRAY['nature:forest'])
RETURNING id
`).Scan(&imageID)
if err != nil {
t.Fatalf("failed to create image: %v", err)
}
handler := http.HandlerFunc(DeleteImage)
w := makeRequestWithContext(handler, "DELETE", "/api/portfolio/images/"+imageID, nil, "user-001", "verified_email")
if w.Code != http.StatusForbidden {
t.Errorf("expected status 403, got %d. body: %s", w.Code, w.Body.String())
}
}
func TestPortfolio_Delete_Unauthenticated(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
// Insert test image
var imageID string
err := db.DB.QueryRow(context.Background(), `
INSERT INTO images (url, thumbnail_url, tag_names)
VALUES ('https://example.com/img1.jpg', 'https://example.com/img1_thumb.jpg', ARRAY['nature:forest'])
RETURNING id
`).Scan(&imageID)
if err != nil {
t.Fatalf("failed to create image: %v", err)
}
handler := http.HandlerFunc(DeleteImage)
w := makeRequest(handler, "DELETE", "/api/portfolio/images/"+imageID, nil)
if w.Code != http.StatusUnauthorized {
t.Errorf("expected status 401, got %d. body: %s", w.Code, w.Body.String())
}
}
@@ -0,0 +1,524 @@
//go:build test
// +build test
package scheduling
import (
"bytes"
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"strconv"
"testing"
"crussell/db"
"crussell/mw"
"crussell/testutils/jwt"
"crussell/testutils/testdb"
"github.com/jackc/pgx/v5/pgxpool"
)
func setupTestDB(t *testing.T) func() {
t.Helper()
pool := testdb.Pool(t)
testdb.Migrate(t, pool)
originalDB := db.DB
db.DB = pool
jwt.Init()
// Seed default working hours
seedDefaultWorkingHours(t, pool)
return func() {
db.DB = originalDB
pool.Close()
}
}
func seedDefaultWorkingHours(t *testing.T, pool *pgxpool.Pool) {
t.Helper()
// Seed 7 days of working hours (Monday=0 to Sunday=6)
hours := []struct {
weekday int
startTime string
endTime string
isOpen bool
}{
{0, "09:00", "17:00", true}, // Monday
{1, "09:00", "17:00", true}, // Tuesday
{2, "09:00", "17:00", true}, // Wednesday
{3, "09:00", "17:00", true}, // Thursday
{4, "09:00", "17:00", true}, // Friday
{5, "10:00", "16:00", true}, // Saturday
{6, "00:00", "00:00", false}, // Sunday
}
for _, h := range hours {
_, err := pool.Exec(context.Background(), `
INSERT INTO working_hours (weekday, start_time, end_time, is_open)
VALUES ($1, $2, $3, $4)
ON CONFLICT (weekday) DO UPDATE SET start_time = $2, end_time = $3, is_open = $4
`, h.weekday, h.startTime, h.endTime, h.isOpen)
if err != nil {
t.Fatalf("failed to seed working hours: %v", err)
}
}
}
func makeRequest(handler http.HandlerFunc, method, path string, body interface{}) *httptest.ResponseRecorder {
var req *http.Request
if body != nil {
bodyBytes, _ := json.Marshal(body)
req = httptest.NewRequest(method, path, bytes.NewReader(bodyBytes))
req.Header.Set("Content-Type", "application/json")
} else {
req = httptest.NewRequest(method, path, nil)
}
w := httptest.NewRecorder()
handler.ServeHTTP(w, req)
return w
}
func makeAuthRequest(handler http.HandlerFunc, method, path, token string, body interface{}) *httptest.ResponseRecorder {
var req *http.Request
if body != nil {
bodyBytes, _ := json.Marshal(body)
req = httptest.NewRequest(method, path, bytes.NewReader(bodyBytes))
req.Header.Set("Content-Type", "application/json")
} else {
req = httptest.NewRequest(method, path, nil)
}
if token != "" {
req.Header.Set("Authorization", "Bearer "+token)
}
w := httptest.NewRecorder()
handler.ServeHTTP(w, req)
return w
}
// --- Tests for GetDefaultHours ---
func TestScheduling_GetDefaultHours(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
handler := http.HandlerFunc(GetDefaultHours)
w := makeRequest(handler, "GET", "/api/scheduling/default-hours", nil)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
var response []DefaultHours
if err := json.Unmarshal(w.Body.Bytes(), &response); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
if len(response) != 7 {
t.Errorf("expected 7 days of hours, got %d", len(response))
}
// Verify Monday (weekday 0) has our seeded hours
var monday *DefaultHours
for i := range response {
if response[i].Weekday == 0 {
monday = &response[i]
break
}
}
if monday == nil {
t.Fatal("expected Monday hours in response")
}
if monday.StartTime != "09:00" {
t.Errorf("expected Monday start time 09:00, got %s", monday.StartTime)
}
if monday.EndTime != "17:00" {
t.Errorf("expected Monday end time 17:00, got %s", monday.EndTime)
}
if !monday.IsOpen {
t.Error("expected Monday to be open")
}
}
// --- Tests for UpdateDefaultHours ---
func TestScheduling_UpdateDefaultHours_Admin(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
adminToken := jwt.GenerateAdminToken()
handler := http.HandlerFunc(UpdateDefaultHours)
newHours := []DefaultHours{
{Weekday: 0, StartTime: "08:00", EndTime: "18:00", IsOpen: true},
{Weekday: 1, StartTime: "08:00", EndTime: "18:00", IsOpen: true},
{Weekday: 2, StartTime: "08:00", EndTime: "18:00", IsOpen: true},
{Weekday: 3, StartTime: "08:00", EndTime: "18:00", IsOpen: true},
{Weekday: 4, StartTime: "08:00", EndTime: "18:00", IsOpen: true},
{Weekday: 5, StartTime: "09:00", EndTime: "17:00", IsOpen: true},
{Weekday: 6, StartTime: "00:00", EndTime: "00:00", IsOpen: false},
}
w := makeAuthRequest(handler, "PUT", "/api/scheduling/default-hours", adminToken, newHours)
if w.Code != http.StatusNoContent {
t.Errorf("expected status 204, got %d. body: %s", w.Code, w.Body.String())
}
// Verify the update persisted
var hours []DefaultHours
rows, err := db.DB.Query(context.Background(), `SELECT weekday, start_time::text, end_time::text, is_open FROM working_hours ORDER BY weekday`)
if err != nil {
t.Fatalf("failed to query hours: %v", err)
}
defer rows.Close()
for rows.Next() {
var h DefaultHours
if err := rows.Scan(&h.Weekday, &h.StartTime, &h.EndTime, &h.IsOpen); err != nil {
t.Fatalf("failed to scan hours: %v", err)
}
hours = append(hours, h)
}
if hours[0].StartTime != "08:00" {
t.Errorf("expected Monday start time 08:00, got %s", hours[0].StartTime)
}
}
func TestScheduling_UpdateDefaultHours_NonAdmin(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
userToken := jwt.GenerateUserToken("user-123")
handler := http.HandlerFunc(UpdateDefaultHours)
newHours := []DefaultHours{
{Weekday: 0, StartTime: "08:00", EndTime: "18:00", IsOpen: true},
{Weekday: 1, StartTime: "08:00", EndTime: "18:00", IsOpen: true},
{Weekday: 2, StartTime: "08:00", EndTime: "18:00", IsOpen: true},
{Weekday: 3, StartTime: "08:00", EndTime: "18:00", IsOpen: true},
{Weekday: 4, StartTime: "08:00", EndTime: "18:00", IsOpen: true},
{Weekday: 5, StartTime: "09:00", EndTime: "17:00", IsOpen: true},
{Weekday: 6, StartTime: "00:00", EndTime: "00:00", IsOpen: false},
}
// Wrap handler with RequireAdmin middleware
w := makeAuthRequest(mw.RequireAdmin(UpdateDefaultHours), "PUT", "/api/scheduling/default-hours", userToken, newHours)
if w.Code != http.StatusForbidden {
t.Errorf("expected status 403, got %d. body: %s", w.Code, w.Body.String())
}
}
// --- Tests for ListExceptionalGroups ---
func TestScheduling_ListExceptionalGroups(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
// Create an exceptional group
_, err := db.DB.Exec(context.Background(), `
INSERT INTO exceptional_working_hours_groups (name, description)
VALUES ('Holiday Hours', 'Christmas holiday schedule')
`)
if err != nil {
t.Fatalf("failed to create group: %v", err)
}
handler := http.HandlerFunc(ListExceptionalGroups)
w := makeRequest(handler, "GET", "/api/scheduling/exceptional-groups", nil)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
var response []ExceptionalGroup
if err := json.Unmarshal(w.Body.Bytes(), &response); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
if len(response) == 0 {
t.Error("expected at least one group in response")
}
if response[0].Name != "Holiday Hours" {
t.Errorf("expected group name 'Holiday Hours', got %s", response[0].Name)
}
}
// --- Tests for CreateExceptionalGroup ---
func TestScheduling_CreateExceptionalGroup_Admin(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
adminToken := jwt.GenerateAdminToken()
handler := http.HandlerFunc(CreateExceptionalGroup)
newGroup := ExceptionalGroup{
Name: "Summer Hours",
Description: "Extended summer schedule",
Hours: []ExceptionalHours{
{Weekday: 0, StartTime: "08:00", EndTime: "18:00", IsOpen: true},
{Weekday: 1, StartTime: "08:00", EndTime: "18:00", IsOpen: true},
{Weekday: 2, StartTime: "08:00", EndTime: "18:00", IsOpen: true},
{Weekday: 3, StartTime: "08:00", EndTime: "18:00", IsOpen: true},
{Weekday: 4, StartTime: "08:00", EndTime: "18:00", IsOpen: true},
{Weekday: 5, StartTime: "09:00", EndTime: "17:00", IsOpen: true},
{Weekday: 6, StartTime: "00:00", EndTime: "00:00", IsOpen: false},
},
WeekStarts: []string{"2026-06-01"},
}
w := makeAuthRequest(handler, "POST", "/api/scheduling/exceptional-groups", adminToken, newGroup)
if w.Code != http.StatusCreated {
t.Errorf("expected status 201, got %d. body: %s", w.Code, w.Body.String())
}
var response ExceptionalGroup
if err := json.Unmarshal(w.Body.Bytes(), &response); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
if response.Name != "Summer Hours" {
t.Errorf("expected group name 'Summer Hours', got %s", response.Name)
}
if len(response.Hours) != 7 {
t.Errorf("expected 7 hours, got %d", len(response.Hours))
}
}
func TestScheduling_CreateExceptionalGroup_NonAdmin(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
userToken := jwt.GenerateUserToken("user-123")
handler := http.HandlerFunc(CreateExceptionalGroup)
newGroup := ExceptionalGroup{
Name: "Summer Hours",
Description: "Extended summer schedule",
Hours: []ExceptionalHours{
{Weekday: 0, StartTime: "08:00", EndTime: "18:00", IsOpen: true},
{Weekday: 1, StartTime: "08:00", EndTime: "18:00", IsOpen: true},
{Weekday: 2, StartTime: "08:00", EndTime: "18:00", IsOpen: true},
{Weekday: 3, StartTime: "08:00", EndTime: "18:00", IsOpen: true},
{Weekday: 4, StartTime: "08:00", EndTime: "18:00", IsOpen: true},
{Weekday: 5, StartTime: "09:00", EndTime: "17:00", IsOpen: true},
{Weekday: 6, StartTime: "00:00", EndTime: "00:00", IsOpen: false},
},
WeekStarts: []string{"2026-06-01"},
}
w := makeAuthRequest(mw.RequireAdmin(CreateExceptionalGroup), "POST", "/api/scheduling/exceptional-groups", userToken, newGroup)
if w.Code != http.StatusForbidden {
t.Errorf("expected status 403, got %d. body: %s", w.Code, w.Body.String())
}
}
// --- Tests for DeleteExceptionalGroup ---
func TestScheduling_DeleteExceptionalGroup_Admin(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
adminToken := jwt.GenerateAdminToken()
// Create a group to delete
var groupID int
err := db.DB.QueryRow(context.Background(), `
INSERT INTO exceptional_working_hours_groups (name, description)
VALUES ('To Delete', 'Will be deleted')
RETURNING id
`).Scan(&groupID)
if err != nil {
t.Fatalf("failed to create group: %v", err)
}
handler := http.HandlerFunc(DeleteExceptionalGroup)
req := httptest.NewRequest("DELETE", "/api/scheduling/exceptional-groups?id="+string(rune(groupID+'0')), nil)
req.Header.Set("Authorization", "Bearer "+adminToken)
w := httptest.NewRecorder()
handler.ServeHTTP(w, req)
// The handler expects id as query param but as a proper int
// Let's use proper URL query
req = httptest.NewRequest("DELETE", "/api/scheduling/exceptional-groups?id="+strconv.Itoa(groupID), nil)
req.Header.Set("Authorization", "Bearer "+adminToken)
w = httptest.NewRecorder()
handler.ServeHTTP(w, req)
if w.Code != http.StatusNoContent {
t.Errorf("expected status 204, got %d. body: %s", w.Code, w.Body.String())
}
// Verify group was deleted
var count int
err = db.DB.QueryRow(context.Background(), `SELECT COUNT(*) FROM exceptional_working_hours_groups WHERE id = $1`, groupID).Scan(&count)
if err != nil {
t.Fatalf("failed to check group: %v", err)
}
if count != 0 {
t.Error("expected group to be deleted")
}
}
func TestScheduling_DeleteExceptionalGroup_NonAdmin(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
userToken := jwt.GenerateUserToken("user-123")
handler := mw.RequireAdmin(DeleteExceptionalGroup)
req := httptest.NewRequest("DELETE", "/api/scheduling/exceptional-groups?id=1", nil)
req.Header.Set("Authorization", "Bearer "+userToken)
w := httptest.NewRecorder()
handler.ServeHTTP(w, req)
if w.Code != http.StatusForbidden {
t.Errorf("expected status 403, got %d. body: %s", w.Code, w.Body.String())
}
}
// --- Tests for GetWorkingHours ---
func TestScheduling_GetWorkingHours(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
handler := http.HandlerFunc(GetWorkingHours)
req := httptest.NewRequest("GET", "/api/scheduling/working-hours?start=2026-02-16&end=2026-02-22", nil)
w := httptest.NewRecorder()
handler.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
var response []DayWorkingHours
if err := json.Unmarshal(w.Body.Bytes(), &response); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
if len(response) == 0 {
t.Error("expected working hours in response")
}
// Verify source is "default" for seeded hours
for _, day := range response {
if day.Source != "default" {
t.Errorf("expected source 'default', got %s", day.Source)
}
break
}
}
// --- Tests for GetAvailableHours ---
func TestScheduling_GetAvailableHours(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
handler := http.HandlerFunc(GetAvailableHours)
req := httptest.NewRequest("GET", "/api/scheduling/available-hours?start=2026-02-16&end=2026-02-22", nil)
w := httptest.NewRecorder()
handler.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
var response []DayAvailableHours
if err := json.Unmarshal(w.Body.Bytes(), &response); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
if len(response) == 0 {
t.Error("expected available hours in response")
}
// Verify we have slots for open days
for _, day := range response {
if day.IsOpen {
if len(day.Slots) == 0 {
t.Error("expected slots for open days")
}
break
}
}
}
// --- Tests for UpdateExceptionalApplications ---
func TestScheduling_UpdateExceptionalApplications_Admin(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
adminToken := jwt.GenerateAdminToken()
// Create a group
var groupID int
err := db.DB.QueryRow(context.Background(), `
INSERT INTO exceptional_working_hours_groups (name, description)
VALUES ('Test Group', 'Test')
RETURNING id
`).Scan(&groupID)
if err != nil {
t.Fatalf("failed to create group: %v", err)
}
handler := http.HandlerFunc(UpdateExceptionalApplications)
reqBody := map[string]interface{}{
"groupId": groupID,
"weekStarts": []string{"2026-03-02", "2026-03-09"},
}
w := makeAuthRequest(handler, "PUT", "/api/scheduling/exceptional-applications", adminToken, reqBody)
if w.Code != http.StatusNoContent {
t.Errorf("expected status 204, got %d. body: %s", w.Code, w.Body.String())
}
// Verify applications were created
var count int
err = db.DB.QueryRow(context.Background(), `
SELECT COUNT(*) FROM exceptional_group_applications WHERE group_id = $1
`, groupID).Scan(&count)
if err != nil {
t.Fatalf("failed to check applications: %v", err)
}
if count != 2 {
t.Errorf("expected 2 applications, got %d", count)
}
}
func TestScheduling_UpdateExceptionalApplications_NonAdmin(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
userToken := jwt.GenerateUserToken("user-123")
handler := mw.RequireAdmin(UpdateExceptionalApplications)
reqBody := map[string]interface{}{
"groupId": 1,
"weekStarts": []string{"2026-03-02"},
}
w := makeAuthRequest(handler, "PUT", "/api/scheduling/exceptional-applications", userToken, reqBody)
if w.Code != http.StatusForbidden {
t.Errorf("expected status 403, got %d. body: %s", w.Code, w.Body.String())
}
}
+276
View File
@@ -0,0 +1,276 @@
//go:build test
// +build test
package services
import (
"bytes"
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
"crussell/db"
"crussell/handlers/user"
"crussell/testutils/jwt"
"crussell/testutils/testdb"
)
func setupTestDB(t *testing.T) func() {
t.Helper()
pool := testdb.Pool(t)
testdb.Migrate(t, pool)
originalDB := db.DB
db.DB = pool
jwt.Init()
return func() {
db.DB = originalDB
pool.Close()
}
}
func makeRequest(handler http.HandlerFunc, method, path string, body interface{}) *httptest.ResponseRecorder {
var req *http.Request
if body != nil {
bodyBytes, _ := json.Marshal(body)
req = httptest.NewRequest(method, path, bytes.NewReader(bodyBytes))
req.Header.Set("Content-Type", "application/json")
} else {
req = httptest.NewRequest(method, path, nil)
}
w := httptest.NewRecorder()
handler.ServeHTTP(w, req)
return w
}
func TestServices_ListAll(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
_, err := db.DB.Exec(context.Background(), `
INSERT INTO services (name, description, price, duration_minutes, is_active, patch_test_duration_hours, minimum_age_required)
VALUES
('Manicure', 'Basic manicure', 25.00, 30, true, 0, 0),
('Pedicure', 'Basic pedicure', 30.00, 45, true, 0, 0),
('Inactive Service', 'Should not appear', 50.00, 60, false, 0, 0)
`)
if err != nil {
t.Fatalf("failed to create services: %v", err)
}
handler := http.HandlerFunc(ServicesHandler)
w := makeRequest(handler, "GET", "/api/services", nil)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
var response []ServiceResponse
if err := json.Unmarshal(w.Body.Bytes(), &response); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
if len(response) != 2 {
t.Errorf("expected 2 services, got %d", len(response))
}
found := map[string]bool{}
for _, s := range response {
found[s.Name] = true
}
if !found["Manicure"] {
t.Error("expected Manicure in response")
}
if !found["Pedicure"] {
t.Error("expected Pedicure in response")
}
if found["Inactive Service"] {
t.Error("should not include inactive service")
}
}
func TestServices_EligibleForUser_AgeFilter(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
dob := "2005-01-01"
userID, err := createUserWithDOB(dob)
if err != nil {
t.Fatalf("failed to create user: %v", err)
}
_, err = db.DB.Exec(context.Background(), `
INSERT INTO services (name, description, price, duration_minutes, is_active, patch_test_duration_hours, minimum_age_required)
VALUES
('Under 18 Service', 'For minors', 20.00, 30, true, 0, 16),
('Adult Only Service', 'For adults only', 50.00, 60, true, 0, 21),
('No Age Restriction', 'Everyone welcome', 30.00, 45, true, 0, 0)
`)
if err != nil {
t.Fatalf("failed to create services: %v", err)
}
handler := http.HandlerFunc(ServicesEligibleForUserHandler)
req := httptest.NewRequest("GET", "/api/services/eligible-for/"+userID, nil)
w := httptest.NewRecorder()
handler.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
var response []ServiceResponse
if err := json.Unmarshal(w.Body.Bytes(), &response); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
if len(response) != 2 {
t.Errorf("expected 2 services, got %d. Response: %s", len(response), w.Body.String())
}
found := map[string]bool{}
for _, s := range response {
found[s.Name] = true
}
if !found["Under 18 Service"] {
t.Error("expected Under 18 Service in response (age 20 >= 16)")
}
if !found["No Age Restriction"] {
t.Error("expected No Age Restriction in response")
}
if found["Adult Only Service"] {
t.Error("should not include Adult Only Service (age 20 < 21)")
}
}
func TestServices_EligibleForUser_PatchTest(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
dob := "2000-01-01"
userID, err := createUserWithDOB(dob)
if err != nil {
t.Fatalf("failed to create user: %v", err)
}
_, err = db.DB.Exec(context.Background(), `
INSERT INTO services (name, description, price, duration_minutes, is_active, patch_test_duration_hours, minimum_age_required)
VALUES
('Regular Service', 'No patch test needed', 30.00, 30, true, 0, 0),
('Patch Test Required', 'Requires patch test', 75.00, 60, true, 48, 0)
`)
if err != nil {
t.Fatalf("failed to create services: %v", err)
}
var patchTestSvcID string
err = db.DB.QueryRow(context.Background(), "SELECT id FROM services WHERE name = 'Patch Test Required'").Scan(&patchTestSvcID)
if err != nil {
t.Fatalf("failed to get patch test service: %v", err)
}
_, err = db.DB.Exec(context.Background(), `
INSERT INTO user_service_patch_tests (user_id, service_id, last_time)
VALUES ($1, $2, NOW() - INTERVAL '24 hours')
`, userID, patchTestSvcID)
if err != nil {
t.Fatalf("failed to create patch test record: %v", err)
}
handler := http.HandlerFunc(ServicesEligibleForUserHandler)
req := httptest.NewRequest("GET", "/api/services/eligible-for/"+userID, nil)
w := httptest.NewRecorder()
handler.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
var response []ServiceResponse
if err := json.Unmarshal(w.Body.Bytes(), &response); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
if len(response) != 2 {
t.Errorf("expected 2 services, got %d. Response: %s", len(response), w.Body.String())
}
var patchTestSvc *ServiceResponse
var regularSvc *ServiceResponse
for i := range response {
if response[i].Name == "Patch Test Required" {
patchTestSvc = &response[i]
}
if response[i].Name == "Regular Service" {
regularSvc = &response[i]
}
}
if patchTestSvc == nil {
t.Fatal("Patch Test Required service not found in response")
}
if patchTestSvc.PatchTestStatus == nil || *patchTestSvc.PatchTestStatus != "ok" {
t.Errorf("expected patch test status 'ok', got %v", patchTestSvc.PatchTestStatus)
}
if regularSvc == nil {
t.Fatal("Regular Service not found in response")
}
if regularSvc.PatchTestStatus != nil {
t.Errorf("expected no patch test status for regular service, got %v", regularSvc.PatchTestStatus)
}
}
func TestContact_ReturnsInfo(t *testing.T) {
cleanup := setupTestDB(t)
defer cleanup()
_, err := db.DB.Exec(context.Background(), `
INSERT INTO users (n_first_name, n_last_name, email, phone, password_hash, account_role, account_type)
VALUES ('John', 'Smith', 'john@test.com', '07700900001', 'hash', 'admin', 'email')
`)
if err != nil {
t.Fatalf("failed to create admin user: %v", err)
}
handler := http.HandlerFunc(user.GetContactInfoHandler)
w := makeRequest(handler, "GET", "/api/contact", nil)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
var response user.ContactInfo
if err := json.Unmarshal(w.Body.Bytes(), &response); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
if response.Name == "" {
t.Error("expected name in response")
}
if response.Phone == "" {
t.Error("expected phone in response")
}
if response.Email == "" {
t.Error("expected email in response")
}
if response.Role == "" {
t.Error("expected role in response")
}
}
func createUserWithDOB(dob string) (string, error) {
ctx := context.Background()
var userID string
err := db.DB.QueryRow(ctx, `
INSERT INTO users (n_first_name, n_last_name, email, password_hash, account_role, account_type, date_of_birth)
VALUES ($1, $2, $3, $4, $5, $6, $7)
RETURNING id
`, "Test", "User", "testuser@test.com", "hash", "verified_email", "email", dob).Scan(&userID)
return userID, err
}
+242
View File
@@ -0,0 +1,242 @@
//go:build test
// +build test
package user
import (
"bytes"
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
"crussell/db"
"crussell/mw"
"crussell/testutils/fixtures"
"crussell/testutils/jwt"
"crussell/testutils/testdb"
"github.com/jackc/pgx/v5/pgxpool"
)
func setupTest(t *testing.T) (func(), *pgxpool.Pool) {
pool := testdb.Pool(t)
testdb.TruncateTables(t, pool)
// Set the global DB pool
db.DB = pool
// Initialize JWT
jwt.Init()
return func() {
pool.Close()
}, pool
}
func TestProfile_Get(t *testing.T) {
cleanup, pool := setupTest(t)
defer cleanup()
userID, err := fixtures.CreateTestUser(pool)
if err != nil {
t.Fatalf("failed to create test user: %v", err)
}
token := jwt.GenerateUserToken(userID)
req := httptest.NewRequest(http.MethodGet, "/api/user/profile", nil)
req = req.WithContext(context.WithValue(context.Background(), mw.UserIDKey, userID))
req.Header.Set("Authorization", "Bearer "+token)
rr := httptest.NewRecorder()
GetProfileHandler(rr, req)
if rr.Code != http.StatusOK {
t.Errorf("expected status 200, got %d", rr.Code)
t.Logf("response body: %s", rr.Body.String())
}
var profile UserProfile
if err := json.Unmarshal(rr.Body.Bytes(), &profile); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
if profile.ID != userID {
t.Errorf("expected user ID %s, got %s", userID, profile.ID)
}
}
func TestProfile_Get_NoAuth(t *testing.T) {
cleanup, _ := setupTest(t)
defer cleanup()
req := httptest.NewRequest(http.MethodGet, "/api/user/profile", nil)
rr := httptest.NewRecorder()
GetProfileHandler(rr, req)
if rr.Code != http.StatusUnauthorized {
t.Errorf("expected status 401, got %d", rr.Code)
}
}
func TestProfile_Update(t *testing.T) {
cleanup, pool := setupTest(t)
defer cleanup()
userID, err := fixtures.CreateTestUser(pool)
if err != nil {
t.Fatalf("failed to create test user: %v", err)
}
token := jwt.GenerateUserToken(userID)
updateReq := UpdateProfileRequest{
FirstName: "John",
LastName: "Doe",
Phone: "07123456789",
}
body, _ := json.Marshal(updateReq)
req := httptest.NewRequest(http.MethodPut, "/api/user/profile", bytes.NewReader(body))
req = req.WithContext(context.WithValue(context.Background(), mw.UserIDKey, userID))
req.Header.Set("Authorization", "Bearer "+token)
req.Header.Set("Content-Type", "application/json")
rr := httptest.NewRecorder()
UpdateProfileHandler(rr, req)
if rr.Code != http.StatusOK {
t.Errorf("expected status 200, got %d", rr.Code)
t.Logf("response body: %s", rr.Body.String())
}
}
func TestPasswordChange_Success(t *testing.T) {
cleanup, pool := setupTest(t)
defer cleanup()
userID, err := fixtures.CreateTestUser(pool)
if err != nil {
t.Fatalf("failed to create test user: %v", err)
}
token := jwt.GenerateUserToken(userID)
changeReq := ChangePasswordRequest{
CurrentPassword: "testpassword123",
NewPassword: "newpassword456",
}
body, _ := json.Marshal(changeReq)
req := httptest.NewRequest(http.MethodPut, "/api/user/change-password", bytes.NewReader(body))
req = req.WithContext(context.WithValue(context.Background(), mw.UserIDKey, userID))
req.Header.Set("Authorization", "Bearer "+token)
req.Header.Set("Content-Type", "application/json")
rr := httptest.NewRecorder()
ChangePasswordHandler(rr, req)
if rr.Code != http.StatusOK {
t.Errorf("expected status 200, got %d", rr.Code)
t.Logf("response body: %s", rr.Body.String())
}
}
func TestPasswordChange_WrongOld(t *testing.T) {
cleanup, pool := setupTest(t)
defer cleanup()
userID, err := fixtures.CreateTestUser(pool)
if err != nil {
t.Fatalf("failed to create test user: %v", err)
}
token := jwt.GenerateUserToken(userID)
changeReq := ChangePasswordRequest{
CurrentPassword: "wrongpassword",
NewPassword: "newpassword456",
}
body, _ := json.Marshal(changeReq)
req := httptest.NewRequest(http.MethodPut, "/api/user/change-password", bytes.NewReader(body))
req = req.WithContext(context.WithValue(context.Background(), mw.UserIDKey, userID))
req.Header.Set("Authorization", "Bearer "+token)
req.Header.Set("Content-Type", "application/json")
rr := httptest.NewRecorder()
ChangePasswordHandler(rr, req)
if rr.Code != http.StatusUnauthorized {
t.Errorf("expected status 401, got %d", rr.Code)
}
}
func TestAccount_Delete(t *testing.T) {
cleanup, pool := setupTest(t)
defer cleanup()
userID, err := fixtures.CreateTestUser(pool)
if err != nil {
t.Fatalf("failed to create test user: %v", err)
}
token := jwt.GenerateUserToken(userID)
req := httptest.NewRequest(http.MethodDelete, "/api/user/account", nil)
req = req.WithContext(context.WithValue(context.Background(), mw.UserIDKey, userID))
req.Header.Set("Authorization", "Bearer "+token)
rr := httptest.NewRecorder()
DeleteAccountHandler(rr, req)
if rr.Code != http.StatusNoContent {
t.Errorf("expected status 204, got %d", rr.Code)
t.Logf("response body: %s", rr.Body.String())
}
}
func TestLoyalty_Get(t *testing.T) {
cleanup, pool := setupTest(t)
defer cleanup()
userID, err := fixtures.CreateTestUser(pool)
if err != nil {
t.Fatalf("failed to create test user: %v", err)
}
// Add some loyalty stamps
_, err = pool.Exec(context.Background(), `UPDATE users SET loyalty_stamps = 10 WHERE id = $1`, userID)
if err != nil {
t.Fatalf("failed to update loyalty stamps: %v", err)
}
token := jwt.GenerateUserToken(userID)
req := httptest.NewRequest(http.MethodGet, "/api/user/loyalty", nil)
req = req.WithContext(context.WithValue(context.Background(), mw.UserIDKey, userID))
req.Header.Set("Authorization", "Bearer "+token)
rr := httptest.NewRecorder()
GetLoyaltyHandler(rr, req)
if rr.Code != http.StatusOK {
t.Errorf("expected status 200, got %d", rr.Code)
t.Logf("response body: %s", rr.Body.String())
}
var loyalty LoyaltyResponse
if err := json.Unmarshal(rr.Body.Bytes(), &loyalty); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
if loyalty.Stamps != 10 {
t.Errorf("expected 10 stamps, got %d", loyalty.Stamps)
}
if loyalty.ReferralCode == "" {
t.Error("expected referral code to be set")
}
}