test: update integration tests for Echo v5 compatibility

Update all integration test files to work with Echo v5 changes.

Changes in new_fixes_test.go:
- Update test helper signatures for *echo.Context
- Fix context handling in test assertions

Changes in security_test.go:
- Update security test signatures for Echo v5

Changes in test_helpers.go:
- Update test setup for Echo v5
- Fix context type usage in test helpers

Changes in websocket_test.go:
- Update WebSocket test for Echo v5 compatibility
- Fix response wrapper usage for v5 API
- Update hijacker interface expectations
  - Echo v5 now properly implements rwUnwrapper
  - WebSocket upgrade works natively without custom wrappers

All tests now properly work with Echo v5's pointer-based context
and improved WebSocket support.
This commit is contained in:
2026-03-06 14:00:56 -05:00
parent a38e4e79da
commit 2cdc2fc913
4 changed files with 119 additions and 102 deletions
+3 -3
View File
@@ -6,7 +6,7 @@ import (
"testing" "testing"
"time" "time"
"github.com/labstack/echo/v4" "github.com/labstack/echo/v5"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
) )
@@ -39,7 +39,7 @@ func TestRateLimiter(t *testing.T) {
e := echo.New() e := echo.New()
// Create a simple handler // Create a simple handler
handler := func(c echo.Context) error { handler := func(c *echo.Context) error {
return c.String(http.StatusOK, "ok") return c.String(http.StatusOK, "ok")
} }
@@ -108,7 +108,7 @@ func (m *mockRateLimiter) Allow(ip string) bool {
func rateLimiterMiddleware(rl *mockRateLimiter) echo.MiddlewareFunc { func rateLimiterMiddleware(rl *mockRateLimiter) echo.MiddlewareFunc {
return func(next echo.HandlerFunc) echo.HandlerFunc { return func(next echo.HandlerFunc) echo.HandlerFunc {
return func(c echo.Context) error { return func(c *echo.Context) error {
ip := c.RealIP() ip := c.RealIP()
if ip == "" { if ip == "" {
ip = c.Request().RemoteAddr ip = c.Request().RemoteAddr
+2 -2
View File
@@ -7,7 +7,7 @@ import (
"testing" "testing"
"time" "time"
"github.com/labstack/echo/v4" "github.com/labstack/echo/v5"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
) )
@@ -94,7 +94,7 @@ func TestRateLimiterSecurity(t *testing.T) {
rl := ratelimit.NewRateLimiter(config) rl := ratelimit.NewRateLimiter(config)
rateLimitMiddleware := ratelimit.RateLimiterMiddleware(rl) rateLimitMiddleware := ratelimit.RateLimiterMiddleware(rl)
handler := func(c echo.Context) error { handler := func(c *echo.Context) error {
return c.String(http.StatusOK, "ok") return c.String(http.StatusOK, "ok")
} }
+8 -7
View File
@@ -26,8 +26,8 @@ import (
"github.com/google/uuid" "github.com/google/uuid"
"github.com/jackc/pgx/v5/pgtype" "github.com/jackc/pgx/v5/pgtype"
"github.com/jackc/pgx/v5/pgxpool" "github.com/jackc/pgx/v5/pgxpool"
"github.com/labstack/echo/v4" "github.com/labstack/echo/v5"
echomiddleware "github.com/labstack/echo/v4/middleware" echomiddleware "github.com/labstack/echo/v5/middleware"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
) )
@@ -486,7 +486,7 @@ func setupTestServer(t *testing.T) *TestServerSetup {
e.Validator = &CustomValidator{validator: v} e.Validator = &CustomValidator{validator: v}
// Middleware // Middleware
e.Use(echomiddleware.Logger()) e.Use(echomiddleware.RequestLogger())
e.Use(echomiddleware.Recover()) e.Use(echomiddleware.Recover())
e.Use(echomiddleware.CORS()) e.Use(echomiddleware.CORS())
@@ -523,12 +523,13 @@ func setupTestServer(t *testing.T) *TestServerSetup {
ln, err := net.Listen("tcp", "127.0.0.1:0") ln, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err, "Failed to create listener") require.NoError(t, err, "Failed to create listener")
// Configure Echo's HTTP server with the listener // Configure Echo's HTTP server with the listener
e.Server.Handler = e serverConfig := &http.Server{
e.Server.Addr = ln.Addr().String() Handler: e,
// Create test server using Echo's server config (supports WebSocket hijacking) Addr: ln.Addr().String(),
}
ts := &httptest.Server{ ts := &httptest.Server{
Listener: ln, Listener: ln,
Config: e.Server, Config: serverConfig,
} }
ts.Start() ts.Start()
+106 -90
View File
@@ -2,8 +2,10 @@ package main
import ( import (
"bookhoard/internal/database" "bookhoard/internal/database"
"bytes"
"context" "context"
"encoding/json" "encoding/json"
"fmt"
"net/http" "net/http"
"strings" "strings"
"testing" "testing"
@@ -68,114 +70,128 @@ func TestWebSocketDeviceAuth(t *testing.T) {
}) })
require.NoError(t, err) require.NoError(t, err)
// Connect to WebSocket with device token // Use gorilla/websocket's RequestHeader support
wsURL := strings.Replace(setup.Server.URL, "http", "ws", 1) + "/ws/sync?token=device-auth-test" wsURL := strings.Replace(setup.Server.URL, "http", "ws", 1) + "/ws/sync?token=device-auth-test"
req, _ := http.NewRequest("GET", wsURL, nil) // Create dialer with custom headers
req.Header.Set("Authorization", "Bearer test-device-token-"+deviceID.String()) dialer := &websocket.Dialer{
HandshakeTimeout: 5 * time.Second,
// We can't easily test WebSocket with custom headers using gorilla/websocket }
// So this test just verifies the device exists headers := http.Header{}
device, err := setup.DB.GetDeviceByAuthToken(context.Background(), "test-device-token-"+deviceID.String()) headers.Set("Authorization", "Bearer test-device-token-"+deviceID.String())
// Connect with device token in header
ws, resp, err := dialer.Dial(wsURL, headers)
require.NoError(t, err, "WebSocket connection with device token should succeed")
defer ws.Close()
if resp != nil {
defer resp.Body.Close()
require.Equal(t, http.StatusSwitchingProtocols, resp.StatusCode, "Should upgrade to WebSocket")
}
// Read initial state message
ws.SetReadDeadline(time.Now().Add(5 * time.Second))
_, msg, err := ws.ReadMessage()
require.NoError(t, err, "Should receive initial state message")
var initialMsg map[string]interface{}
err = json.Unmarshal(msg, &initialMsg)
require.NoError(t, err) require.NoError(t, err)
assert.Equal(t, "Test KOReader", device.DeviceName) assert.Equal(t, "initial_state", initialMsg["type"])
} }
// TestWebSocketProgressBroadcast tests that progress updates are broadcast to connected clients // TestWebSocketProgressBroadcast tests that progress updates are broadcast to connected clients
/* func TestWebSocketProgressBroadcast(t *testing.T) { func TestWebSocketProgressBroadcast(t *testing.T) {
setup := setupTestServer(t) setup := setupTestServer(t)
// Get JWT token // Get JWT token
token := setup.Token token := setup.Token
// Create a test media item // Create a test media item
userID := getTestUserID(t, setup.DB) userID := getTestUserID(t, setup.DB)
mediaID := createTestMediaItem(t, setup.DB, userID) mediaID := createTestMediaItem(t, setup.DB, userID)
// Connect WebSocket client // Connect WebSocket client
wsURL := strings.Replace(setup.Server.URL, "http", "ws", 1) + "/ws/sync?token=" + token wsURL := strings.Replace(setup.Server.URL, "http", "ws", 1) + "/ws/sync?token=" + token
ws, _, err := websocket.DefaultDialer.Dial(wsURL, nil) ws, _, err := websocket.DefaultDialer.Dial(wsURL, nil)
require.NoError(t, err) require.NoError(t, err)
defer ws.Close() defer ws.Close()
// Read and discard initial state message // Read and discard initial state message
ws.SetReadDeadline(time.Now().Add(5 * time.Second)) ws.SetReadDeadline(time.Now().Add(5 * time.Second))
_, _, _ = ws.ReadMessage() _, _, _ = ws.ReadMessage()
// Update progress via HTTP API // Update progress via HTTP API
progressReq := map[string]interface{}{ progressReq := map[string]interface{}{
"source": "test", "source": "test",
"location": map[string]interface{}{ "location": map[string]interface{}{
"percentage": 0.5, "percentage": 0.5,
}, },
"device_metadata": map[string]interface{}{ "device_metadata": map[string]interface{}{
"device_type": "web", "device_type": "web",
}, },
}
body, _ := json.Marshal(progressReq)
req, _ := http.NewRequest("POST", setup.Server.URL+"/api/progress/"+mediaID, strings.NewReader(string(body)))
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Authorization", "Bearer "+token)
client := &http.Client{}
resp, err := client.Do(req)
require.NoError(t, err)
defer resp.Body.Close()
assert.Equal(t, http.StatusOK, resp.StatusCode)
// Read the broadcast message from WebSocket
_, msg, err := ws.ReadMessage()
require.NoError(t, err, "Failed to read broadcast message")
var broadcastMsg map[string]interface{}
err = json.Unmarshal(msg, &broadcastMsg)
require.NoError(t, err)
assert.Equal(t, "progress_update", broadcastMsg["type"])
data := broadcastMsg["data"].(map[string]interface{})
assert.Equal(t, mediaID, data["book_id"])
assert.InDelta(t, 0.5, data["percentage"], 0.01)
sourceDevice := broadcastMsg["source_device"].(map[string]interface{})
assert.Equal(t, "web", sourceDevice["type"])
} }
body, _ := json.Marshal(progressReq)
req, _ := http.NewRequest("POST", setup.Server.URL+"/api/progress/"+mediaID, strings.NewReader(string(body)))
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Authorization", "Bearer "+token)
client := &http.Client{}
resp, err := client.Do(req)
require.NoError(t, err)
defer resp.Body.Close()
assert.Equal(t, http.StatusOK, resp.StatusCode)
// Read the broadcast message from WebSocket
_, msg, err := ws.ReadMessage()
require.NoError(t, err, "Failed to read broadcast message")
var broadcastMsg map[string]interface{}
err = json.Unmarshal(msg, &broadcastMsg)
require.NoError(t, err)
assert.Equal(t, "progress_update", broadcastMsg["type"])
data := broadcastMsg["data"].(map[string]interface{})
assert.Equal(t, mediaID, data["book_id"])
assert.InDelta(t, 0.5, data["percentage"], 0.01)
sourceDevice := broadcastMsg["source_device"].(map[string]interface{})
assert.Equal(t, "web", sourceDevice["type"])
} */
// TestWebSocketPingPong tests that ping/pong messages work correctly // TestWebSocketPingPong tests that ping/pong messages work correctly
/* func TestWebSocketPingPong(t *testing.T) { func TestWebSocketPingPong(t *testing.T) {
setup := setupTestServer(t) setup := setupTestServer(t)
token := setup.Token token := setup.Token
wsURL := strings.Replace(setup.Server.URL, "http", "ws", 1) + "/ws/sync?token=" + token wsURL := strings.Replace(setup.Server.URL, "http", "ws", 1) + "/ws/sync?token=" + token
ws, _, err := websocket.DefaultDialer.Dial(wsURL, nil) ws, _, err := websocket.DefaultDialer.Dial(wsURL, nil)
require.NoError(t, err) require.NoError(t, err)
defer ws.Close() defer ws.Close()
// Read initial message // Read initial message
ws.SetReadDeadline(time.Now().Add(5 * time.Second)) ws.SetReadDeadline(time.Now().Add(5 * time.Second))
_, _, _ = ws.ReadMessage() _, _, _ = ws.ReadMessage()
// Send a ping message (as a text message for testing) // Send a ping message (as a text message for testing)
err = ws.WriteMessage(websocket.TextMessage, []byte(`{"type":"ping"}`)) err = ws.WriteMessage(websocket.TextMessage, []byte(`{"type":"ping"}`))
require.NoError(t, err) require.NoError(t, err)
// Server should respond with pong // Server should respond with pong
ws.SetReadDeadline(time.Now().Add(2 * time.Second)) ws.SetReadDeadline(time.Now().Add(2 * time.Second))
_, msg, err := ws.ReadMessage() _, msg, err := ws.ReadMessage()
if err == nil {
var pongMsg map[string]interface{}
err = json.Unmarshal(msg, &pongMsg)
if err == nil { if err == nil {
// Server might respond with pong var pongMsg map[string]interface{}
assert.Equal(t, "pong", pongMsg["type"]) err = json.Unmarshal(msg, &pongMsg)
if err == nil {
// Server might respond with pong
assert.Equal(t, "pong", pongMsg["type"])
}
} }
} }
} */
// TestWebSocketConnectionLimit tests that the server handles multiple connections // TestWebSocketConnectionLimit tests that the server handles multiple connections
/* func TestWebSocketConnectionLimit(t *testing.T) { func TestWebSocketConnectionLimit(t *testing.T) {
setup := setupTestServer(t) setup := setupTestServer(t)
token := setup.Token token := setup.Token
@@ -197,9 +213,9 @@ if err == nil {
for _, ws := range connections { for _, ws := range connections {
ws.Close() ws.Close()
} }
} */ }
/* // TestWebSocketInvalidToken tests that invalid tokens are rejected // TestWebSocketInvalidToken tests that invalid tokens are rejected
func TestWebSocketInvalidToken(t *testing.T) { func TestWebSocketInvalidToken(t *testing.T) {
setup := setupTestServer(t) setup := setupTestServer(t)
@@ -215,7 +231,7 @@ func TestWebSocketInvalidToken(t *testing.T) {
return return
} }
assert.Error(t, err) assert.Error(t, err)
} */ }
// Helper function to create a test media item // Helper function to create a test media item
func createTestMediaItem(t *testing.T, db *database.Queries, userID uuid.UUID) string { func createTestMediaItem(t *testing.T, db *database.Queries, userID uuid.UUID) string {
@@ -241,7 +257,7 @@ func createTestMediaItem(t *testing.T, db *database.Queries, userID uuid.UUID) s
return uuid.UUID(mediaID.ID.Bytes).String() return uuid.UUID(mediaID.ID.Bytes).String()
} }
/* // TestWebSocketUserScopedBroadcast tests that broadcasts only go to the user who made changes // TestWebSocketUserScopedBroadcast tests that broadcasts only go to the user who made changes
func TestWebSocketUserScopedBroadcast(t *testing.T) { func TestWebSocketUserScopedBroadcast(t *testing.T) {
setup := setupTestServer(t) setup := setupTestServer(t)
client := &http.Client{} client := &http.Client{}
@@ -307,7 +323,7 @@ func TestWebSocketUserScopedBroadcast(t *testing.T) {
if err == nil { if err == nil {
t.Errorf("Regular user should not receive collection_updated message") t.Errorf("Regular user should not receive collection_updated message")
} }
} */ }
// Helper: connectWebSocketToServer establishes WebSocket connection with auth token // Helper: connectWebSocketToServer establishes WebSocket connection with auth token
func connectWebSocketToServer(t *testing.T, serverURL string, token string) *websocket.Conn { func connectWebSocketToServer(t *testing.T, serverURL string, token string) *websocket.Conn {