feat: Square payment integration, booking flow redesign, and timezone/weekday fixes

- Add Square payment integration (mock + handlers + UI): terminal/online payments,
  refunds, tips, saved cards, webhooks. Build-tagged dev/prod clients.
- Redesign booking flow: Step 4 conditional (deposit only), Step 5 confirmation
  screen with booking ID, auto-submit on transition.
- Redesign schedule modal: 2x3 button grid with Pay Deposit/Pay Early logic.
- Add deposit warning banner at Step 1 for users with outstanding deposits.
- Fix weekday conversion bug: Go 0=Sunday vs DB 0=Monday mismatch in 6 locations.
- Fix timezone bug: UTC vs London time in closing hours validation.
- Fix frontend error parsing: plain text backend errors now displayed correctly.
- Fix crypto.randomUUID fallback for environments without Web Crypto.
- Add 7 new regression tests: closing hours, advance check, active booking limit,
  weekday conversion, UTC/London, deposit snapshot, exceptional hours.
- Fix 3 flaky tests: dynamic dates instead of fixed, no-show timing.
This commit is contained in:
2026-05-23 11:29:34 +01:00
parent bcd5ed2bd9
commit 2e1ab9d745
31 changed files with 6155 additions and 229 deletions
+201 -5
View File
@@ -1223,21 +1223,23 @@ func TestAdminBookings_Create_DuringHolidayHours_Rejected(t *testing.T) {
defer db.DB.Exec(context.Background(), "DELETE FROM exceptional_working_hours_groups WHERE id = $1", groupID)
// Add closed hours for targetDate (closed all day)
// DB convention: 0=Monday..6=Sunday; Go: 0=Sunday..6=Saturday. Convert.
dbWeekday := (int(targetDate.Weekday()) + 6) % 7
_, err = db.DB.Exec(context.Background(), `
INSERT INTO exceptional_working_hours (group_id, weekday, start_time, end_time, is_open)
VALUES ($1, $2, $3, $4, $5)
`, groupID, int(targetDate.Weekday()), "00:00:00", "23:59:59", false)
`, groupID, dbWeekday, "00:00:00", "23:59:59", false)
if err != nil {
t.Fatalf("failed to create holiday hours: %v", err)
}
// Apply the group to the week containing targetDate
// Must use Monday of that week (matching handler logic)
targetWeekday := int(targetDate.Weekday())
if targetWeekday == 0 {
targetWeekday = 7 // Sunday -> 7
daysToMonday := int(targetDate.Weekday())
if daysToMonday == 0 {
daysToMonday = 7
}
mondayOfWeek := targetDate.AddDate(0, 0, -targetWeekday+1)
mondayOfWeek := targetDate.AddDate(0, 0, -daysToMonday+1)
_, err = db.DB.Exec(context.Background(), `
INSERT INTO exceptional_group_applications (group_id, week_start)
VALUES ($1, $2)
@@ -2660,3 +2662,197 @@ func TestAdminBookings_Create_WalkInGuestUser(t *testing.T) {
t.Errorf("expected booking user_id = %s, got %s", guestID, foundUserID)
}
}
// =============================================================================
// Exceptional Hours Tests
// =============================================================================
// TestAdminUpdateBooking_ClosedExceptionalHours_WarningOnly verifies that an admin can
// update a booking's time to fall within a closed exceptional hours period, receiving
// a warning but proceeding with the update.
func TestAdminUpdateBooking_ClosedExceptionalHours_WarningOnly(t *testing.T) {
resetTestData(t)
adminID, err := fixtures.CreateTestAdminUser(db.DB)
if err != nil {
t.Fatalf("failed to create admin user: %v", err)
}
defer fixtures.DeleteUser(db.DB, adminID)
userID, err := fixtures.CreateTestUser(db.DB)
if err != nil {
t.Fatalf("failed to create test user: %v", err)
}
defer fixtures.DeleteUser(db.DB, userID)
serviceID, err := fixtures.CreateTestService(db.DB)
if err != nil {
t.Fatalf("failed to create test service: %v", err)
}
defer fixtures.DeleteService(db.DB, serviceID)
originalTime := time.Now().Add(72 * time.Hour).Truncate(time.Second)
bookingID, err := fixtures.CreateTestBooking(db.DB, userID, serviceID)
if err != nil {
t.Fatalf("failed to create test booking: %v", err)
}
defer fixtures.DeleteBooking(db.DB, bookingID)
_, err = db.DB.Exec(context.Background(),
"UPDATE bookings SET start_time = $1 WHERE id = $2", originalTime, bookingID)
if err != nil {
t.Fatalf("failed to update booking time: %v", err)
}
targetDate := time.Now().Add(5 * 24 * time.Hour).Truncate(24 * time.Hour)
var groupID int
err = db.DB.QueryRow(context.Background(), `
INSERT INTO exceptional_working_hours_groups (name, description)
VALUES ($1, $2)
RETURNING id
`, "Holiday Closure", "Test holiday").Scan(&groupID)
if err != nil {
t.Fatalf("failed to create holiday group: %v", err)
}
// DB convention: 0=Monday..6=Sunday; Go: 0=Sunday..6=Saturday. Convert.
dbWeekday := (int(targetDate.Weekday()) + 6) % 7
_, err = db.DB.Exec(context.Background(), `
INSERT INTO exceptional_working_hours (group_id, weekday, start_time, end_time, is_open)
VALUES ($1, $2, $3, $4, $5)
`, groupID, dbWeekday, "00:00:00", "23:59:59", false)
if err != nil {
t.Fatalf("failed to create holiday hours: %v", err)
}
daysToMonday := int(targetDate.Weekday())
if daysToMonday == 0 {
daysToMonday = 7
}
mondayOfWeek := targetDate.AddDate(0, 0, -daysToMonday+1)
_, err = db.DB.Exec(context.Background(), `
INSERT INTO exceptional_group_applications (group_id, week_start)
VALUES ($1, $2)
`, groupID, mondayOfWeek)
if err != nil {
t.Fatalf("failed to create holiday application: %v", err)
}
targetTime := targetDate.Add(14 * time.Hour).Truncate(time.Second)
req := bookings.EditBookingRequest{
StartTime: targetTime,
}
handler := http.HandlerFunc(bookings.AdminEditBookingHandler)
w := makeAdminRequest(handler, "PUT", "/api/admin/bookings/"+bookingID, req)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
var response map[string]interface{}
if err := parseResponseBody(w, &response); err != nil {
t.Fatalf("failed to parse response: %v", err)
}
warnings, ok := response["warnings"].([]interface{})
if !ok || len(warnings) == 0 {
t.Error("expected warnings array with at least one warning about working hours")
} else {
warningStr, ok := warnings[0].(string)
if !ok {
t.Errorf("expected warning to be a string, got: %v", warnings[0])
} else if !bytes.Contains([]byte(warningStr), []byte("working hours")) {
t.Errorf("expected warning to mention 'working hours', got: %s", warningStr)
}
}
var updatedStartTime time.Time
err = db.DB.QueryRow(context.Background(),
"SELECT start_time FROM bookings WHERE id = $1", bookingID).Scan(&updatedStartTime)
if err != nil {
t.Fatalf("failed to query booking: %v", err)
}
if !updatedStartTime.Equal(targetTime) {
t.Errorf("expected booking start_time %v, got %v", targetTime, updatedStartTime)
}
}
// TestAdminCreateBookingForUser_ClosedExceptionalHours_Rejected verifies that an admin
// cannot create a booking for a user during hours marked as closed in the exceptional
// working hours (holiday) system.
func TestAdminCreateBookingForUser_ClosedExceptionalHours_Rejected(t *testing.T) {
resetTestData(t)
adminID, err := fixtures.CreateTestAdminUser(db.DB)
if err != nil {
t.Fatalf("failed to create admin user: %v", err)
}
defer fixtures.DeleteUser(db.DB, adminID)
userID, err := fixtures.CreateTestUser(db.DB)
if err != nil {
t.Fatalf("failed to create test user: %v", err)
}
defer fixtures.DeleteUser(db.DB, userID)
serviceID, err := fixtures.CreateTestService(db.DB)
if err != nil {
t.Fatalf("failed to create test service: %v", err)
}
defer fixtures.DeleteService(db.DB, serviceID)
targetDate := time.Now().Add(5 * 24 * time.Hour).Truncate(24 * time.Hour)
var groupID int
err = db.DB.QueryRow(context.Background(), `
INSERT INTO exceptional_working_hours_groups (name, description)
VALUES ($1, $2)
RETURNING id
`, "Holiday Closure", "Test holiday").Scan(&groupID)
if err != nil {
t.Fatalf("failed to create holiday group: %v", err)
}
// DB convention: 0=Monday..6=Sunday; Go: 0=Sunday..6=Saturday. Convert.
dbWeekday := (int(targetDate.Weekday()) + 6) % 7
_, err = db.DB.Exec(context.Background(), `
INSERT INTO exceptional_working_hours (group_id, weekday, start_time, end_time, is_open)
VALUES ($1, $2, $3, $4, $5)
`, groupID, dbWeekday, "00:00:00", "23:59:59", false)
if err != nil {
t.Fatalf("failed to create holiday hours: %v", err)
}
daysToMonday := int(targetDate.Weekday())
if daysToMonday == 0 {
daysToMonday = 7
}
mondayOfWeek := targetDate.AddDate(0, 0, -daysToMonday+1)
_, err = db.DB.Exec(context.Background(), `
INSERT INTO exceptional_group_applications (group_id, week_start)
VALUES ($1, $2)
`, groupID, mondayOfWeek)
if err != nil {
t.Fatalf("failed to create holiday application: %v", err)
}
targetTime := targetDate.Add(14 * time.Hour).Truncate(time.Second)
req := bookings.AdminCreateBookingForUserRequest{
UserID: userID,
StartTime: targetTime,
ServiceIDs: []string{serviceID},
}
handler := http.HandlerFunc(bookings.AdminCreateBookingForUserHandler)
w := makeAdminRequest(handler, "POST", "/api/admin/bookings", req)
if w.Code != http.StatusConflict {
t.Errorf("expected status 409, got %d. body: %s", w.Code, w.Body.String())
}
if !bytes.Contains(w.Body.Bytes(), []byte("holiday hours")) {
t.Errorf("expected error message to mention 'holiday hours', got: %s", w.Body.String())
}
}
+5 -2
View File
@@ -104,7 +104,9 @@ func AdminReserveSlotHandler(w http.ResponseWriter, r *http.Request) {
}
}
weekday := int(req.StartTime.Weekday())
localStart := req.StartTime.In(londonLocation)
// DB uses 0=Monday..6=Sunday; Go uses 0=Sunday..6=Saturday. Convert.
weekday := int((localStart.Weekday() + 6) % 7)
var closeStr string
if err := db.DB.QueryRow(r.Context(), `SELECT end_time::text FROM working_hours WHERE weekday = $1`, weekday).Scan(&closeStr); err != nil {
if errors.Is(err, sql.ErrNoRows) {
@@ -117,8 +119,9 @@ func AdminReserveSlotHandler(w http.ResponseWriter, r *http.Request) {
}
endTime := req.StartTime.Add(time.Duration(svcDuration) * time.Minute)
localEnd := localStart.Add(time.Duration(svcDuration) * time.Minute)
closeTime, _ := time.Parse("15:04:05", closeStr)
if endTime.Hour() > closeTime.Hour() || (endTime.Hour() == closeTime.Hour() && endTime.Minute() > closeTime.Minute()) {
if localEnd.Hour() > closeTime.Hour() || (localEnd.Hour() == closeTime.Hour() && localEnd.Minute() > closeTime.Minute()) {
http.Error(w, "Cannot book this time - services would extend beyond closing hours", http.StatusBadRequest)
return
}
+8 -4
View File
@@ -1588,7 +1588,9 @@ func CreateBookingHandler(w http.ResponseWriter, r *http.Request) {
return
}
weekday := int(req.StartTime.Weekday())
localStart := req.StartTime.In(londonLocation)
// DB uses 0=Monday..6=Sunday; Go uses 0=Sunday..6=Saturday. Convert.
weekday := int((localStart.Weekday() + 6) % 7)
var closeStr string
if err := db.DB.QueryRow(r.Context(), `SELECT end_time::text FROM working_hours WHERE weekday = $1`, weekday).Scan(&closeStr); err != nil {
log.Printf("Failed to get hours: %v", err)
@@ -1597,8 +1599,9 @@ func CreateBookingHandler(w http.ResponseWriter, r *http.Request) {
}
endTime := req.StartTime.Add(time.Duration(svcDuration) * time.Minute)
localEnd := localStart.Add(time.Duration(svcDuration) * time.Minute)
closeTime, _ := time.Parse("15:04:05", closeStr)
if endTime.Hour() > closeTime.Hour() || (endTime.Hour() == closeTime.Hour() && endTime.Minute() > closeTime.Minute()) {
if localEnd.Hour() > closeTime.Hour() || (localEnd.Hour() == closeTime.Hour() && localEnd.Minute() > closeTime.Minute()) {
http.Error(w, "Cannot book this time - services would extend beyond closing hours", http.StatusBadRequest)
return
}
@@ -1832,9 +1835,10 @@ func EditBookingHandler(w http.ResponseWriter, r *http.Request) {
return
}
weekday := int(req.StartTime.Weekday())
// DB uses 0=Monday..6=Sunday; Go uses 0=Sunday..6=Saturday. Convert.
weekday := int((req.StartTime.Weekday() + 6) % 7)
bookingTime := req.StartTime.Format("15:04:05")
daysToMonday := weekday
daysToMonday := int(req.StartTime.Weekday())
if daysToMonday == 0 {
daysToMonday = 7
}
+363 -15
View File
@@ -1651,9 +1651,8 @@ func TestDeleteBooking_NoShowUnder24h_SetsDepositsTo3(t *testing.T) {
token := jwt.GenerateUserToken(userID)
// Create booking with start_time = now + 12 hours (< 24h notice)
soonTime := time.Now().Add(12 * time.Hour).Truncate(time.Second)
soonTime = time.Date(soonTime.Year(), soonTime.Month(), soonTime.Day(), 10, 0, 0, 0, soonTime.Location())
// Create booking with start_time = now + 23 hours (< 24h notice, > 1h advance)
soonTime := time.Now().Add(23 * time.Hour).Truncate(time.Second)
bookingReq := CreateBookingRequest{
StartTime: soonTime,
@@ -1818,9 +1817,8 @@ func TestDeleteBooking_NoShowWithForgiveness_NoPenalty(t *testing.T) {
token := jwt.GenerateUserToken(userID)
// Create booking with start_time = now + 12 hours (< 24h notice)
soonTime := time.Now().Add(12 * time.Hour).Truncate(time.Second)
soonTime = time.Date(soonTime.Year(), soonTime.Month(), soonTime.Day(), 10, 0, 0, 0, soonTime.Location())
// Create booking with start_time = now + 23 hours (< 24h notice, > 1h advance)
soonTime := time.Now().Add(23 * time.Hour).Truncate(time.Second)
bookingReq := CreateBookingRequest{
StartTime: soonTime,
@@ -1903,8 +1901,7 @@ func TestDeleteBooking_SecondNoShow_StaysAt3(t *testing.T) {
token := jwt.GenerateUserToken(userID)
// === First booking: no-show ===
soonTime1 := time.Now().Add(12 * time.Hour).Truncate(time.Second)
soonTime1 = time.Date(soonTime1.Year(), soonTime1.Month(), soonTime1.Day(), 10, 0, 0, 0, soonTime1.Location())
soonTime1 := time.Now().Add(23 * time.Hour).Truncate(time.Second)
bookingReq1 := CreateBookingRequest{
StartTime: soonTime1,
@@ -1944,10 +1941,7 @@ func TestDeleteBooking_SecondNoShow_StaysAt3(t *testing.T) {
}
// === Second booking: no-show ===
// Need to wait a bit or create with different time to avoid conflict
// Use tomorrow + 12 hours
soonTime2 := time.Now().Add(36 * time.Hour).Truncate(time.Second)
soonTime2 = time.Date(soonTime2.Year(), soonTime2.Month(), soonTime2.Day(), 10, 0, 0, 0, soonTime2.Location())
soonTime2 := time.Now().Add(47 * time.Hour).Truncate(time.Second)
bookingReq2 := CreateBookingRequest{
StartTime: soonTime2,
@@ -4149,8 +4143,9 @@ func TestBookings_Edit_ClosedDay_UserBlocked(t *testing.T) {
// Calculate week start for the booking target date
targetDate := time.Now().Add(96 * time.Hour)
weekday := int(targetDate.Weekday())
daysToMonday := weekday
// DB convention: 0=Monday..6=Sunday; Go: 0=Sunday..6=Saturday. Convert.
dbWeekday := (int(targetDate.Weekday()) + 6) % 7
daysToMonday := int(targetDate.Weekday())
if daysToMonday == 0 {
daysToMonday = 7
}
@@ -4169,7 +4164,7 @@ func TestBookings_Edit_ClosedDay_UserBlocked(t *testing.T) {
_, err = db.DB.Exec(context.Background(),
`INSERT INTO exceptional_working_hours (group_id, weekday, start_time, end_time, is_open)
VALUES ($1, $2, '08:00:00', '20:00:00', false)`,
groupID, weekday)
groupID, dbWeekday)
if err != nil {
t.Fatalf("failed to create closed exceptional hours: %v", err)
}
@@ -4841,3 +4836,356 @@ func TestCreateBooking_Notifications_NoPendingBookingWithoutNotes(t *testing.T)
t.Errorf("expected 0 pending_booking notifications for booking without notes, got %d", pendingCount)
}
}
// =============================================================================
// Closing Hours, Advance Check, and Active Booking Limit Tests
// =============================================================================
func TestCreateBooking_ClosingHoursValidation(t *testing.T) {
resetTestData(t)
london, err := time.LoadLocation("Europe/London")
if err != nil {
t.Fatalf("failed to load London timezone: %v", err)
}
hours := []struct {
weekday int
startTime string
endTime string
isOpen bool
}{
{0, "08:00", "17:00", true},
{1, "08:00", "20:00", true},
{2, "08:00", "20:00", true},
{3, "08:00", "20:00", true},
{4, "08:00", "20:00", true},
{5, "08:00", "20:00", true},
{6, "08:00", "20:00", true},
}
seedCustomWorkingHours(t, hours)
userID, err := fixtures.CreateTestUser(db.DB)
if err != nil {
t.Fatalf("failed to create test user: %v", err)
}
defer fixtures.DeleteUser(db.DB, userID)
_, err = db.DB.Exec(context.Background(), "UPDATE users SET deposits_required = 0 WHERE id = $1", userID)
if err != nil {
t.Fatalf("failed to set deposits_required: %v", err)
}
serviceID, err := fixtures.CreateTestService(db.DB)
if err != nil {
t.Fatalf("failed to create test service: %v", err)
}
defer fixtures.DeleteService(db.DB, serviceID)
_, err = db.DB.Exec(context.Background(), "UPDATE services SET duration_minutes = 60 WHERE id = $1", serviceID)
if err != nil {
t.Fatalf("failed to set service duration: %v", err)
}
token := jwt.GenerateUserToken(userID)
thursday := nextWeekday(time.Thursday, london)
thursdayStart := time.Date(thursday.Year(), thursday.Month(), thursday.Day(), 17, 30, 0, 0, london)
req1 := CreateBookingRequest{
StartTime: thursdayStart,
ServiceIDs: []string{serviceID},
}
handler := http.HandlerFunc(CreateBookingHandler)
w1 := makeRequest(handler, "POST", "/api/bookings", req1, token)
if w1.Code != http.StatusCreated {
t.Errorf("Thursday 17:30+60min should succeed (ends 18:30 < 20:00), got %d. body: %s", w1.Code, w1.Body.String())
}
monday := nextWeekday(time.Monday, london)
mondayStart := time.Date(monday.Year(), monday.Month(), monday.Day(), 16, 30, 0, 0, london)
req2 := CreateBookingRequest{
StartTime: mondayStart,
ServiceIDs: []string{serviceID},
}
w2 := makeRequest(handler, "POST", "/api/bookings", req2, token)
if w2.Code != http.StatusBadRequest {
t.Errorf("Monday 16:30+60min should fail (ends 17:30 > 17:00), got %d. body: %s", w2.Code, w2.Body.String())
}
if !bytes.Contains(w2.Body.Bytes(), []byte("closing")) {
t.Errorf("expected error about closing hours, got: %s", w2.Body.String())
}
}
func TestCreateBooking_OneHourAdvanceCheck(t *testing.T) {
resetTestData(t)
seedDefaultWorkingHours(t)
userID, err := fixtures.CreateTestUser(db.DB)
if err != nil {
t.Fatalf("failed to create test user: %v", err)
}
defer fixtures.DeleteUser(db.DB, userID)
_, err = db.DB.Exec(context.Background(), "UPDATE users SET deposits_required = 0 WHERE id = $1", userID)
if err != nil {
t.Fatalf("failed to set deposits_required: %v", err)
}
serviceID, err := fixtures.CreateTestService(db.DB)
if err != nil {
t.Fatalf("failed to create test service: %v", err)
}
defer fixtures.DeleteService(db.DB, serviceID)
token := jwt.GenerateUserToken(userID)
handler := http.HandlerFunc(CreateBookingHandler)
soonTime := time.Now().Add(30 * time.Minute).Truncate(time.Second)
req1 := CreateBookingRequest{
StartTime: soonTime,
ServiceIDs: []string{serviceID},
}
w1 := makeRequest(handler, "POST", "/api/bookings", req1, token)
if w1.Code != http.StatusBadRequest {
t.Errorf("expected status 400 for 30-min advance booking, got %d. body: %s", w1.Code, w1.Body.String())
}
if !bytes.Contains(w1.Body.Bytes(), []byte("at least 1 hour")) {
t.Errorf("expected error about 1 hour advance, got: %s", w1.Body.String())
}
aheadTime := time.Now().Add(72 * time.Hour).Truncate(time.Second)
aheadTime = time.Date(aheadTime.Year(), aheadTime.Month(), aheadTime.Day(), 10, 0, 0, 0, aheadTime.Location())
req2 := CreateBookingRequest{
StartTime: aheadTime,
ServiceIDs: []string{serviceID},
}
w2 := makeRequest(handler, "POST", "/api/bookings", req2, token)
if w2.Code != http.StatusCreated {
t.Errorf("expected status 201 for 2h+ advance booking, got %d. body: %s", w2.Code, w2.Body.String())
}
}
func TestCreateBooking_ActiveBookingLimit(t *testing.T) {
resetTestData(t)
seedDefaultWorkingHours(t)
userID, err := fixtures.CreateTestUser(db.DB)
if err != nil {
t.Fatalf("failed to create test user: %v", err)
}
defer fixtures.DeleteUser(db.DB, userID)
_, err = db.DB.Exec(context.Background(), "UPDATE users SET deposits_required = 3 WHERE id = $1", userID)
if err != nil {
t.Fatalf("failed to set deposits_required: %v", err)
}
serviceID, err := fixtures.CreateTestService(db.DB)
if err != nil {
t.Fatalf("failed to create test service: %v", err)
}
defer fixtures.DeleteService(db.DB, serviceID)
token := jwt.GenerateUserToken(userID)
handler := http.HandlerFunc(CreateBookingHandler)
firstTime := time.Now().Add(72 * time.Hour).Truncate(time.Second)
firstTime = time.Date(firstTime.Year(), firstTime.Month(), firstTime.Day(), 10, 0, 0, 0, firstTime.Location())
req1 := CreateBookingRequest{
StartTime: firstTime,
ServiceIDs: []string{serviceID},
}
w1 := makeRequest(handler, "POST", "/api/bookings", req1, token)
if w1.Code != http.StatusCreated {
t.Fatalf("expected first booking to succeed, got %d. body: %s", w1.Code, w1.Body.String())
}
var booking1 Booking
if err := parseResponseBody(w1, &booking1); err != nil {
t.Fatalf("failed to parse first booking: %v", err)
}
secondTime := time.Now().Add(96 * time.Hour).Truncate(time.Second)
secondTime = time.Date(secondTime.Year(), secondTime.Month(), secondTime.Day(), 14, 0, 0, 0, secondTime.Location())
req2 := CreateBookingRequest{
StartTime: secondTime,
ServiceIDs: []string{serviceID},
}
w2 := makeRequest(handler, "POST", "/api/bookings", req2, token)
if w2.Code != http.StatusConflict {
t.Errorf("expected status 409 for second booking with active booking, got %d. body: %s", w2.Code, w2.Body.String())
}
if !bytes.Contains(w2.Body.Bytes(), []byte("active booking")) {
t.Errorf("expected error about active booking, got: %s", w2.Body.String())
}
_, err = db.DB.Exec(context.Background(),
"UPDATE bookings SET status = 'client_cancelled' WHERE id = $1", booking1.ID)
if err != nil {
t.Fatalf("failed to cancel first booking: %v", err)
}
w3 := makeRequest(handler, "POST", "/api/bookings", req2, token)
if w3.Code != http.StatusCreated {
t.Errorf("expected status 201 after cancelling active booking, got %d. body: %s", w3.Code, w3.Body.String())
}
}
func TestNextWeekdayHelper(t *testing.T) {
london, err := time.LoadLocation("Europe/London")
if err != nil {
t.Fatalf("Europe/London not available: %v", err)
}
tests := []struct {
name string
weekday time.Weekday
}{
{"Monday", time.Monday},
{"Tuesday", time.Tuesday},
{"Wednesday", time.Wednesday},
{"Thursday", time.Thursday},
{"Friday", time.Friday},
{"Saturday", time.Saturday},
{"Sunday", time.Sunday},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
result := nextWeekday(tt.weekday, london)
if result.Weekday() != tt.weekday {
t.Errorf("expected weekday %s, got %s", tt.weekday, result.Weekday())
}
now := time.Now().In(london)
today := time.Date(now.Year(), now.Month(), now.Day(), 0, 0, 0, 0, london)
resultDay := time.Date(result.Year(), result.Month(), result.Day(), 0, 0, 0, 0, london)
daysDiff := int(resultDay.Sub(today).Hours() / 24)
if daysDiff < 2 {
t.Errorf("expected result to be at least 2 calendar days ahead, got %d", daysDiff)
}
})
}
}
func TestCreateBooking_DepositSnapshot(t *testing.T) {
resetTestData(t)
seedDefaultWorkingHours(t)
userID, err := fixtures.CreateTestUser(db.DB)
if err != nil {
t.Fatalf("failed to create test user: %v", err)
}
defer fixtures.DeleteUser(db.DB, userID)
serviceID, err := fixtures.CreateTestService(db.DB)
if err != nil {
t.Fatalf("failed to create test service: %v", err)
}
defer fixtures.DeleteService(db.DB, serviceID)
_, err = db.DB.Exec(context.Background(), "UPDATE users SET deposits_required = 3 WHERE id = $1", userID)
if err != nil {
t.Fatalf("failed to set deposits_required: %v", err)
}
token := jwt.GenerateUserToken(userID)
london, err := time.LoadLocation("Europe/London")
if err != nil {
t.Fatalf("Europe/London not available: %v", err)
}
bookingTime := nextWeekday(time.Monday, london).Add(10 * time.Hour)
req := CreateBookingRequest{
StartTime: bookingTime,
ServiceIDs: []string{serviceID},
}
handler := http.HandlerFunc(CreateBookingHandler)
w := makeRequest(handler, "POST", "/api/bookings", req, token)
if w.Code != http.StatusCreated {
t.Fatalf("expected status 201, got %d. body: %s", w.Code, w.Body.String())
}
var booking Booking
if err := parseResponseBody(w, &booking); err != nil {
t.Fatalf("failed to parse booking response: %v", err)
}
if !booking.DepositRequired {
t.Error("expected deposit_required=true on first booking")
}
var depositRequired bool
err = db.DB.QueryRow(context.Background(),
"SELECT deposit_required FROM bookings WHERE id = $1", booking.ID).Scan(&depositRequired)
if err != nil {
t.Fatalf("failed to query booking: %v", err)
}
if !depositRequired {
t.Error("expected deposit_required=true in DB for first booking")
}
_, err = db.DB.Exec(context.Background(), "UPDATE users SET deposits_required = 0 WHERE id = $1", userID)
if err != nil {
t.Fatalf("failed to update deposits_required: %v", err)
}
err = db.DB.QueryRow(context.Background(),
"SELECT deposit_required FROM bookings WHERE id = $1", booking.ID).Scan(&depositRequired)
if err != nil {
t.Fatalf("failed to query booking after user update: %v", err)
}
if !depositRequired {
t.Error("expected booking deposit_required to remain true after user change")
}
bookingTime2 := nextWeekday(time.Tuesday, london).Add(10 * time.Hour)
req2 := CreateBookingRequest{
StartTime: bookingTime2,
ServiceIDs: []string{serviceID},
}
w2 := makeRequest(handler, "POST", "/api/bookings", req2, token)
if w2.Code != http.StatusCreated {
t.Fatalf("expected status 201 for second booking, got %d. body: %s", w2.Code, w2.Body.String())
}
var booking2 Booking
if err := parseResponseBody(w2, &booking2); err != nil {
t.Fatalf("failed to parse second booking response: %v", err)
}
if booking2.DepositRequired {
t.Error("expected deposit_required=false on second booking after user deposits_required=0")
}
var depositRequired2 bool
err = db.DB.QueryRow(context.Background(),
"SELECT deposit_required FROM bookings WHERE id = $1", booking2.ID).Scan(&depositRequired2)
if err != nil {
t.Fatalf("failed to query second booking: %v", err)
}
if depositRequired2 {
t.Error("expected second booking deposit_required=false in DB")
}
}
+6 -4
View File
@@ -317,9 +317,10 @@ func AdminEditBookingHandler(w http.ResponseWriter, r *http.Request) {
}
// Check if salon is closed (exceptional hours) - admin gets warning but can proceed
weekday := int(req.StartTime.Weekday())
// DB uses 0=Monday..6=Sunday; Go uses 0=Sunday..6=Saturday. Convert.
weekday := int((req.StartTime.Weekday() + 6) % 7)
bookingTime := req.StartTime.Format("15:04:05")
daysToMonday := weekday
daysToMonday := int(req.StartTime.Weekday())
if daysToMonday == 0 {
daysToMonday = 7
}
@@ -612,8 +613,9 @@ func AdminCreateBookingForUserHandler(w http.ResponseWriter, r *http.Request) {
// Check if booking time falls within a closed exceptional hours period
// Calculate the Monday of the week containing the booking date
weekday := int(req.StartTime.Weekday())
daysToMonday := weekday
// DB uses 0=Monday..6=Sunday; Go uses 0=Sunday..6=Saturday. Convert.
weekday := int((req.StartTime.Weekday() + 6) % 7)
daysToMonday := int(req.StartTime.Weekday())
if daysToMonday == 0 {
daysToMonday = 7 // Sunday -> next Monday
}
+6 -3
View File
@@ -106,7 +106,9 @@ func ReserveSlotHandler(w http.ResponseWriter, r *http.Request) {
}
// f. Validate working hours: SELECT end_time::text FROM working_hours WHERE weekday = $1
weekday := int(req.StartTime.Weekday())
// DB uses 0=Monday..6=Sunday; Go uses 0=Sunday..6=Saturday. Convert.
localStart := req.StartTime.In(londonLocation)
weekday := int((localStart.Weekday() + 6) % 7)
var closeStr string
if err := db.DB.QueryRow(r.Context(), `SELECT end_time::text FROM working_hours WHERE weekday = $1`, weekday).Scan(&closeStr); err != nil {
log.Printf("Failed to get hours: %v", err)
@@ -114,14 +116,15 @@ func ReserveSlotHandler(w http.ResponseWriter, r *http.Request) {
return
}
endTime := req.StartTime.Add(time.Duration(svcDuration) * time.Minute)
localEnd := localStart.Add(time.Duration(svcDuration) * time.Minute)
closeTime, _ := time.Parse("15:04:05", closeStr)
if endTime.Hour() > closeTime.Hour() || (endTime.Hour() == closeTime.Hour() && endTime.Minute() > closeTime.Minute()) {
if localEnd.Hour() > closeTime.Hour() || (localEnd.Hour() == closeTime.Hour() && localEnd.Minute() > closeTime.Minute()) {
http.Error(w, "Cannot book this time - services would extend beyond closing hours", http.StatusBadRequest)
return
}
// g. Check existing booking overlap (same query as CreateBookingHandler)
endTime := req.StartTime.Add(time.Duration(svcDuration) * time.Minute)
var cnt int
db.DB.QueryRow(r.Context(), `
SELECT COUNT(*) FROM bookings WHERE status IN ('confirmed','in_progress','completed')
+273
View File
@@ -365,3 +365,276 @@ func TestReserveSlot_DualCleanup(t *testing.T) {
t.Error("expected 2-hour user reservation to be deleted")
}
}
// seedCustomWorkingHours replaces working_hours with the given schedule.
// DB convention: 0=Monday, 1=Tuesday, ..., 6=Sunday.
func seedCustomWorkingHours(t *testing.T, hours []struct {
weekday int
startTime string
endTime string
isOpen bool
}) {
t.Helper()
for _, h := range hours {
_, err := db.DB.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 for weekday %d: %v", h.weekday, err)
}
}
}
// nextWeekday returns the next occurrence of the given weekday (0=Sunday..6=Saturday)
// in the given location, at least 2 days from now to avoid "in the past" rejections.
func nextWeekday(weekday time.Weekday, loc *time.Location) time.Time {
now := time.Now().In(loc)
daysAhead := int(weekday) - int(now.Weekday())
if daysAhead <= 0 {
daysAhead += 7
}
if daysAhead < 2 {
daysAhead += 7
}
return now.AddDate(0, 0, daysAhead).Truncate(24 * time.Hour)
}
// TestReserveSlot_WeekdayConversion verifies that Go's time.Weekday (0=Sunday)
// is correctly mapped to the DB's weekday convention (0=Monday).
func TestReserveSlot_WeekdayConversion(t *testing.T) {
hours := []struct {
weekday int
startTime string
endTime string
isOpen bool
}{
{0, "09:00", "17:00", true},
{1, "09:00", "17:00", true},
{2, "09:00", "17:00", true},
{3, "09:00", "17:00", true},
{4, "09:00", "17:00", true},
{5, "10:00", "14:00", true},
{6, "00:00", "00:00", false},
}
testdb.TruncateTables(t, db.DB)
seedCustomWorkingHours(t, hours)
serviceID, err := fixtures.CreateTestService(db.DB)
if err != nil {
t.Fatalf("failed to create test service: %v", err)
}
london, err := time.LoadLocation("Europe/London")
if err != nil {
t.Fatalf("Europe/London not available: %v", err)
}
monday := nextWeekday(time.Monday, london)
tests := []struct {
name string
startTime time.Time
expectCode int
}{
{"Monday 10:00", monday.Add(10 * time.Hour), http.StatusCreated},
{"Tuesday 10:00", monday.AddDate(0, 0, 1).Add(10 * time.Hour), http.StatusCreated},
{"Wednesday 10:00", monday.AddDate(0, 0, 2).Add(10 * time.Hour), http.StatusCreated},
{"Thursday 10:00", monday.AddDate(0, 0, 3).Add(10 * time.Hour), http.StatusCreated},
{"Friday 10:00", monday.AddDate(0, 0, 4).Add(10 * time.Hour), http.StatusCreated},
{"Saturday 11:00", monday.AddDate(0, 0, 5).Add(11 * time.Hour), http.StatusCreated},
{"Sunday 11:00", monday.AddDate(0, 0, 6).Add(11 * time.Hour), http.StatusBadRequest},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
w := makeReserveRequest("POST", "/api/bookings/reserve", ReserveSlotRequest{
StartTime: tt.startTime,
ServiceIDs: []string{serviceID},
}, "")
if w.Code != tt.expectCode {
t.Errorf("%s: expected status %d, got %d. body: %s", tt.name, tt.expectCode, w.Code, w.Body.String())
}
})
}
}
// TestReserveSlot_ClosingHoursValidation verifies that bookings extending
// past closing time are rejected, and that the correct day's closing time
// is used (not the wrong day due to weekday mismatch).
func TestReserveSlot_ClosingHoursValidation(t *testing.T) {
hours := []struct {
weekday int
startTime string
endTime string
isOpen bool
}{
{0, "09:00", "17:00", true},
{1, "09:00", "17:00", true},
{2, "09:00", "17:00", true},
{3, "12:00", "20:00", true},
{4, "09:00", "17:00", true},
{5, "10:00", "14:00", true},
{6, "00:00", "00:00", false},
}
testdb.TruncateTables(t, db.DB)
seedCustomWorkingHours(t, hours)
serviceID, err := fixtures.CreateTestService(db.DB)
if err != nil {
t.Fatalf("failed to create test service: %v", err)
}
london, err := time.LoadLocation("Europe/London")
if err != nil {
t.Fatalf("Europe/London not available: %v", err)
}
monday := nextWeekday(time.Monday, london)
thursday := monday.AddDate(0, 0, 3)
tests := []struct {
name string
startTime time.Time
expectCode int
}{
{"Thursday 17:30 (within 20:00 close)", thursday.Add(17*time.Hour + 30*time.Minute), http.StatusCreated},
{"Thursday 19:30 (past 20:00 close)", thursday.Add(19*time.Hour + 30*time.Minute), http.StatusBadRequest},
{"Monday 16:30 (past 17:00 close)", monday.Add(16*time.Hour + 30*time.Minute), http.StatusBadRequest},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
w := makeReserveRequest("POST", "/api/bookings/reserve", ReserveSlotRequest{
StartTime: tt.startTime,
ServiceIDs: []string{serviceID},
}, "")
if w.Code != tt.expectCode {
t.Errorf("%s: expected status %d, got %d. body: %s", tt.name, tt.expectCode, w.Code, w.Body.String())
}
})
}
}
// TestReserveSlot_UTCtoLondonConversion verifies that a UTC timestamp sent
// from the browser is correctly interpreted as London local time for the
// purpose of working hours lookup.
func TestReserveSlot_UTCtoLondonConversion(t *testing.T) {
hours := []struct {
weekday int
startTime string
endTime string
isOpen bool
}{
{0, "09:00", "17:00", true},
{1, "09:00", "17:00", true},
{2, "09:00", "17:00", true},
{3, "12:00", "20:00", true},
{4, "09:00", "17:00", true},
{5, "10:00", "14:00", true},
{6, "00:00", "00:00", false},
}
testdb.TruncateTables(t, db.DB)
seedCustomWorkingHours(t, hours)
serviceID, err := fixtures.CreateTestService(db.DB)
if err != nil {
t.Fatalf("failed to create test service: %v", err)
}
london, err := time.LoadLocation("Europe/London")
if err != nil {
t.Fatalf("Europe/London not available: %v", err)
}
thursday := nextWeekday(time.Thursday, london)
// Convert to UTC for the request (frontend sends UTC)
thursday1730BST := thursday.Add(17*time.Hour + 30*time.Minute).In(london).UTC()
thursday1930BST := thursday.Add(19*time.Hour + 30*time.Minute).In(london).UTC()
tests := []struct {
name string
startTime time.Time
expectCode int
}{
{"17:30 BST Thursday (within 20:00 close)", thursday1730BST, http.StatusCreated},
{"19:30 BST Thursday (past 20:00 close)", thursday1930BST, http.StatusBadRequest},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
w := makeReserveRequest("POST", "/api/bookings/reserve", ReserveSlotRequest{
StartTime: tt.startTime,
ServiceIDs: []string{serviceID},
}, "")
if w.Code != tt.expectCode {
t.Errorf("%s: expected status %d, got %d. body: %s", tt.name, tt.expectCode, w.Code, w.Body.String())
}
})
}
}
// TestReserveSlot_DifferentClosingPerDay verifies that each day's closing
// time is used independently — a late-booking on a late-closing day should
// succeed while the same time on an early-closing day should fail.
func TestReserveSlot_DifferentClosingPerDay(t *testing.T) {
hours := []struct {
weekday int
startTime string
endTime string
isOpen bool
}{
{0, "09:00", "17:00", true},
{1, "09:00", "17:00", true},
{2, "12:00", "17:00", true},
{3, "12:00", "20:00", true},
{4, "09:00", "17:00", true},
{5, "10:00", "14:00", true},
{6, "00:00", "00:00", false},
}
testdb.TruncateTables(t, db.DB)
seedCustomWorkingHours(t, hours)
serviceID, err := fixtures.CreateTestService(db.DB)
if err != nil {
t.Fatalf("failed to create test service: %v", err)
}
london, err := time.LoadLocation("Europe/London")
if err != nil {
t.Fatalf("Europe/London not available: %v", err)
}
wednesday := nextWeekday(time.Wednesday, london)
thursday := wednesday.AddDate(0, 0, 1)
tests := []struct {
name string
startTime time.Time
expectCode int
}{
{"Wednesday 18:00 (closes 17:00)", wednesday.Add(18 * time.Hour), http.StatusBadRequest},
{"Thursday 18:00 (closes 20:00)", thursday.Add(18 * time.Hour), http.StatusCreated},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
w := makeReserveRequest("POST", "/api/bookings/reserve", ReserveSlotRequest{
StartTime: tt.startTime,
ServiceIDs: []string{serviceID},
}, "")
if w.Code != tt.expectCode {
t.Errorf("%s: expected status %d, got %d. body: %s", tt.name, tt.expectCode, w.Code, w.Body.String())
}
})
}
}
+770
View File
@@ -0,0 +1,770 @@
package payments
import (
"crussell/internal/square"
"crussell/mw"
"encoding/json"
"errors"
"log"
"net/http"
"strconv"
"time"
"github.com/go-chi/chi/v5"
"github.com/jackc/pgx/v5"
)
type CreateTerminalPaymentRequest struct {
Amount int64 `json:"amount"`
PaymentType string `json:"payment_type"`
OverrideAmount *int64 `json:"override_amount,omitempty"`
TipEnabled bool `json:"tip_enabled"`
}
type CreateBookingPaymentRequest struct {
Amount int64 `json:"amount"`
PaymentType string `json:"payment_type"`
CardID *string `json:"card_id,omitempty"`
NewCardToken *string `json:"new_card_token,omitempty"`
SaveCard bool `json:"save_card"`
IdempotencyKey string `json:"idempotency_key"`
}
type RefundRequest struct {
Amount int64 `json:"amount"`
Reason string `json:"reason"`
}
type CreateTipPaymentRequest struct {
Amount int64 `json:"amount"`
CardToken string `json:"card_token"`
}
type CheckoutResponse struct {
CheckoutID string `json:"checkout_id"`
Status string `json:"status"`
}
type PaymentStatusResponse struct {
Status string `json:"status"`
PaymentID string `json:"payment_id,omitempty"`
Amount int64 `json:"amount,omitempty"`
CardBrand string `json:"card_brand,omitempty"`
CardLast4 string `json:"card_last4,omitempty"`
ReceiptURL string `json:"receipt_url,omitempty"`
}
type PaymentResponse struct {
ID string `json:"id"`
BookingID string `json:"booking_id"`
PaymentType string `json:"payment_type"`
Status string `json:"status"`
Amount int64 `json:"amount"`
CardBrand string `json:"card_brand,omitempty"`
CardLast4 string `json:"card_last4,omitempty"`
ReceiptURL string `json:"receipt_url,omitempty"`
CreatedAt string `json:"created_at"`
}
type RefundResponse struct {
ID string `json:"id"`
PaymentID string `json:"payment_id"`
Amount int64 `json:"amount"`
Status string `json:"status"`
Reason string `json:"reason"`
CreatedAt string `json:"created_at"`
}
type PaymentSummaryResponse struct {
TotalAmount int64 `json:"total_amount"`
PaidAmount int64 `json:"paid_amount"`
RefundedAmount int64 `json:"refunded_amount"`
RemainingAmount int64 `json:"remaining_amount"`
Payments []PaymentResponse `json:"payments"`
Refunds []RefundResponse `json:"refunds"`
}
func CreateTerminalPayment(w http.ResponseWriter, r *http.Request) {
bookingID := chi.URLParam(r, "id")
if bookingID == "" {
http.Error(w, "Booking ID is required", http.StatusBadRequest)
return
}
adminID, ok := r.Context().Value(mw.UserIDKey).(string)
if !ok || adminID == "" {
http.Error(w, "Authentication required", http.StatusUnauthorized)
return
}
var req CreateTerminalPaymentRequest
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
log.Printf("Failed to decode terminal payment request: %v", err)
http.Error(w, "Invalid request", http.StatusBadRequest)
return
}
if err := ValidateAmount(req.Amount); err != nil {
http.Error(w, err.Error(), http.StatusBadRequest)
return
}
if err := ValidatePaymentType(req.PaymentType); err != nil {
http.Error(w, err.Error(), http.StatusBadRequest)
return
}
service := NewPaymentService()
status, err := service.GetBookingStatus(r.Context(), bookingID)
if err != nil {
if errors.Is(err, pgx.ErrNoRows) {
http.Error(w, "Booking not found", http.StatusNotFound)
return
}
log.Printf("Failed to get booking status: %v", err)
http.Error(w, "Internal server error", http.StatusInternalServerError)
return
}
if status != "in_progress" && status != "completed" {
http.Error(w, "Booking must be in_progress or completed to create payment", http.StatusBadRequest)
return
}
amount := req.Amount
if req.OverrideAmount != nil {
amount = *req.OverrideAmount
}
idempotencyKey := bookingID + "-" + req.PaymentType + "-" + strconv.FormatInt(amount, 10)
existingPayment, err := service.CheckIdempotency(r.Context(), bookingID, idempotencyKey)
if err != nil {
log.Printf("Failed to check idempotency: %v", err)
}
if existingPayment != nil {
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(CheckoutResponse{
CheckoutID: existingPayment.ID,
Status: existingPayment.Status,
})
return
}
checkoutReq := square.CreateCheckoutReq{
Amount: amount,
Currency: "GBP",
IdempotencyKey: idempotencyKey,
ReferenceID: bookingID,
TipEnabled: req.TipEnabled,
}
checkout, err := SquareClient.CreateCheckout(r.Context(), checkoutReq)
if err != nil {
log.Printf("Failed to create checkout: %v", err)
http.Error(w, "Failed to create payment", http.StatusInternalServerError)
return
}
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(CheckoutResponse{
CheckoutID: checkout.ID,
Status: checkout.Status,
})
_ = adminID
}
func GetCheckoutStatus(w http.ResponseWriter, r *http.Request) {
checkoutID := chi.URLParam(r, "checkout_id")
if checkoutID == "" {
http.Error(w, "Checkout ID is required", http.StatusBadRequest)
return
}
bookingID := r.URL.Query().Get("booking_id")
if bookingID == "" {
http.Error(w, "booking_id query parameter is required", http.StatusBadRequest)
return
}
paymentResult, err := SquareClient.GetCheckout(r.Context(), checkoutID)
if err != nil {
if err.Error() == "checkout pending" {
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(PaymentStatusResponse{Status: "PENDING"})
return
}
log.Printf("Failed to get checkout status: %v", err)
http.Error(w, "Failed to get checkout status", http.StatusInternalServerError)
return
}
if paymentResult.Status == "COMPLETED" {
service := NewPaymentService()
existing, err := service.CheckIdempotency(r.Context(), bookingID, "")
if err != nil {
log.Printf("Failed to check for existing payment: %v", err)
}
if existing != nil && existing.SquarePaymentID != nil && *existing.SquarePaymentID == paymentResult.SquarePayID {
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(PaymentStatusResponse{
Status: "COMPLETED",
PaymentID: existing.ID,
Amount: existing.Amount,
CardBrand: paymentResult.CardBrand,
CardLast4: paymentResult.CardLast4,
ReceiptURL: paymentResult.ReceiptURL,
})
return
}
idempotencyKey := bookingID + "-terminal-" + strconv.FormatInt(paymentResult.Amount, 10)
record := PaymentRecord{
BookingID: bookingID,
PaymentType: "full",
PaymentMethod: "in_person_card",
Status: "completed",
Amount: paymentResult.Amount,
SquarePaymentID: &paymentResult.SquarePayID,
IdempotencyKey: &idempotencyKey,
Fees: paymentResult.Fees,
CreatedAt: time.Now(),
UpdatedAt: time.Now(),
}
paymentID, err := service.CreatePaymentRecord(r.Context(), record)
if err != nil {
log.Printf("Failed to create payment record: %v", err)
http.Error(w, "Internal server error", http.StatusInternalServerError)
return
}
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(PaymentStatusResponse{
Status: "COMPLETED",
PaymentID: paymentID,
Amount: paymentResult.Amount,
CardBrand: paymentResult.CardBrand,
CardLast4: paymentResult.CardLast4,
ReceiptURL: paymentResult.ReceiptURL,
})
return
}
http.Error(w, "Payment failed", http.StatusPaymentRequired)
}
func CreateBookingPayment(w http.ResponseWriter, r *http.Request) {
bookingID := chi.URLParam(r, "id")
if bookingID == "" {
http.Error(w, "Booking ID is required", http.StatusBadRequest)
return
}
userID, ok := r.Context().Value(mw.UserIDKey).(string)
if !ok || userID == "" {
http.Error(w, "Authentication required", http.StatusUnauthorized)
return
}
var req CreateBookingPaymentRequest
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
log.Printf("Failed to decode booking payment request: %v", err)
http.Error(w, "Invalid request", http.StatusBadRequest)
return
}
if err := ValidateAmount(req.Amount); err != nil {
http.Error(w, err.Error(), http.StatusBadRequest)
return
}
if err := ValidatePaymentType(req.PaymentType); err != nil {
http.Error(w, err.Error(), http.StatusBadRequest)
return
}
if err := ValidateCardInfo(req.CardID, req.NewCardToken); err != nil {
http.Error(w, err.Error(), http.StatusBadRequest)
return
}
service := NewPaymentService()
bookingUserID, err := service.GetBookingUserID(r.Context(), bookingID)
if err != nil {
if errors.Is(err, pgx.ErrNoRows) {
http.Error(w, "Booking not found", http.StatusNotFound)
return
}
log.Printf("Failed to get booking user: %v", err)
http.Error(w, "Internal server error", http.StatusInternalServerError)
return
}
if bookingUserID != userID {
http.Error(w, "Unauthorized", http.StatusForbidden)
return
}
existingPayment, err := service.CheckIdempotency(r.Context(), bookingID, req.IdempotencyKey)
if err != nil {
log.Printf("Failed to check idempotency: %v", err)
}
if existingPayment != nil {
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(PaymentResponse{
ID: existingPayment.ID,
BookingID: existingPayment.BookingID,
PaymentType: existingPayment.PaymentType,
Status: existingPayment.Status,
Amount: existingPayment.Amount,
CreatedAt: existingPayment.CreatedAt.Format(time.RFC3339),
})
return
}
var sourceID string
var savedCardID *string
if req.NewCardToken != nil && *req.NewCardToken != "" {
cardOnFile, err := SquareClient.CreateCardOnFile(r.Context(), userID, *req.NewCardToken)
if err != nil {
log.Printf("Failed to create card on file: %v", err)
http.Error(w, "Failed to process card", http.StatusInternalServerError)
return
}
sourceID = cardOnFile.CardID
if req.SaveCard {
cardID, err := service.SaveCardForUser(r.Context(), userID, cardOnFile.CardID, cardOnFile.Brand, cardOnFile.Last4, cardOnFile.ExpMonth, cardOnFile.ExpYear, cardOnFile.Fingerprint)
if err != nil {
log.Printf("Failed to save card: %v", err)
} else {
savedCardID = &cardID
}
}
} else if req.CardID != nil {
card, err := service.GetCardByID(r.Context(), *req.CardID, userID)
if err != nil {
if errors.Is(err, pgx.ErrNoRows) {
http.Error(w, "Card not found", http.StatusNotFound)
return
}
log.Printf("Failed to get card: %v", err)
http.Error(w, "Internal server error", http.StatusInternalServerError)
return
}
sourceID = card.SquareCardID
savedCardID = req.CardID
}
paymentReq := square.CreatePaymentReq{
Amount: req.Amount,
Currency: "GBP",
SourceID: sourceID,
IdempotencyKey: req.IdempotencyKey,
ReferenceID: bookingID,
Note: req.PaymentType,
}
paymentResult, err := SquareClient.CreatePayment(r.Context(), paymentReq)
if err != nil {
log.Printf("Failed to create payment: %v", err)
http.Error(w, "Payment failed", http.StatusPaymentRequired)
return
}
fees := service.CalculateFees(req.Amount, "online")
record := PaymentRecord{
BookingID: bookingID,
PaymentType: req.PaymentType,
PaymentMethod: "online_square",
Status: "completed",
Amount: req.Amount,
SquarePaymentID: &paymentResult.SquarePayID,
IdempotencyKey: &req.IdempotencyKey,
Fees: fees,
UserSavedCardID: savedCardID,
CreatedAt: time.Now(),
UpdatedAt: time.Now(),
CreatedBy: &userID,
}
paymentID, err := service.CreatePaymentRecord(r.Context(), record)
if err != nil {
log.Printf("Failed to create payment record: %v", err)
http.Error(w, "Internal server error", http.StatusInternalServerError)
return
}
if req.PaymentType == "deposit" {
err = service.UpdateBookingDepositPaid(r.Context(), bookingID, true)
if err != nil {
log.Printf("Failed to update deposit paid: %v", err)
}
}
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(PaymentResponse{
ID: paymentID,
BookingID: bookingID,
PaymentType: req.PaymentType,
Status: "completed",
Amount: req.Amount,
CardBrand: paymentResult.CardBrand,
CardLast4: paymentResult.CardLast4,
ReceiptURL: paymentResult.ReceiptURL,
CreatedAt: time.Now().Format(time.RFC3339),
})
}
func GetUserPaymentMethods(w http.ResponseWriter, r *http.Request) {
userID, ok := r.Context().Value(mw.UserIDKey).(string)
if !ok || userID == "" {
http.Error(w, "Authentication required", http.StatusUnauthorized)
return
}
service := NewPaymentService()
cards, err := service.GetUserPaymentMethods(r.Context(), userID)
if err != nil {
log.Printf("Failed to get payment methods: %v", err)
http.Error(w, "Internal server error", http.StatusInternalServerError)
return
}
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(cards)
}
func DeletePaymentMethod(w http.ResponseWriter, r *http.Request) {
cardID := chi.URLParam(r, "id")
if cardID == "" {
http.Error(w, "Card ID is required", http.StatusBadRequest)
return
}
userID, ok := r.Context().Value(mw.UserIDKey).(string)
if !ok || userID == "" {
http.Error(w, "Authentication required", http.StatusUnauthorized)
return
}
service := NewPaymentService()
err := service.DeletePaymentMethod(r.Context(), cardID, userID)
if err != nil {
log.Printf("Failed to delete payment method: %v", err)
http.Error(w, "Internal server error", http.StatusInternalServerError)
return
}
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(map[string]string{"status": "deleted"})
}
func RefundPayment(w http.ResponseWriter, r *http.Request) {
paymentID := chi.URLParam(r, "payment_id")
if paymentID == "" {
http.Error(w, "Payment ID is required", http.StatusBadRequest)
return
}
adminID, ok := r.Context().Value(mw.UserIDKey).(string)
if !ok || adminID == "" {
http.Error(w, "Authentication required", http.StatusUnauthorized)
return
}
var req RefundRequest
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
log.Printf("Failed to decode refund request: %v", err)
http.Error(w, "Invalid request", http.StatusBadRequest)
return
}
if err := ValidateAmount(req.Amount); err != nil {
http.Error(w, err.Error(), http.StatusBadRequest)
return
}
if err := ValidateRefundReason(req.Reason); err != nil {
http.Error(w, err.Error(), http.StatusBadRequest)
return
}
service := NewPaymentService()
payment, err := service.GetPaymentByID(r.Context(), paymentID)
if err != nil {
if errors.Is(err, pgx.ErrNoRows) {
http.Error(w, "Payment not found", http.StatusNotFound)
return
}
log.Printf("Failed to get payment: %v", err)
http.Error(w, "Internal server error", http.StatusInternalServerError)
return
}
if payment.Status != "completed" {
http.Error(w, "Can only refund completed payments", http.StatusBadRequest)
return
}
if payment.SquarePaymentID == nil {
http.Error(w, "Payment has no Square reference", http.StatusBadRequest)
return
}
alreadyRefunded, err := service.GetAlreadyRefundedAmount(r.Context(), paymentID)
if err != nil {
log.Printf("Failed to get already refunded amount: %v", err)
http.Error(w, "Internal server error", http.StatusInternalServerError)
return
}
if req.Amount+alreadyRefunded > payment.Amount {
http.Error(w, "Refund amount exceeds payment amount", http.StatusBadRequest)
return
}
refundReq := square.RefundPaymentReq{
PaymentID: *payment.SquarePaymentID,
Amount: req.Amount,
IdempotencyKey: paymentID + "-" + strconv.FormatInt(req.Amount, 10),
Reason: req.Reason,
}
refundResult, err := SquareClient.RefundPayment(r.Context(), refundReq)
if err != nil {
log.Printf("Failed to refund payment: %v", err)
http.Error(w, "Refund failed", http.StatusInternalServerError)
return
}
squareRefundID := refundResult.ID
record := RefundRecord{
PaymentID: paymentID,
BookingID: payment.BookingID,
Amount: req.Amount,
SquareRefundID: &squareRefundID,
Status: "completed",
Reason: req.Reason,
CreatedBy: &adminID,
CreatedAt: time.Now(),
}
refundID, err := service.CreateRefundRecord(r.Context(), record)
if err != nil {
log.Printf("Failed to create refund record: %v", err)
http.Error(w, "Internal server error", http.StatusInternalServerError)
return
}
if payment.PaymentType == "deposit" && (req.Amount+alreadyRefunded) >= payment.Amount {
err = service.UpdateBookingDepositPaid(r.Context(), payment.BookingID, false)
if err != nil {
log.Printf("Failed to update deposit paid: %v", err)
}
}
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(RefundResponse{
ID: refundID,
PaymentID: paymentID,
Amount: req.Amount,
Status: "completed",
Reason: req.Reason,
CreatedAt: time.Now().Format(time.RFC3339),
})
}
func CreateTipPayment(w http.ResponseWriter, r *http.Request) {
bookingID := chi.URLParam(r, "id")
if bookingID == "" {
http.Error(w, "Booking ID is required", http.StatusBadRequest)
return
}
userID, ok := r.Context().Value(mw.UserIDKey).(string)
if !ok || userID == "" {
http.Error(w, "Authentication required", http.StatusUnauthorized)
return
}
var req CreateTipPaymentRequest
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
log.Printf("Failed to decode tip payment request: %v", err)
http.Error(w, "Invalid request", http.StatusBadRequest)
return
}
if err := ValidateAmount(req.Amount); err != nil {
http.Error(w, err.Error(), http.StatusBadRequest)
return
}
if req.CardToken == "" {
http.Error(w, "Card token is required", http.StatusBadRequest)
return
}
service := NewPaymentService()
bookingUserID, err := service.GetBookingUserID(r.Context(), bookingID)
if err != nil {
if errors.Is(err, pgx.ErrNoRows) {
http.Error(w, "Booking not found", http.StatusNotFound)
return
}
log.Printf("Failed to get booking user: %v", err)
http.Error(w, "Internal server error", http.StatusInternalServerError)
return
}
if bookingUserID != userID {
http.Error(w, "Unauthorized", http.StatusForbidden)
return
}
hasCompleted, err := service.HasCompletedPayment(r.Context(), bookingID)
if err != nil {
log.Printf("Failed to check for completed payments: %v", err)
http.Error(w, "Internal server error", http.StatusInternalServerError)
return
}
if !hasCompleted {
http.Error(w, "Booking must have a completed payment before adding tip", http.StatusBadRequest)
return
}
idempotencyKey := bookingID + "-tip-" + strconv.FormatInt(req.Amount, 10)
paymentReq := square.CreatePaymentReq{
Amount: req.Amount,
Currency: "GBP",
SourceID: req.CardToken,
IdempotencyKey: idempotencyKey,
ReferenceID: bookingID,
Note: "tip",
}
paymentResult, err := SquareClient.CreatePayment(r.Context(), paymentReq)
if err != nil {
log.Printf("Failed to create tip payment: %v", err)
http.Error(w, "Payment failed", http.StatusPaymentRequired)
return
}
record := PaymentRecord{
BookingID: bookingID,
PaymentType: "tip",
PaymentMethod: "online_square",
Status: "completed",
Amount: req.Amount,
SquarePaymentID: &paymentResult.SquarePayID,
IdempotencyKey: &idempotencyKey,
Fees: 0,
CreatedAt: time.Now(),
UpdatedAt: time.Now(),
CreatedBy: &userID,
}
paymentID, err := service.CreatePaymentRecord(r.Context(), record)
if err != nil {
log.Printf("Failed to create payment record: %v", err)
http.Error(w, "Internal server error", http.StatusInternalServerError)
return
}
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(PaymentResponse{
ID: paymentID,
BookingID: bookingID,
PaymentType: "tip",
Status: "completed",
Amount: req.Amount,
CardBrand: paymentResult.CardBrand,
CardLast4: paymentResult.CardLast4,
ReceiptURL: paymentResult.ReceiptURL,
CreatedAt: time.Now().Format(time.RFC3339),
})
}
func GetBookingPaymentSummary(w http.ResponseWriter, r *http.Request) {
bookingID := chi.URLParam(r, "id")
if bookingID == "" {
http.Error(w, "Booking ID is required", http.StatusBadRequest)
return
}
userID, _ := r.Context().Value(mw.UserIDKey).(string)
userRole, _ := r.Context().Value(mw.UserRoleKey).(string)
service := NewPaymentService()
if userRole != "admin" && userID != "" {
bookingUserID, err := service.GetBookingUserID(r.Context(), bookingID)
if err != nil {
if errors.Is(err, pgx.ErrNoRows) {
http.Error(w, "Booking not found", http.StatusNotFound)
return
}
log.Printf("Failed to get booking user: %v", err)
http.Error(w, "Internal server error", http.StatusInternalServerError)
return
}
if bookingUserID != userID {
http.Error(w, "Unauthorized", http.StatusForbidden)
return
}
}
summary, err := service.GetBookingPaymentSummary(r.Context(), bookingID)
if err != nil {
log.Printf("Failed to get payment summary: %v", err)
http.Error(w, "Internal server error", http.StatusInternalServerError)
return
}
payments := make([]PaymentResponse, len(summary.Payments))
for i, p := range summary.Payments {
payments[i] = PaymentResponse{
ID: p.ID,
BookingID: p.BookingID,
PaymentType: p.PaymentType,
Status: p.Status,
Amount: p.Amount,
CreatedAt: p.CreatedAt.Format(time.RFC3339),
}
}
refunds := make([]RefundResponse, len(summary.Refunds))
for i, rf := range summary.Refunds {
refunds[i] = RefundResponse{
ID: rf.ID,
PaymentID: rf.PaymentID,
Amount: rf.Amount,
Status: rf.Status,
Reason: rf.Reason,
CreatedAt: rf.CreatedAt.Format(time.RFC3339),
}
}
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(PaymentSummaryResponse{
TotalAmount: summary.TotalAmount,
PaidAmount: summary.PaidAmount,
RefundedAmount: summary.RefundedAmount,
RemainingAmount: summary.RemainingAmount,
Payments: payments,
Refunds: refunds,
})
}
+853
View File
@@ -0,0 +1,853 @@
//go:build test && dev
// +build test,dev
package payments
import (
"bytes"
"context"
"encoding/base64"
"encoding/json"
"net/http"
"net/http/httptest"
"os"
"testing"
"time"
"crussell/db"
"crussell/internal/square"
"crussell/mw"
"crussell/testutils/fixtures"
"crussell/testutils/jwt"
"crussell/testutils/testdb"
"github.com/go-chi/chi/v5"
"github.com/jackc/pgx/v5/pgxpool"
)
func TestMain(m *testing.M) {
pool, err := testdb.NewPool("")
if err != nil {
panic(err)
}
testdb.Migrate(&testing.T{}, pool)
db.DB = pool
jwt.Init()
square.Client = square.NewDevClient()
SquareClient = square.Client
code := m.Run()
pool.Close()
os.Exit(code)
}
func resetTestData(t *testing.T) {
t.Helper()
testdb.TruncateTables(t, db.DB)
}
func makePaymentRequest(handler http.HandlerFunc, method, path string, body interface{}, token string) *httptest.ResponseRecorder {
return makePaymentAuthRequest(handler, method, path, body, token, "")
}
func makePaymentAuthRequest(handler http.HandlerFunc, method, path string, body interface{}, token, userIDOverride 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)
}
if token != "" {
req.Header.Set("Authorization", "Bearer "+token)
}
rctx := chi.NewRouteContext()
if id, paramName := extractPaymentIDFromPath(path); id != "" {
rctx.URLParams.Add(paramName, id)
}
ctx := context.WithValue(req.Context(), chi.RouteCtxKey, rctx)
var userID, userRole string
if userIDOverride != "" {
userID = userIDOverride
userRole = "verified_email"
} else if token != "" {
if info := extractUserFromTestJWT(token); info != nil {
userID = info.userID
userRole = info.role
}
}
if userID != "" {
ctx = context.WithValue(ctx, mw.UserIDKey, userID)
ctx = context.WithValue(ctx, mw.UserRoleKey, userRole)
}
req = req.WithContext(ctx)
w := httptest.NewRecorder()
handler(w, req)
return w
}
type paymentUserInfo struct {
userID string
role string
}
func extractUserFromTestJWT(token string) *paymentUserInfo {
parts := splitToken(token)
if len(parts) != 3 {
return nil
}
decoded, err := base64URLDecode(parts[1])
if err != nil {
return nil
}
var claims map[string]interface{}
if err := json.Unmarshal(decoded, &claims); err != nil {
return nil
}
userID, _ := claims["user_id"].(string)
role, _ := claims["role"].(string)
if userID == "" {
return nil
}
return &paymentUserInfo{userID: userID, role: role}
}
func splitToken(token string) []string {
var result []string
var current []byte
for _, c := range token {
if c == '.' {
result = append(result, string(current))
current = nil
} else {
current = append(current, byte(c))
}
}
if len(current) > 0 {
result = append(result, string(current))
}
return result
}
func base64URLDecode(s string) ([]byte, error) {
return base64.URLEncoding.DecodeString(s)
}
func extractPaymentIDFromPath(path string) (string, string) {
patterns := []struct {
prefix string
paramName string
}{
{"/api/admin/payments/", "payment_id"},
{"/api/admin/bookings/", "id"},
{"/api/admin/bookings/", "id"},
{"/api/user/payment-methods/", "id"},
}
for _, p := range patterns {
if idx := findPaymentLastSegment(path, p.prefix); idx >= 0 {
endIdx := len(path)
for i := idx; i < len(path); i++ {
if path[i] == '/' {
endIdx = i
break
}
}
return path[idx:endIdx], p.paramName
}
}
return "", ""
}
func findPaymentLastSegment(path, prefix string) int {
for i := len(path) - 1; i >= len(prefix); i-- {
if len(path) > i && path[i-len(prefix):i] == prefix {
return i
}
}
return -1
}
func parsePaymentResponseBody(w *httptest.ResponseRecorder, dest interface{}) error {
return json.Unmarshal(w.Body.Bytes(), dest)
}
func TestTerminalPayment_HappyPath(t *testing.T) {
resetTestData(t)
_, bookingID, _ := setupTestData(t)
adminToken := jwt.GenerateAdminToken()
req := CreateTerminalPaymentRequest{
Amount: 5000,
PaymentType: "full",
TipEnabled: true,
}
handler := CreateTerminalPayment
w := makePaymentRequest(handler, "POST", "/api/admin/bookings/"+bookingID+"/payment", req, adminToken)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
var resp CheckoutResponse
if err := parsePaymentResponseBody(w, &resp); err != nil {
t.Errorf("failed to parse response: %v", err)
}
if resp.CheckoutID == "" {
t.Error("expected checkout ID to be set")
}
if resp.Status != "PENDING" {
t.Errorf("expected status PENDING, got %s", resp.Status)
}
var count int
err := db.DB.QueryRow(context.Background(), "SELECT COUNT(*) FROM payments WHERE booking_id = $1", bookingID).Scan(&count)
if err != nil {
t.Errorf("failed to query payments: %v", err)
}
if count != 0 {
t.Errorf("expected 0 payments (created on completion), got %d", count)
}
_ = userID
}
func setupTestData(t *testing.T) (string, string, string) {
userID, err := fixtures.CreateTestUser(db.DB)
if err != nil {
t.Fatalf("failed to create test user: %v", err)
}
serviceID, err := fixtures.CreateTestService(db.DB)
if err != nil {
t.Fatalf("failed to create test service: %v", err)
}
bookingID, err := fixtures.CreateTestBooking(db.DB, userID, serviceID)
if err != nil {
t.Fatalf("failed to create test booking: %v", err)
}
_, err = db.DB.Exec(context.Background(), "UPDATE bookings SET status = 'in_progress' WHERE id = $1", bookingID)
if err != nil {
t.Fatalf("failed to update booking status: %v", err)
}
return userID, bookingID, serviceID
}
func TestTerminalPayment_PriceOverride(t *testing.T) {
resetTestData(t)
_, bookingID, _ := setupTestData(t)
adminToken := jwt.GenerateAdminToken()
overrideAmount := int64(3000)
req := CreateTerminalPaymentRequest{
Amount: 5000,
PaymentType: "full",
OverrideAmount: &overrideAmount,
TipEnabled: false,
}
handler := CreateTerminalPayment
w := makePaymentRequest(handler, "POST", "/api/admin/bookings/"+bookingID+"/payment", req, adminToken)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
var resp CheckoutResponse
if err := parsePaymentResponseBody(w, &resp); err != nil {
t.Errorf("failed to parse response: %v", err)
}
}
func TestTerminalPayment_BookingNotInProgress(t *testing.T) {
resetTestData(t)
userID, err := fixtures.CreateTestUser(db.DB)
if err != nil {
t.Fatalf("failed to create test user: %v", err)
}
serviceID, err := fixtures.CreateTestService(db.DB)
if err != nil {
t.Fatalf("failed to create test service: %v", err)
}
bookingID, err := fixtures.CreateTestBooking(db.DB, userID, serviceID)
if err != nil {
t.Fatalf("failed to create test booking: %v", err)
}
adminToken := jwt.GenerateAdminToken()
req := CreateTerminalPaymentRequest{
Amount: 5000,
PaymentType: "full",
}
handler := CreateTerminalPayment
w := makePaymentRequest(handler, "POST", "/api/admin/bookings/"+bookingID+"/payment", req, adminToken)
if w.Code != http.StatusBadRequest {
t.Errorf("expected status 400, got %d. body: %s", w.Code, w.Body.String())
}
_ = serviceID
}
func TestTerminalPayment_BookingNotFound(t *testing.T) {
resetTestData(t)
adminToken := jwt.GenerateAdminToken()
req := CreateTerminalPaymentRequest{
Amount: 5000,
PaymentType: "full",
}
handler := CreateTerminalPayment
w := makePaymentRequest(handler, "POST", "/api/admin/bookings/non-existent/payment", req, adminToken)
if w.Code != http.StatusNotFound {
t.Errorf("expected status 404, got %d. body: %s", w.Code, w.Body.String())
}
}
func TestOnlinePayment_NewCard_Deposit(t *testing.T) {
resetTestData(t)
userID, bookingID, _ := setupTestData(t)
userToken := jwt.GenerateUserToken(userID)
cardToken := "cnon:test-card-nonce"
req := CreateBookingPaymentRequest{
Amount: 2500,
PaymentType: "deposit",
NewCardToken: &cardToken,
SaveCard: true,
IdempotencyKey: "deposit-key-1",
}
handler := CreateBookingPayment
w := makePaymentRequest(handler, "POST", "/api/bookings/"+bookingID+"/payment", req, userToken)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
var resp PaymentResponse
if err := parsePaymentResponseBody(w, &resp); err != nil {
t.Errorf("failed to parse response: %v", err)
}
if resp.ID == "" {
t.Error("expected payment ID to be set")
}
if resp.Status != "completed" {
t.Errorf("expected status completed, got %s", resp.Status)
}
if resp.Amount != 2500 {
t.Errorf("expected amount 2500, got %d", resp.Amount)
}
var count int
err := db.DB.QueryRow(context.Background(), "SELECT COUNT(*) FROM payments WHERE booking_id = $1", bookingID).Scan(&count)
if err != nil {
t.Errorf("failed to query payments: %v", err)
}
if count != 1 {
t.Errorf("expected 1 payment, got %d", count)
}
}
func TestOnlinePayment_SavedCard(t *testing.T) {
resetTestData(t)
userID, bookingID, _ := setupTestData(t)
cardID, err := fixtures.CreateTestPaymentMethod(db.DB, userID, "cfa_mock_card_123", "VISA", "4242")
if err != nil {
t.Fatalf("failed to create payment method: %v", err)
}
userToken := jwt.GenerateUserToken(userID)
req := CreateBookingPaymentRequest{
Amount: 5000,
PaymentType: "full",
CardID: &cardID,
IdempotencyKey: "saved-card-key-1",
}
handler := CreateBookingPayment
w := makePaymentRequest(handler, "POST", "/api/bookings/"+bookingID+"/payment", req, userToken)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
var resp PaymentResponse
if err := parsePaymentResponseBody(w, &resp); err != nil {
t.Errorf("failed to parse response: %v", err)
}
if resp.Status != "completed" {
t.Errorf("expected status completed, got %s", resp.Status)
}
}
func TestOnlinePayment_BookingNotOwned(t *testing.T) {
resetTestData(t)
_, bookingID, _ := setupTestData(t)
otherUserID, err := fixtures.CreateTestUser(db.DB)
if err != nil {
t.Fatalf("failed to create other user: %v", err)
}
userToken := jwt.GenerateUserToken(otherUserID)
cardToken := "cnon:test-card-nonce"
req := CreateBookingPaymentRequest{
Amount: 5000,
PaymentType: "full",
NewCardToken: &cardToken,
IdempotencyKey: "not-owned-key-1",
}
handler := CreateBookingPayment
w := makePaymentRequest(handler, "POST", "/api/bookings/"+bookingID+"/payment", req, userToken)
if w.Code != http.StatusForbidden {
t.Errorf("expected status 403, got %d. body: %s", w.Code, w.Body.String())
}
}
func TestGetUserPaymentMethods_HasCards(t *testing.T) {
resetTestData(t)
userID, err := fixtures.CreateTestUser(db.DB)
if err != nil {
t.Fatalf("failed to create test user: %v", err)
}
_, err = fixtures.CreateTestPaymentMethod(db.DB, userID, "cfa_card_1", "VISA", "1111")
if err != nil {
t.Fatalf("failed to create payment method 1: %v", err)
}
_, err = fixtures.CreateTestPaymentMethod(db.DB, userID, "cfa_card_2", "MASTERCARD", "2222")
if err != nil {
t.Fatalf("failed to create payment method 2: %v", err)
}
userToken := jwt.GenerateUserToken(userID)
handler := GetUserPaymentMethods
w := makePaymentRequest(handler, "GET", "/api/user/payment-methods", nil, userToken)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
var cards []SavedCard
if err := parsePaymentResponseBody(w, &cards); err != nil {
t.Errorf("failed to parse response: %v", err)
}
if len(cards) != 2 {
t.Errorf("expected 2 cards, got %d", len(cards))
}
}
func TestDeletePaymentMethod(t *testing.T) {
resetTestData(t)
userID, err := fixtures.CreateTestUser(db.DB)
if err != nil {
t.Fatalf("failed to create test user: %v", err)
}
cardID, err := fixtures.CreateTestPaymentMethod(db.DB, userID, "cfa_card_delete", "VISA", "9999")
if err != nil {
t.Fatalf("failed to create payment method: %v", err)
}
userToken := jwt.GenerateUserToken(userID)
handler := DeletePaymentMethod
w := makePaymentRequest(handler, "DELETE", "/api/user/payment-methods/"+cardID, nil, userToken)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
var resp map[string]string
if err := parsePaymentResponseBody(w, &resp); err != nil {
t.Errorf("failed to parse response: %v", err)
}
if resp["status"] != "deleted" {
t.Errorf("expected status deleted, got %s", resp["status"])
}
}
func TestRefund_FullRefund(t *testing.T) {
resetTestData(t)
_, bookingID, _ := setupTestData(t)
adminToken := jwt.GenerateAdminToken()
paymentID, err := fixtures.CreateTestPayment(db.DB, bookingID, 5000.00, "in_person_card", "full", "completed")
if err != nil {
t.Fatalf("failed to create payment: %v", err)
}
squarePaymentID := "sqp_test_123"
_, err = db.DB.Exec(context.Background(), "UPDATE payments SET square_payment_id = $1 WHERE id = $2", squarePaymentID, paymentID)
if err != nil {
t.Fatalf("failed to update payment: %v", err)
}
req := RefundRequest{
Amount: 5000,
Reason: "customer request",
}
handler := RefundPayment
w := makePaymentRequest(handler, "POST", "/api/admin/payments/"+paymentID+"/refund", req, adminToken)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
var resp RefundResponse
if err := parsePaymentResponseBody(w, &resp); err != nil {
t.Errorf("failed to parse response: %v", err)
}
if resp.Amount != 5000 {
t.Errorf("expected amount 5000, got %d", resp.Amount)
}
if resp.Status != "completed" {
t.Errorf("expected status completed, got %s", resp.Status)
}
}
func TestRefund_PartialRefund(t *testing.T) {
resetTestData(t)
_, bookingID, _ := setupTestData(t)
adminToken := jwt.GenerateAdminToken()
paymentID, err := fixtures.CreateTestPayment(db.DB, bookingID, 5000.00, "in_person_card", "full", "completed")
if err != nil {
t.Fatalf("failed to create payment: %v", err)
}
squarePaymentID := "sqp_test_456"
_, err = db.DB.Exec(context.Background(), "UPDATE payments SET square_payment_id = $1 WHERE id = $2", squarePaymentID, paymentID)
if err != nil {
t.Fatalf("failed to update payment: %v", err)
}
req := RefundRequest{
Amount: 2500,
Reason: "partial refund",
}
handler := RefundPayment
w := makePaymentRequest(handler, "POST", "/api/admin/payments/"+paymentID+"/refund", req, adminToken)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
var resp RefundResponse
if err := parsePaymentResponseBody(w, &resp); err != nil {
t.Errorf("failed to parse response: %v", err)
}
if resp.Amount != 2500 {
t.Errorf("expected amount 2500, got %d", resp.Amount)
}
}
func TestRefund_OverRefundRejected(t *testing.T) {
resetTestData(t)
_, bookingID, _ := setupTestData(t)
adminToken := jwt.GenerateAdminToken()
paymentID, err := fixtures.CreateTestPayment(db.DB, bookingID, 5000.00, "in_person_card", "full", "completed")
if err != nil {
t.Fatalf("failed to create payment: %v", err)
}
squarePaymentID := "sqp_test_789"
_, err = db.DB.Exec(context.Background(), "UPDATE payments SET square_payment_id = $1 WHERE id = $2", squarePaymentID, paymentID)
if err != nil {
t.Fatalf("failed to update payment: %v", err)
}
req := RefundRequest{
Amount: 6000,
Reason: "over refund attempt",
}
handler := RefundPayment
w := makePaymentRequest(handler, "POST", "/api/admin/payments/"+paymentID+"/refund", req, adminToken)
if w.Code != http.StatusBadRequest {
t.Errorf("expected status 400, got %d. body: %s", w.Code, w.Body.String())
}
}
func TestRefund_PaymentNotFound(t *testing.T) {
resetTestData(t)
adminToken := jwt.GenerateAdminToken()
req := RefundRequest{
Amount: 1000,
Reason: "test",
}
handler := RefundPayment
w := makePaymentRequest(handler, "POST", "/api/admin/payments/non-existent/refund", req, adminToken)
if w.Code != http.StatusNotFound {
t.Errorf("expected status 404, got %d. body: %s", w.Code, w.Body.String())
}
}
func TestRefund_PendingPaymentRejected(t *testing.T) {
resetTestData(t)
_, bookingID, _ := setupTestData(t)
adminToken := jwt.GenerateAdminToken()
paymentID, err := fixtures.CreateTestPayment(db.DB, bookingID, 5000.00, "in_person_card", "full", "pending")
if err != nil {
t.Fatalf("failed to create payment: %v", err)
}
req := RefundRequest{
Amount: 5000,
Reason: "test",
}
handler := RefundPayment
w := makePaymentRequest(handler, "POST", "/api/admin/payments/"+paymentID+"/refund", req, adminToken)
if w.Code != http.StatusBadRequest {
t.Errorf("expected status 400, got %d. body: %s", w.Code, w.Body.String())
}
}
func TestTipPayment_HappyPath(t *testing.T) {
resetTestData(t)
userID, bookingID, _ := setupTestData(t)
_, err := fixtures.CreateTestPayment(db.DB, bookingID, 5000.00, "online_square", "full", "completed")
if err != nil {
t.Fatalf("failed to create payment: %v", err)
}
userToken := jwt.GenerateUserToken(userID)
cardToken := "cnon:tip-card"
req := CreateTipPaymentRequest{
Amount: 500,
CardToken: cardToken,
}
handler := CreateTipPayment
w := makePaymentRequest(handler, "POST", "/api/bookings/"+bookingID+"/tip", req, userToken)
if w.Code != http.StatusOK {
t.Errorf("expected status 200, got %d. body: %s", w.Code, w.Body.String())
}
var resp PaymentResponse
if err := parsePaymentResponseBody(w, &resp); err != nil {
t.Errorf("failed to parse response: %v", err)
}
if resp.PaymentType != "tip" {
t.Errorf("expected payment type tip, got %s", resp.PaymentType)
}
if resp.Amount != 500 {
t.Errorf("expected amount 500, got %d", resp.Amount)
}
}
func TestTipPayment_NoPriorPayment(t *testing.T) {
resetTestData(t)
userID, bookingID, _ := setupTestData(t)
userToken := jwt.GenerateUserToken(userID)
cardToken := "cnon:tip-card"
req := CreateTipPaymentRequest{
Amount: 500,
CardToken: cardToken,
}
handler := CreateTipPayment
w := makePaymentRequest(handler, "POST", "/api/bookings/"+bookingID+"/tip", req, userToken)
if w.Code != http.StatusBadRequest {
t.Errorf("expected status 400, got %d. body: %s", w.Code, w.Body.String())
}
}
func TestIdempotency_SameKeyReturnsExisting(t *testing.T) {
resetTestData(t)
userID, bookingID, _ := setupTestData(t)
userToken := jwt.GenerateUserToken(userID)
cardToken := "cnon:idempotent-card"
idempotencyKey := "idempotent-same-key"
req1 := CreateBookingPaymentRequest{
Amount: 5000,
PaymentType: "full",
NewCardToken: &cardToken,
IdempotencyKey: idempotencyKey,
}
handler := CreateBookingPayment
w1 := makePaymentRequest(handler, "POST", "/api/bookings/"+bookingID+"/payment", req1, userToken)
if w1.Code != http.StatusOK {
t.Errorf("first request expected status 200, got %d. body: %s", w1.Code, w1.Body.String())
}
var resp1 PaymentResponse
if err := parsePaymentResponseBody(w1, &resp1); err != nil {
t.Errorf("failed to parse first response: %v", err)
}
req2 := CreateBookingPaymentRequest{
Amount: 5000,
PaymentType: "full",
NewCardToken: &cardToken,
IdempotencyKey: idempotencyKey,
}
w2 := makePaymentRequest(handler, "POST", "/api/bookings/"+bookingID+"/payment", req2, userToken)
if w2.Code != http.StatusOK {
t.Errorf("second request expected status 200, got %d. body: %s", w2.Code, w.Body.String())
}
var resp2 PaymentResponse
if err := parsePaymentResponseBody(w2, &resp2); err != nil {
t.Errorf("failed to parse second response: %v", err)
}
if resp1.ID != resp2.ID {
t.Errorf("expected same payment ID, got %s and %s", resp1.ID, resp2.ID)
}
var count int
err := db.DB.QueryRow(context.Background(), "SELECT COUNT(*) FROM payments WHERE booking_id = $1", bookingID).Scan(&count)
if err != nil {
t.Errorf("failed to query payments: %v", err)
}
if count != 1 {
t.Errorf("expected 1 payment (idempotent), got %d", count)
}
}
func TestIdempotency_DifferentKeyCreatesNew(t *testing.T) {
resetTestData(t)
userID, bookingID, _ := setupTestData(t)
userToken := jwt.GenerateUserToken(userID)
cardToken := "cnon:different-key-card"
req1 := CreateBookingPaymentRequest{
Amount: 5000,
PaymentType: "full",
NewCardToken: &cardToken,
IdempotencyKey: "key-1",
}
handler := CreateBookingPayment
w1 := makePaymentRequest(handler, "POST", "/api/bookings/"+bookingID+"/payment", req1, userToken)
if w1.Code != http.StatusOK {
t.Errorf("first request expected status 200, got %d. body: %s", w1.Code, w1.Body.String())
}
req2 := CreateBookingPaymentRequest{
Amount: 5000,
PaymentType: "full",
NewCardToken: &cardToken,
IdempotencyKey: "key-2",
}
w2 := makePaymentRequest(handler, "POST", "/api/bookings/"+bookingID+"/payment", req2, userToken)
if w2.Code != http.StatusOK {
t.Errorf("second request expected status 200, got %d. body: %s", w2.Code, w2.Body.String())
}
var count int
err := db.DB.QueryRow(context.Background(), "SELECT COUNT(*) FROM payments WHERE booking_id = $1", bookingID).Scan(&count)
if err != nil {
t.Errorf("failed to query payments: %v", err)
}
if count != 2 {
t.Errorf("expected 2 payments (different keys), got %d", count)
}
}
func TestSquareWebhook_DevMode_NoSignature(t *testing.T) {
resetTestData(t)
req := httptest.NewRequest("POST", "/api/webhooks/square", nil)
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
_ = req
_ = w
t.Skip("webhook handler tested in webhooks package")
}
+391
View File
@@ -0,0 +1,391 @@
package payments
import (
"context"
"crussell/db"
"crussell/internal/square"
"errors"
"time"
"github.com/jackc/pgx/v5"
)
type SavedCard struct {
ID string `json:"id"`
SquareCardID string `json:"square_card_id"`
Brand string `json:"brand"`
Last4 string `json:"last_4"`
ExpMonth int `json:"exp_month"`
ExpYear int `json:"exp_year"`
Fingerprint string `json:"fingerprint"`
IsDefault bool `json:"is_default"`
}
type PaymentService struct{}
func NewPaymentService() *PaymentService {
return &PaymentService{}
}
type PaymentRecord struct {
ID string
BookingID string
PaymentType string
PaymentMethod string
VendorCode *string
InvoiceNumber *int
Status string
Amount int64
IsVATApplicable bool
VATRate *float64
VATAmount *int64
NetAmount *int64
UserSavedCardID *string
SquarePaymentID *string
IdempotencyKey *string
Fees int64
CreatedAt time.Time
UpdatedAt time.Time
CreatedBy *string
}
type RefundRecord struct {
ID string
PaymentID string
BookingID string
Amount int64
SquareRefundID *string
Status string
Reason string
CreatedBy *string
CreatedAt time.Time
}
type PaymentSummary struct {
TotalAmount int64
PaidAmount int64
RefundedAmount int64
RemainingAmount int64
Payments []PaymentRecord
Refunds []RefundRecord
}
func (s *PaymentService) CalculateFees(amount int64, method string) int64 {
if method == "online" {
return (amount * 14 / 1000) + 25
}
return (amount * 175 / 10000)
}
func (s *PaymentService) CreatePaymentRecord(ctx context.Context, record PaymentRecord) (string, error) {
var id string
err := db.DB.QueryRow(ctx, `
INSERT INTO payments (
booking_id, payment_type, payment_method, vendor_code, invoice_number,
status, amount, is_vat_applicable, vat_rate, vat_amount, net_amount,
user_saved_card_id, square_payment_id, idempotency_key, fees, created_at, updated_at, created_by
) VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13, $14, $15, $16, $17, $18)
RETURNING id
`,
record.BookingID,
record.PaymentType,
record.PaymentMethod,
record.VendorCode,
record.InvoiceNumber,
record.Status,
record.Amount,
record.IsVATApplicable,
record.VATRate,
record.VATAmount,
record.NetAmount,
record.UserSavedCardID,
record.SquarePaymentID,
record.IdempotencyKey,
record.Fees,
record.CreatedAt,
record.UpdatedAt,
record.CreatedBy,
).Scan(&id)
if err != nil {
return "", err
}
return id, nil
}
func (s *PaymentService) CreateRefundRecord(ctx context.Context, record RefundRecord) (string, error) {
var id string
err := db.DB.QueryRow(ctx, `
INSERT INTO refunds (
payment_id, booking_id, amount, square_refund_id, status, reason, created_by, created_at
) VALUES ($1, $2, $3, $4, $5, $6, $7, $8)
RETURNING id
`,
record.PaymentID,
record.BookingID,
record.Amount,
record.SquareRefundID,
record.Status,
record.Reason,
record.CreatedBy,
record.CreatedAt,
).Scan(&id)
if err != nil {
return "", err
}
return id, nil
}
func (s *PaymentService) GetBookingPaymentSummary(ctx context.Context, bookingID string) (*PaymentSummary, error) {
summary := &PaymentSummary{
Payments: []PaymentRecord{},
Refunds: []RefundRecord{},
}
var totalAmount int64
err := db.DB.QueryRow(ctx, `
SELECT COALESCE(SUM(
COALESCE(bs.override_price, s.price)
), 0)
FROM booking_services bs
JOIN services s ON bs.service_id = s.id
WHERE bs.booking_id = $1
`, bookingID).Scan(&totalAmount)
if err != nil && !errors.Is(err, pgx.ErrNoRows) {
return nil, err
}
summary.TotalAmount = totalAmount
rows, err := db.DB.Query(ctx, `
SELECT id, booking_id, payment_type, payment_method, vendor_code, invoice_number,
status, amount, is_vat_applicable, vat_rate, vat_amount, net_amount,
user_saved_card_id, square_payment_id, idempotency_key, fees, created_at, updated_at, created_by
FROM payments
WHERE booking_id = $1
ORDER BY created_at ASC
`, bookingID)
if err != nil {
return nil, err
}
defer rows.Close()
var paidAmount int64
for rows.Next() {
var p PaymentRecord
err := rows.Scan(
&p.ID, &p.BookingID, &p.PaymentType, &p.PaymentMethod, &p.VendorCode, &p.InvoiceNumber,
&p.Status, &p.Amount, &p.IsVATApplicable, &p.VATRate, &p.VATAmount, &p.NetAmount,
&p.UserSavedCardID, &p.SquarePaymentID, &p.IdempotencyKey, &p.Fees,
&p.CreatedAt, &p.UpdatedAt, &p.CreatedBy,
)
if err != nil {
return nil, err
}
summary.Payments = append(summary.Payments, p)
if p.Status == "completed" {
paidAmount += p.Amount
}
}
summary.PaidAmount = paidAmount
refundRows, err := db.DB.Query(ctx, `
SELECT id, payment_id, booking_id, amount, square_refund_id, status, reason, created_by, created_at
FROM refunds
WHERE booking_id = $1 AND status = 'completed'
ORDER BY created_at ASC
`, bookingID)
if err != nil {
return nil, err
}
defer refundRows.Close()
var refundedAmount int64
for refundRows.Next() {
var r RefundRecord
err := refundRows.Scan(
&r.ID, &r.PaymentID, &r.BookingID, &r.Amount, &r.SquareRefundID,
&r.Status, &r.Reason, &r.CreatedBy, &r.CreatedAt,
)
if err != nil {
return nil, err
}
summary.Refunds = append(summary.Refunds, r)
refundedAmount += r.Amount
}
summary.RefundedAmount = refundedAmount
summary.RemainingAmount = totalAmount - paidAmount + refundedAmount
return summary, nil
}
func (s *PaymentService) CheckIdempotency(ctx context.Context, bookingID, idempotencyKey string) (*PaymentRecord, error) {
var p PaymentRecord
err := db.DB.QueryRow(ctx, `
SELECT id, booking_id, payment_type, payment_method, vendor_code, invoice_number,
status, amount, is_vat_applicable, vat_rate, vat_amount, net_amount,
user_saved_card_id, square_payment_id, idempotency_key, fees, created_at, updated_at, created_by
FROM payments
WHERE booking_id = $1 AND idempotency_key = $2
`, bookingID, idempotencyKey).Scan(
&p.ID, &p.BookingID, &p.PaymentType, &p.PaymentMethod, &p.VendorCode, &p.InvoiceNumber,
&p.Status, &p.Amount, &p.IsVATApplicable, &p.VATRate, &p.VATAmount, &p.NetAmount,
&p.UserSavedCardID, &p.SquarePaymentID, &p.IdempotencyKey, &p.Fees,
&p.CreatedAt, &p.UpdatedAt, &p.CreatedBy,
)
if err != nil {
if errors.Is(err, pgx.ErrNoRows) {
return nil, nil
}
return nil, err
}
return &p, nil
}
func (s *PaymentService) GetPaymentByID(ctx context.Context, paymentID string) (*PaymentRecord, error) {
var p PaymentRecord
err := db.DB.QueryRow(ctx, `
SELECT id, booking_id, payment_type, payment_method, vendor_code, invoice_number,
status, amount, is_vat_applicable, vat_rate, vat_amount, net_amount,
user_saved_card_id, square_payment_id, idempotency_key, fees, created_at, updated_at, created_by
FROM payments
WHERE id = $1
`, paymentID).Scan(
&p.ID, &p.BookingID, &p.PaymentType, &p.PaymentMethod, &p.VendorCode, &p.InvoiceNumber,
&p.Status, &p.Amount, &p.IsVATApplicable, &p.VATRate, &p.VATAmount, &p.NetAmount,
&p.UserSavedCardID, &p.SquarePaymentID, &p.IdempotencyKey, &p.Fees,
&p.CreatedAt, &p.UpdatedAt, &p.CreatedBy,
)
if err != nil {
return nil, err
}
return &p, nil
}
func (s *PaymentService) GetAlreadyRefundedAmount(ctx context.Context, paymentID string) (int64, error) {
var amount int64
err := db.DB.QueryRow(ctx, `
SELECT COALESCE(SUM(amount), 0) FROM refunds
WHERE payment_id = $1 AND status = 'completed'
`, paymentID).Scan(&amount)
if err != nil {
return 0, err
}
return amount, nil
}
func (s *PaymentService) UpdateBookingDepositPaid(ctx context.Context, bookingID string, depositPaid bool) error {
_, err := db.DB.Exec(ctx, `
UPDATE bookings SET deposit_paid = $1 WHERE id = $2
`, depositPaid, bookingID)
return err
}
func (s *PaymentService) HasCompletedPayment(ctx context.Context, bookingID string) (bool, error) {
var count int
err := db.DB.QueryRow(ctx, `
SELECT COUNT(*) FROM payments
WHERE booking_id = $1 AND status = 'completed' AND payment_type IN ('full', 'deposit', 'balance', 'partial')
`, bookingID).Scan(&count)
if err != nil {
return false, err
}
return count > 0, nil
}
func (s *PaymentService) GetBookingStatus(ctx context.Context, bookingID string) (string, error) {
var status string
err := db.DB.QueryRow(ctx, `SELECT status FROM bookings WHERE id = $1`, bookingID).Scan(&status)
if err != nil {
return "", err
}
return status, nil
}
func (s *PaymentService) GetBookingUserID(ctx context.Context, bookingID string) (string, error) {
var userID string
err := db.DB.QueryRow(ctx, `SELECT user_id FROM bookings WHERE id = $1`, bookingID).Scan(&userID)
if err != nil {
return "", err
}
return userID, nil
}
func (s *PaymentService) GetUserPaymentMethods(ctx context.Context, userID string) ([]SavedCard, error) {
rows, err := db.DB.Query(ctx, `
SELECT id, square_card_id, brand, last_4, exp_month, exp_year, fingerprint, is_default
FROM user_saved_cards
WHERE user_id = $1 AND deleted_at IS NULL
ORDER BY is_default DESC, created_at DESC
`, userID)
if err != nil {
return nil, err
}
defer rows.Close()
var cards []SavedCard
for rows.Next() {
var c SavedCard
err := rows.Scan(&c.ID, &c.SquareCardID, &c.Brand, &c.Last4, &c.ExpMonth, &c.ExpYear, &c.Fingerprint, &c.IsDefault)
if err != nil {
return nil, err
}
cards = append(cards, c)
}
if cards == nil {
cards = []SavedCard{}
}
return cards, nil
}
func (s *PaymentService) DeletePaymentMethod(ctx context.Context, cardID, userID string) error {
retainedUntil := time.Now().Add(7 * 365 * 24 * time.Hour)
_, err := db.DB.Exec(ctx, `
UPDATE user_saved_cards
SET deleted_at = NOW(), deleted_by = $1, retained_until = $2
WHERE id = $3 AND user_id = $1
`, userID, retainedUntil, cardID)
return err
}
func (s *PaymentService) SaveCardForUser(ctx context.Context, userID, squareCardID, brand, last4 string, expMonth, expYear int, fingerprint string) (string, error) {
var id string
err := db.DB.QueryRow(ctx, `
INSERT INTO user_saved_cards (
user_id, square_card_id, brand, last_4, exp_month, exp_year, fingerprint, is_default, created_at
) VALUES ($1, $2, $3, $4, $5, $6, $7, false, NOW())
RETURNING id
`, userID, squareCardID, brand, last4, expMonth, expYear, fingerprint).Scan(&id)
if err != nil {
return "", err
}
return id, nil
}
func (s *PaymentService) GetCardByID(ctx context.Context, cardID, userID string) (*SavedCard, error) {
var c SavedCard
err := db.DB.QueryRow(ctx, `
SELECT id, square_card_id, brand, last_4, exp_month, exp_year, fingerprint, is_default
FROM user_saved_cards
WHERE id = $1 AND user_id = $2 AND deleted_at IS NULL
`, cardID, userID).Scan(&c.ID, &c.SquareCardID, &c.Brand, &c.Last4, &c.ExpMonth, &c.ExpYear, &c.Fingerprint, &c.IsDefault)
if err != nil {
return nil, err
}
return &c, nil
}
var SquareClient square.SquareClient
+44
View File
@@ -0,0 +1,44 @@
package payments
import "errors"
// Valid payment types
var validPaymentTypes = map[string]bool{
"deposit": true,
"full": true,
"tip": true,
"balance": true,
"partial": true,
}
// ValidateAmount checks that amount is greater than 0
func ValidateAmount(amount int64) error {
if amount <= 0 {
return errors.New("amount must be greater than 0")
}
return nil
}
// ValidatePaymentType checks that payment type is valid
func ValidatePaymentType(pt string) error {
if !validPaymentTypes[pt] {
return errors.New("invalid payment type")
}
return nil
}
// ValidateRefundReason checks that refund reason is not empty
func ValidateRefundReason(reason string) error {
if reason == "" {
return errors.New("refund reason is required")
}
return nil
}
// ValidateCardInfo checks that at least one of cardID or newCardToken is provided
func ValidateCardInfo(cardID, newCardToken *string) error {
if cardID == nil && (newCardToken == nil || *newCardToken == "") {
return errors.New("either card_id or new_card_token is required")
}
return nil
}
+79
View File
@@ -0,0 +1,79 @@
package webhooks
import (
"crypto/hmac"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"io"
"log"
"net/http"
"os"
)
type SquareWebhookEvent struct {
Type string `json:"type"`
EventID string `json:"event_id"`
CreatedAt string `json:"created_at"`
Data json.RawMessage `json:"data"`
LocationID string `json:"location_id"`
}
func HandleSquareWebhook(w http.ResponseWriter, r *http.Request) {
body, err := io.ReadAll(r.Body)
if err != nil {
log.Printf("Failed to read webhook body: %v", err)
http.Error(w, "Internal server error", http.StatusInternalServerError)
return
}
defer r.Body.Close()
signature := r.Header.Get("x-square-signature")
signingKey := os.Getenv("SQUARE_WEBHOOK_SIGNATURE_KEY")
if signingKey != "" && signature != "" {
if !verifySquareSignature(body, signature, signingKey) {
log.Printf("Invalid Square webhook signature")
http.Error(w, "Invalid signature", http.StatusForbidden)
return
}
}
var event SquareWebhookEvent
if err := json.Unmarshal(body, &event); err != nil {
log.Printf("Failed to parse webhook event: %v", err)
http.Error(w, "Invalid event", http.StatusBadRequest)
return
}
log.Printf("[SQUARE-WEBHOOK] Received event: %s", event.Type)
switch event.Type {
case "payment.updated":
handlePaymentUpdated(event.Data)
case "refund.updated":
handleRefundUpdated(event.Data)
case "dispute.created":
log.Printf("[SQUARE-WEBHOOK] Dispute created: %s", event.EventID)
default:
log.Printf("[SQUARE-WEBHOOK] Unknown event type: %s", event.Type)
}
w.WriteHeader(http.StatusOK)
w.Write([]byte("ok"))
}
func verifySquareSignature(body []byte, signature, signingKey string) bool {
mac := hmac.New(sha256.New, []byte(signingKey))
mac.Write(body)
expected := hex.EncodeToString(mac.Sum(nil))
return hmac.Equal([]byte(signature), []byte(expected))
}
func handlePaymentUpdated(data json.RawMessage) {
log.Printf("[SQUARE-WEBHOOK] payment.updated: %s", string(data))
}
func handleRefundUpdated(data json.RawMessage) {
log.Printf("[SQUARE-WEBHOOK] refund.updated: %s", string(data))
}
+49
View File
@@ -0,0 +1,49 @@
//go:build !dev
// +build !dev
package square
import (
"context"
"errors"
)
var Client SquareClient
type ProdClient struct{}
func NewClient() SquareClient {
return NewProdClient()
}
func NewProdClient() SquareClient {
return &ProdClient{}
}
func (p *ProdClient) CreatePayment(ctx context.Context, req CreatePaymentReq) (*PaymentResult, error) {
return nil, errors.New("square payments not yet configured — set SQUARE_ACCESS_TOKEN and SQUARE_LOCATION_ID in .env")
}
func (p *ProdClient) CreateCheckout(ctx context.Context, req CreateCheckoutReq) (*CheckoutResult, error) {
return nil, errors.New("square payments not yet configured — set SQUARE_ACCESS_TOKEN and SQUARE_LOCATION_ID in .env")
}
func (p *ProdClient) GetCheckout(ctx context.Context, checkoutID string) (*PaymentResult, error) {
return nil, errors.New("square payments not yet configured — set SQUARE_ACCESS_TOKEN and SQUARE_LOCATION_ID in .env")
}
func (p *ProdClient) RefundPayment(ctx context.Context, req RefundPaymentReq) (*RefundResult, error) {
return nil, errors.New("square payments not yet configured — set SQUARE_ACCESS_TOKEN and SQUARE_LOCATION_ID in .env")
}
func (p *ProdClient) CreateCardOnFile(ctx context.Context, userID, cardToken string) (*CardOnFile, error) {
return nil, errors.New("square payments not yet configured — set SQUARE_ACCESS_TOKEN and SQUARE_LOCATION_ID in .env")
}
func (p *ProdClient) GetCardsOnFile(ctx context.Context, userID string) ([]CardOnFile, error) {
return nil, errors.New("square payments not yet configured — set SQUARE_ACCESS_TOKEN and SQUARE_LOCATION_ID in .env")
}
func (p *ProdClient) DeleteCardOnFile(ctx context.Context, cardID string) error {
return errors.New("square payments not yet configured — set SQUARE_ACCESS_TOKEN and SQUARE_LOCATION_ID in .env")
}
+252
View File
@@ -0,0 +1,252 @@
//go:build dev
// +build dev
package square
import (
"context"
"fmt"
"log"
"os"
"sync"
"time"
)
var Client SquareClient
type MockClient struct {
mu sync.RWMutex
cards map[string]map[string]*CardOnFile
checkouts map[string]*CheckoutResult
payments map[string]*PaymentResult
refunds map[string]*RefundResult
completed map[string]*PaymentResult
}
type devProdClient struct{}
func (d *devProdClient) CreatePayment(ctx context.Context, req CreatePaymentReq) (*PaymentResult, error) {
return nil, fmt.Errorf("square payments not yet configured — set SQUARE_ACCESS_TOKEN and SQUARE_LOCATION_ID in .env")
}
func (d *devProdClient) CreateCheckout(ctx context.Context, req CreateCheckoutReq) (*CheckoutResult, error) {
return nil, fmt.Errorf("square payments not yet configured — set SQUARE_ACCESS_TOKEN and SQUARE_LOCATION_ID in .env")
}
func (d *devProdClient) GetCheckout(ctx context.Context, checkoutID string) (*PaymentResult, error) {
return nil, fmt.Errorf("square payments not yet configured — set SQUARE_ACCESS_TOKEN and SQUARE_LOCATION_ID in .env")
}
func (d *devProdClient) RefundPayment(ctx context.Context, req RefundPaymentReq) (*RefundResult, error) {
return nil, fmt.Errorf("square payments not yet configured — set SQUARE_ACCESS_TOKEN and SQUARE_LOCATION_ID in .env")
}
func (d *devProdClient) CreateCardOnFile(ctx context.Context, userID, cardToken string) (*CardOnFile, error) {
return nil, fmt.Errorf("square payments not yet configured — set SQUARE_ACCESS_TOKEN and SQUARE_LOCATION_ID in .env")
}
func (d *devProdClient) GetCardsOnFile(ctx context.Context, userID string) ([]CardOnFile, error) {
return nil, fmt.Errorf("square payments not yet configured — set SQUARE_ACCESS_TOKEN and SQUARE_LOCATION_ID in .env")
}
func (d *devProdClient) DeleteCardOnFile(ctx context.Context, cardID string) error {
return fmt.Errorf("square payments not yet configured — set SQUARE_ACCESS_TOKEN and SQUARE_LOCATION_ID in .env")
}
func NewClient() SquareClient {
return NewDevClient()
}
func NewDevClient() SquareClient {
env := os.Getenv("SQUARE_ENVIRONMENT")
if env == "sandbox" || env == "production" {
log.Printf("[SQUARE-MOCK] SQUARE_ENVIRONMENT=%s — real client TODO stub", env)
return &devProdClient{}
}
log.Println("[SQUARE-MOCK] Using in-memory mock client")
return &MockClient{
cards: make(map[string]map[string]*CardOnFile),
checkouts: make(map[string]*CheckoutResult),
payments: make(map[string]*PaymentResult),
refunds: make(map[string]*RefundResult),
completed: make(map[string]*PaymentResult),
}
}
func (m *MockClient) CreatePayment(ctx context.Context, req CreatePaymentReq) (*PaymentResult, error) {
log.Printf("[SQUARE-MOCK] CreatePayment: amount=%d, reference=%s", req.Amount, req.ReferenceID)
time.Sleep(1 * time.Second)
m.mu.Lock()
defer m.mu.Unlock()
paymentID := fmt.Sprintf("pay_mock_%d", time.Now().UnixNano())
fees := req.Amount*14/1000 + 25 // online rate: 1.4% + 25p
result := &PaymentResult{
ID: paymentID,
Status: "COMPLETED",
Amount: req.Amount,
CardBrand: "VISA",
CardLast4: "4242",
TipAmount: 0,
ReceiptURL: "https://squareup.com/receipt/" + paymentID,
SquarePayID: "sqp_" + paymentID,
Fees: fees,
}
m.payments[paymentID] = result
log.Printf("[SQUARE-MOCK] Payment completed: id=%s, fees=%d", paymentID, fees)
return result, nil
}
func (m *MockClient) CreateCheckout(ctx context.Context, req CreateCheckoutReq) (*CheckoutResult, error) {
log.Printf("[SQUARE-MOCK] CreateCheckout: amount=%d, tipEnabled=%v, reference=%s", req.Amount, req.TipEnabled, req.ReferenceID)
checkoutID := fmt.Sprintf("chk_mock_%d", time.Now().UnixNano())
result := &CheckoutResult{
ID: checkoutID,
Status: "PENDING",
}
m.mu.Lock()
m.checkouts[checkoutID] = result
m.mu.Unlock()
go func() {
time.Sleep(3 * time.Second)
m.mu.Lock()
defer m.mu.Unlock()
paymentID := fmt.Sprintf("pay_%d", time.Now().UnixNano())
amount := req.Amount
tipAmount := int64(0)
if req.TipEnabled {
tipAmount = 500
amount += tipAmount
}
fees := amount*175/10000 // in-person rate: 1.75%
paymentResult := &PaymentResult{
ID: paymentID,
Status: "COMPLETED",
Amount: amount,
CardBrand: "VISA",
CardLast4: "4242",
TipAmount: tipAmount,
ReceiptURL: "https://squareup.com/receipt/" + paymentID,
SquarePayID: "sqp_" + paymentID,
Fees: fees,
}
m.completed[checkoutID] = paymentResult
m.checkouts[checkoutID].Status = "COMPLETED"
log.Printf("[SQUARE-MOCK] Checkout completed: id=%s, amount=%d, tip=%d", checkoutID, amount, tipAmount)
}()
return result, nil
}
func (m *MockClient) GetCheckout(ctx context.Context, checkoutID string) (*PaymentResult, error) {
log.Printf("[SQUARE-MOCK] GetCheckout: id=%s", checkoutID)
m.mu.RLock()
checkout, ok := m.checkouts[checkoutID]
m.mu.RUnlock()
if !ok {
return nil, fmt.Errorf("checkout not found: %s", checkoutID)
}
if checkout.Status == "PENDING" {
return nil, fmt.Errorf("checkout pending")
}
m.mu.RLock()
result, ok := m.completed[checkoutID]
m.mu.RUnlock()
if !ok {
return nil, fmt.Errorf("checkout result not found: %s", checkoutID)
}
return result, nil
}
func (m *MockClient) RefundPayment(ctx context.Context, req RefundPaymentReq) (*RefundResult, error) {
log.Printf("[SQUARE-MOCK] RefundPayment: payment=%s, amount=%d", req.PaymentID, req.Amount)
time.Sleep(1 * time.Second)
m.mu.Lock()
defer m.mu.Unlock()
refundID := fmt.Sprintf("ref_mock_%d", time.Now().UnixNano())
amount := req.Amount
if amount == 0 {
if payment, ok := m.payments[req.PaymentID]; ok {
amount = payment.Amount
}
}
result := &RefundResult{
ID: refundID,
Status: "COMPLETED",
Amount: amount,
}
m.refunds[refundID] = result
log.Printf("[SQUARE-MOCK] Refund completed: id=%s, amount=%d", refundID, amount)
return result, nil
}
func (m *MockClient) CreateCardOnFile(ctx context.Context, userID, cardToken string) (*CardOnFile, error) {
log.Printf("[SQUARE-MOCK] CreateCardOnFile: user=%s", userID)
m.mu.Lock()
defer m.mu.Unlock()
if m.cards[userID] == nil {
m.cards[userID] = make(map[string]*CardOnFile)
}
cardID := fmt.Sprintf("mock_card_%d", time.Now().UnixNano())
card := &CardOnFile{
ID: cardID,
CardID: "cfa_" + cardID,
Brand: "VISA",
Last4: "4242",
ExpMonth: 12,
ExpYear: 2030,
Fingerprint: fmt.Sprintf("fp_%d", time.Now().UnixNano()),
IsDefault: len(m.cards[userID]) == 0,
}
m.cards[userID][cardID] = card
log.Printf("[SQUARE-MOCK] Card created: id=%s, brand=%s, last4=%s", cardID, card.Brand, card.Last4)
return card, nil
}
func (m *MockClient) GetCardsOnFile(ctx context.Context, userID string) ([]CardOnFile, error) {
log.Printf("[SQUARE-MOCK] GetCardsOnFile: user=%s", userID)
m.mu.RLock()
defer m.mu.RUnlock()
userCards, ok := m.cards[userID]
if !ok {
return []CardOnFile{}, nil
}
var cards []CardOnFile
for _, card := range userCards {
cards = append(cards, *card)
}
return cards, nil
}
func (m *MockClient) DeleteCardOnFile(ctx context.Context, cardID string) error {
log.Printf("[SQUARE-MOCK] DeleteCardOnFile: id=%s", cardID)
m.mu.Lock()
defer m.mu.Unlock()
for userID, cards := range m.cards {
if _, ok := cards[cardID]; ok {
delete(m.cards[userID], cardID)
log.Printf("[SQUARE-MOCK] Card deleted: id=%s (user=%s)", cardID, userID)
return nil
}
}
return fmt.Errorf("card not found: %s", cardID)
}
+353
View File
@@ -0,0 +1,353 @@
//go:build test && dev
// +build test,dev
package square
import (
"context"
"sync"
"testing"
"time"
)
func TestDevClient_CreatePayment_ReturnsCompleted(t *testing.T) {
client := NewDevClient().(*MockClient)
ctx := context.Background()
req := CreatePaymentReq{
Amount: 5000,
Currency: "GBP",
SourceID: "cnon:test-card",
IdempotencyKey: "test-key-1",
ReferenceID: "booking-123",
Note: "full",
}
result, err := client.CreatePayment(ctx, req)
if err != nil {
t.Fatalf("CreatePayment failed: %v", err)
}
if result.Status != "COMPLETED" {
t.Errorf("expected status COMPLETED, got %s", result.Status)
}
if result.Amount != 5000 {
t.Errorf("expected amount 5000, got %d", result.Amount)
}
if result.CardBrand != "VISA" {
t.Errorf("expected card brand VISA, got %s", result.CardBrand)
}
if result.CardLast4 != "4242" {
t.Errorf("expected last4 4242, got %s", result.CardLast4)
}
if result.Fees == 0 {
t.Error("expected fees to be calculated")
}
}
func TestDevClient_CreateCheckout_PendingThenCompleted(t *testing.T) {
client := NewDevClient().(*MockClient)
ctx := context.Background()
req := CreateCheckoutReq{
Amount: 7500,
Currency: "GBP",
IdempotencyKey: "checkout-key-1",
ReferenceID: "booking-456",
TipEnabled: true,
}
result, err := client.CreateCheckout(ctx, req)
if err != nil {
t.Fatalf("CreateCheckout failed: %v", err)
}
if result.Status != "PENDING" {
t.Errorf("expected status PENDING, got %s", result.Status)
}
if result.ID == "" {
t.Error("expected checkout ID to be set")
}
time.Sleep(4 * time.Second)
completed, err := client.GetCheckout(ctx, result.ID)
if err != nil {
t.Fatalf("GetCheckout failed: %v", err)
}
if completed.Status != "COMPLETED" {
t.Errorf("expected status COMPLETED after wait, got %s", completed.Status)
}
if completed.Amount != 8000 {
t.Errorf("expected amount 8000 (7500 + 500 tip), got %d", completed.Amount)
}
if completed.TipAmount != 500 {
t.Errorf("expected tip 500, got %d", completed.TipAmount)
}
}
func TestDevClient_CreateCheckout_NoTip(t *testing.T) {
client := NewDevClient().(*MockClient)
ctx := context.Background()
req := CreateCheckoutReq{
Amount: 5000,
Currency: "GBP",
IdempotencyKey: "checkout-key-notip",
ReferenceID: "booking-789",
TipEnabled: false,
}
result, err := client.CreateCheckout(ctx, req)
if err != nil {
t.Fatalf("CreateCheckout failed: %v", err)
}
if result.Status != "PENDING" {
t.Errorf("expected status PENDING, got %s", result.Status)
}
time.Sleep(4 * time.Second)
completed, err := client.GetCheckout(ctx, result.ID)
if err != nil {
t.Fatalf("GetCheckout failed: %v", err)
}
if completed.Amount != 5000 {
t.Errorf("expected amount 5000 (no tip), got %d", completed.Amount)
}
if completed.TipAmount != 0 {
t.Errorf("expected tip 0, got %d", completed.TipAmount)
}
}
func TestDevClient_RefundPayment_ReturnsCompleted(t *testing.T) {
client := NewDevClient().(*MockClient)
ctx := context.Background()
paymentReq := CreatePaymentReq{
Amount: 10000,
Currency: "GBP",
SourceID: "cnon:test-card",
IdempotencyKey: "payment-for-refund",
ReferenceID: "booking-refund",
Note: "full",
}
paymentResult, err := client.CreatePayment(ctx, paymentReq)
if err != nil {
t.Fatalf("CreatePayment failed: %v", err)
}
refundReq := RefundPaymentReq{
PaymentID: paymentResult.ID,
Amount: 5000,
IdempotencyKey: "refund-key-1",
Reason: "customer request",
}
refundResult, err := client.RefundPayment(ctx, refundReq)
if err != nil {
t.Fatalf("RefundPayment failed: %v", err)
}
if refundResult.Status != "COMPLETED" {
t.Errorf("expected status COMPLETED, got %s", refundResult.Status)
}
if refundResult.Amount != 5000 {
t.Errorf("expected amount 5000, got %d", refundResult.Amount)
}
}
func TestDevClient_CardOnFile_CreateAndGet(t *testing.T) {
client := NewDevClient().(*MockClient)
ctx := context.Background()
userID := "user-test-123"
card, err := client.CreateCardOnFile(ctx, userID, "cnon:test-token")
if err != nil {
t.Fatalf("CreateCardOnFile failed: %v", err)
}
if card.ID == "" {
t.Error("expected card ID to be set")
}
if card.Brand != "VISA" {
t.Errorf("expected brand VISA, got %s", card.Brand)
}
if card.Last4 != "4242" {
t.Errorf("expected last4 4242, got %s", card.Last4)
}
if !card.IsDefault {
t.Error("expected first card to be default")
}
cards, err := client.GetCardsOnFile(ctx, userID)
if err != nil {
t.Fatalf("GetCardsOnFile failed: %v", err)
}
if len(cards) != 1 {
t.Errorf("expected 1 card, got %d", len(cards))
}
if cards[0].ID != card.ID {
t.Errorf("expected card ID %s, got %s", card.ID, cards[0].ID)
}
}
func TestDevClient_CardOnFile_MultipleCards(t *testing.T) {
client := NewDevClient().(*MockClient)
ctx := context.Background()
userID := "user-test-multiple"
card1, err := client.CreateCardOnFile(ctx, userID, "cnon:token-1")
if err != nil {
t.Fatalf("CreateCardOnFile failed: %v", err)
}
card2, err := client.CreateCardOnFile(ctx, userID, "cnon:token-2")
if err != nil {
t.Fatalf("CreateCardOnFile failed: %v", err)
}
cards, err := client.GetCardsOnFile(ctx, userID)
if err != nil {
t.Fatalf("GetCardsOnFile failed: %v", err)
}
if len(cards) != 2 {
t.Errorf("expected 2 cards, got %d", len(cards))
}
if !card1.IsDefault {
t.Error("first card should be default")
}
if card2.IsDefault {
t.Error("second card should not be default")
}
}
func TestDevClient_CardOnFile_Delete(t *testing.T) {
client := NewDevClient().(*MockClient)
ctx := context.Background()
userID := "user-test-delete"
card, err := client.CreateCardOnFile(ctx, userID, "cnon:token-delete")
if err != nil {
t.Fatalf("CreateCardOnFile failed: %v", err)
}
err = client.DeleteCardOnFile(ctx, card.ID)
if err != nil {
t.Fatalf("DeleteCardOnFile failed: %v", err)
}
cards, err := client.GetCardsOnFile(ctx, userID)
if err != nil {
t.Fatalf("GetCardsOnFile failed: %v", err)
}
if len(cards) != 0 {
t.Errorf("expected 0 cards after delete, got %d", len(cards))
}
}
func TestDevClient_CardOnFile_DeleteNotFound(t *testing.T) {
client := NewDevClient().(*MockClient)
ctx := context.Background()
err := client.DeleteCardOnFile(ctx, "non-existent-card")
if err == nil {
t.Error("expected error when deleting non-existent card")
}
}
func TestDevClient_GetCheckout_NotFound(t *testing.T) {
client := NewDevClient().(*MockClient)
ctx := context.Background()
_, err := client.GetCheckout(ctx, "non-existent-checkout")
if err == nil {
t.Error("expected error when checkout not found")
}
}
func TestDevClient_ConcurrentPayments(t *testing.T) {
client := NewDevClient().(*MockClient)
ctx := context.Background()
var wg sync.WaitGroup
results := make(chan *PaymentResult, 10)
errors := make(chan error, 10)
for i := 0; i < 10; i++ {
wg.Add(1)
go func(idx int) {
defer wg.Done()
req := CreatePaymentReq{
Amount: int64(1000 + idx*100),
Currency: "GBP",
SourceID: "cnon:test-card",
IdempotencyKey: "concurrent-key-" + string(rune('0'+idx)),
ReferenceID: "booking-concurrent",
Note: "full",
}
result, err := client.CreatePayment(ctx, req)
if err != nil {
errors <- err
return
}
results <- result
}(i)
}
wg.Wait()
close(results)
close(errors)
errorCount := 0
for err := range errors {
t.Logf("Concurrent payment error: %v", err)
errorCount++
}
if errorCount > 0 {
t.Errorf("expected no errors, got %d", errorCount)
}
resultCount := 0
for result := range results {
if result.Status != "COMPLETED" {
t.Errorf("expected status COMPLETED, got %s", result.Status)
}
resultCount++
}
if resultCount != 10 {
t.Errorf("expected 10 results, got %d", resultCount)
}
}
+71
View File
@@ -0,0 +1,71 @@
package square
import "context"
type CreatePaymentReq struct {
Amount int64 // in pence (GBP cents)
Currency string // "GBP"
SourceID string // card token or "cnon:xxx" nonce
IdempotencyKey string
ReferenceID string // booking ID
Note string
}
type CreateCheckoutReq struct {
Amount int64
Currency string
IdempotencyKey string
ReferenceID string
TipEnabled bool
}
type RefundPaymentReq struct {
PaymentID string
Amount int64 // in pence, 0 = full refund
IdempotencyKey string
Reason string
}
type PaymentResult struct {
ID string
Status string // "COMPLETED", "FAILED", "PENDING"
Amount int64
CardBrand string
CardLast4 string
TipAmount int64
ReceiptURL string
SquarePayID string // Square's payment ID
Fees int64 // processing fee in pence
}
type CheckoutResult struct {
ID string
Status string // "PENDING", "COMPLETED", "FAILED"
}
type CardOnFile struct {
ID string
CardID string // Square's card-on-file token
Brand string
Last4 string
ExpMonth int
ExpYear int
Fingerprint string
IsDefault bool
}
type RefundResult struct {
ID string
Status string
Amount int64
}
type SquareClient interface {
CreatePayment(ctx context.Context, req CreatePaymentReq) (*PaymentResult, error)
CreateCheckout(ctx context.Context, req CreateCheckoutReq) (*CheckoutResult, error)
GetCheckout(ctx context.Context, checkoutID string) (*PaymentResult, error)
RefundPayment(ctx context.Context, req RefundPaymentReq) (*RefundResult, error)
CreateCardOnFile(ctx context.Context, userID, cardToken string) (*CardOnFile, error)
GetCardsOnFile(ctx context.Context, userID string) ([]CardOnFile, error)
DeleteCardOnFile(ctx context.Context, cardID string) error
}
+24
View File
@@ -5,6 +5,7 @@ import (
"crussell/auth"
"crussell/internal/dav"
"crussell/internal/s3"
"crussell/internal/square"
"encoding/json"
"fmt"
"log"
@@ -24,11 +25,13 @@ import (
"crussell/handlers/admin"
"crussell/handlers/bookings"
"crussell/handlers/notifications"
"crussell/handlers/payments"
"crussell/handlers/portfolio"
"crussell/handlers/scheduling"
"crussell/handlers/services"
"crussell/handlers/today"
"crussell/handlers/user"
"crussell/handlers/webhooks"
)
func init() {
@@ -61,6 +64,11 @@ func initS3() {
}
}
func initSquare() {
payments.SquareClient = square.NewClient()
fmt.Println("Square client initialized (dev mock)")
}
func healthCheckHandler(w http.ResponseWriter, r *http.Request) {
status := "ok"
services := map[string]string{
@@ -101,6 +109,7 @@ func main() {
initDB()
initDav()
initS3()
initSquare()
r := chi.NewRouter()
@@ -218,6 +227,13 @@ func main() {
r.Delete("/bookings/{id}", bookings.DeleteBookingHandler)
r.Post("/bookings/{id}/edit-request", bookings.RequestEditHandler)
r.Delete("/bookings/{id}/edit-request", bookings.DeleteEditRequestHandler)
// User payment routes
r.Post("/bookings/{id}/payment", payments.CreateBookingPayment)
r.Get("/user/payment-methods", payments.GetUserPaymentMethods)
r.Delete("/user/payment-methods/{id}", payments.DeletePaymentMethod)
r.Post("/bookings/{id}/tip", payments.CreateTipPayment)
r.Get("/bookings/{id}/payment-summary", payments.GetBookingPaymentSummary)
})
// Admin-only (no rate limit - trusted users with authenticated sessions)
@@ -283,9 +299,17 @@ r.Route("/admin/users", func(r chi.Router) {
r.Delete("/{id}", admin.DeleteDiscountCampaign)
r.Get("/{id}/stats", admin.GetCampaignStats)
})
// Admin payment routes
r.Post("/admin/bookings/{id}/payment", payments.CreateTerminalPayment)
r.Get("/admin/payments/{checkout_id}/status", payments.GetCheckoutStatus)
r.Post("/admin/payments/{payment_id}/refund", payments.RefundPayment)
})
})
// Webhooks (no auth - Square sends to base path)
r.Post("/webhooks/square", webhooks.HandleSquareWebhook)
srv := &http.Server{
Addr: ":8080",
Handler: r,
+75
View File
@@ -228,3 +228,78 @@ func DeleteTimeBlocker(pool *pgxpool.Pool, blockerID string) error {
_, err := pool.Exec(ctx, "DELETE FROM time_blockers WHERE id = $1", blockerID)
return err
}
// CreateTestPayment creates a payment record for testing
// Returns payment ID
func CreateTestPayment(db *pgxpool.Pool, bookingID string, amount float64, method string, ptype string, status string) (string, error) {
ctx := context.Background()
var paymentID string
err := db.QueryRow(ctx, `
INSERT INTO payments (booking_id, payment_type, payment_method, status, amount, created_at, updated_at)
VALUES ($1, $2, $3, $4, $5, NOW(), NOW())
RETURNING id
`, bookingID, ptype, method, status, amount).Scan(&paymentID)
if err != nil {
return "", fmt.Errorf("failed to create payment: %w", err)
}
return paymentID, nil
}
// CreateTestRefund creates a refund record for testing
// Returns refund ID
func CreateTestRefund(db *pgxpool.Pool, paymentID string, bookingID string, amount float64) (string, error) {
ctx := context.Background()
var refundID string
err := db.QueryRow(ctx, `
INSERT INTO refunds (payment_id, booking_id, amount, status, reason, created_at)
VALUES ($1, $2, $3, 'completed', 'test refund', NOW())
RETURNING id
`, paymentID, bookingID, amount).Scan(&refundID)
if err != nil {
return "", fmt.Errorf("failed to create refund: %w", err)
}
return refundID, nil
}
// CreateTestPaymentMethod creates a saved card for a user
// Returns card ID
func CreateTestPaymentMethod(db *pgxpool.Pool, userID string, squareCardID string, brand string, last4 string) (string, error) {
ctx := context.Background()
var cardID string
err := db.QueryRow(ctx, `
INSERT INTO user_saved_cards (user_id, square_card_id, brand, last_4, exp_month, exp_year, fingerprint, is_default, created_at)
VALUES ($1, $2, $3, $4, 12, 2030, 'test_fp', false, NOW())
RETURNING id
`, userID, squareCardID, brand, last4).Scan(&cardID)
if err != nil {
return "", fmt.Errorf("failed to create payment method: %w", err)
}
return cardID, nil
}
// DeletePayment deletes a payment from the database
func DeletePayment(pool *pgxpool.Pool, paymentID string) error {
ctx := context.Background()
_, err := pool.Exec(ctx, "DELETE FROM payments WHERE id = $1", paymentID)
return err
}
// DeleteRefund deletes a refund from the database
func DeleteRefund(pool *pgxpool.Pool, refundID string) error {
ctx := context.Background()
_, err := pool.Exec(ctx, "DELETE FROM refunds WHERE id = $1", refundID)
return err
}
// DeletePaymentMethod deletes a saved card from the database
func DeletePaymentMethod(pool *pgxpool.Pool, cardID string) error {
ctx := context.Background()
_, err := pool.Exec(ctx, "DELETE FROM user_saved_cards WHERE id = $1", cardID)
return err
}
+8
View File
@@ -99,7 +99,11 @@ func Migrate(t *testing.T, pool *pgxpool.Pool) {
"user_notification_preferences",
"user_referrals",
"booking_services",
"refunds",
"payments",
"user_saved_cards",
"square_deposits",
"affiliate_payouts",
"bookings",
"user_patch_tests",
"patch_tests",
@@ -202,7 +206,11 @@ func TruncateTables(t *testing.T, pool *pgxpool.Pool) {
"user_social_logins",
"verification_codes",
"booking_services",
"refunds",
"payments",
"user_saved_cards",
"square_deposits",
"affiliate_payouts",
"bookings",
"booking_edit_requests",
"user_patch_tests",