package main import ( "bookhoard/internal/database" "bytes" "context" "encoding/json" "fmt" "net/http" "strings" "testing" "time" "github.com/google/uuid" "github.com/gorilla/websocket" "github.com/jackc/pgx/v5/pgtype" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) // TestWebSocketConnection tests basic WebSocket connection and authentication /* func TestWebSocketConnection(t *testing.T) { // Setup test server with WebSocket setup := setupTestServer(t) // Get JWT token for a test user token := setup.Token // Connect to WebSocket endpoint wsURL := strings.Replace(setup.Server.URL, "http", "ws", 1) + "/ws/sync?token=" + token ws, _, err := websocket.DefaultDialer.Dial(wsURL, nil) require.NoError(t, err, "Failed to connect to WebSocket") defer ws.Close() // Set read deadline ws.SetReadDeadline(time.Now().Add(5 * time.Second)) // Wait for initial state message _, msg, err := ws.ReadMessage() require.NoError(t, err, "Failed to read initial message") var initialMsg map[string]interface{} err = json.Unmarshal(msg, &initialMsg) require.NoError(t, err) assert.Equal(t, "initial_state", initialMsg["type"]) assert.Contains(t, initialMsg["data"], "progress") assert.Contains(t, initialMsg["data"], "devices") } */ // TestWebSocketDeviceAuth tests device authentication via WebSocket func TestWebSocketDeviceAuth(t *testing.T) { setup := setupTestServer(t) // Create a test device userID := getTestUserID(t, setup.DB) deviceID := uuid.New() _, err := setup.DB.CreateDevice(context.Background(), database.CreateDeviceParams{ UserID: pgtype.UUID{Bytes: [16]byte(userID), Valid: true}, DeviceName: "Test KOReader", DeviceType: "koreader", DeviceIdentifier: "test-device-123", AuthToken: "test-device-token-" + deviceID.String(), SyncEnabled: pgtype.Bool{Bool: true, Valid: true}, AutoSync: pgtype.Bool{Bool: true, Valid: true}, SyncFrequencyMinutes: pgtype.Int4{Int32: 5, Valid: true}, DeviceMetadata: []byte("{}"), }) require.NoError(t, err) // Use gorilla/websocket's RequestHeader support wsURL := strings.Replace(setup.Server.URL, "http", "ws", 1) + "/ws/sync?token=device-auth-test" // Create dialer with custom headers dialer := &websocket.Dialer{ HandshakeTimeout: 5 * time.Second, } headers := http.Header{} 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) assert.Equal(t, "initial_state", initialMsg["type"]) } // TestWebSocketProgressBroadcast tests that progress updates are broadcast to connected clients func TestWebSocketProgressBroadcast(t *testing.T) { setup := setupTestServer(t) // Get JWT token token := setup.Token // Create a test media item userID := getTestUserID(t, setup.DB) mediaID := createTestMediaItem(t, setup.DB, userID) // Connect WebSocket client wsURL := strings.Replace(setup.Server.URL, "http", "ws", 1) + "/ws/sync?token=" + token ws, _, err := websocket.DefaultDialer.Dial(wsURL, nil) require.NoError(t, err) defer ws.Close() // Read and discard initial state message ws.SetReadDeadline(time.Now().Add(5 * time.Second)) _, _, _ = ws.ReadMessage() // Update progress via HTTP API progressReq := map[string]interface{}{ "source": "test", "location": map[string]interface{}{ "percentage": 0.5, }, "device_metadata": map[string]interface{}{ "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"]) } // TestWebSocketPingPong tests that ping/pong messages work correctly func TestWebSocketPingPong(t *testing.T) { setup := setupTestServer(t) token := setup.Token wsURL := strings.Replace(setup.Server.URL, "http", "ws", 1) + "/ws/sync?token=" + token ws, _, err := websocket.DefaultDialer.Dial(wsURL, nil) require.NoError(t, err) defer ws.Close() // Read initial message ws.SetReadDeadline(time.Now().Add(5 * time.Second)) _, _, _ = ws.ReadMessage() // Send a ping message (as a text message for testing) err = ws.WriteMessage(websocket.TextMessage, []byte(`{"type":"ping"}`)) require.NoError(t, err) // Server should respond with pong ws.SetReadDeadline(time.Now().Add(2 * time.Second)) _, msg, err := ws.ReadMessage() if err == nil { var pongMsg map[string]interface{} 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 func TestWebSocketConnectionLimit(t *testing.T) { setup := setupTestServer(t) token := setup.Token // Create multiple connections connections := make([]*websocket.Conn, 5) for i := 0; i < 5; i++ { wsURL := strings.Replace(setup.Server.URL, "http", "ws", 1) + fmt.Sprintf("/ws/sync?token=%s&conn=%d", token, i) ws, _, err := websocket.DefaultDialer.Dial(wsURL, nil) require.NoError(t, err, "Failed to create connection %d", i) connections[i] = ws // Read initial message ws.SetReadDeadline(time.Now().Add(5 * time.Second)) _, _, _ = ws.ReadMessage() } // Close all connections for _, ws := range connections { ws.Close() } } // TestWebSocketInvalidToken tests that invalid tokens are rejected func TestWebSocketInvalidToken(t *testing.T) { setup := setupTestServer(t) // Try to connect with invalid token wsURL := strings.Replace(setup.Server.URL, "http", "ws", 1) + "/ws/sync?token=invalid-token" _, _, err := websocket.DefaultDialer.Dial(wsURL, nil) assert.Error(t, err, "Expected error when connecting with invalid token") // Check if it's a websocket close error if websocket.IsCloseError(err, 1000, 1001, 1002, 1003, 1005, 1006, 1007, 1008, 1009, 1010, 1011) { // Expected close error return } assert.Error(t, err) } // Helper function to create a test media item func createTestMediaItem(t *testing.T, db *database.Queries, userID uuid.UUID) string { // First create a test library libID, err := db.CreateLibrary(context.Background(), database.CreateLibraryParams{ Name: "Test Library", LibraryTypeID: pgtype.UUID{Bytes: [16]byte(uuid.UUID{}), Valid: true}, CreatedByAdminID: pgtype.UUID{Bytes: [16]byte(userID), Valid: true}, }) require.NoError(t, err) // Create a test media item mediaID, err := db.CreateMediaItem(context.Background(), database.CreateMediaItemParams{ LibraryID: libID.ID, Title: "Test Book", FilePath: "/tmp/test.epub", FileSize: pgtype.Int8{Int64: 1024, Valid: true}, MimeType: pgtype.Text{String: "application/epub+zip", Valid: true}, AddedByAdminID: pgtype.UUID{Bytes: [16]byte(userID), Valid: true}, }) require.NoError(t, err) return uuid.UUID(mediaID.ID.Bytes).String() } // TestWebSocketUserScopedBroadcast tests that broadcasts only go to the user who made changes func TestWebSocketUserScopedBroadcast(t *testing.T) { setup := setupTestServer(t) client := &http.Client{} // Use pre-created admin user (setup.Token) and regular user (setup.RegularToken) // Both users already created by setupTestServer() // Create a collection for admin user collectionReq := map[string]interface{}{ "name": "Admin Collection", "description": "Test collection for WebSocket test", } collectionBody, _ := json.Marshal(collectionReq) collectionHTTP, _ := http.NewRequest("POST", setup.Server.URL+"/api/collections", bytes.NewBuffer(collectionBody)) collectionHTTP.Header.Set("Content-Type", "application/json") collectionHTTP.Header.Set("Authorization", "Bearer "+setup.Token) collectionResp, err := client.Do(collectionHTTP) require.NoError(t, err) defer collectionResp.Body.Close() require.Equal(t, http.StatusCreated, collectionResp.StatusCode) var collectionResult map[string]interface{} json.NewDecoder(collectionResp.Body).Decode(&collectionResult) collectionID := collectionResult["id"].(string) // Create a test book via API bookID := createTestMediaItemID(t, setup) // Connect admin user via WebSocket wsAdmin := connectWebSocketToServer(t, setup.Server.URL, setup.Token) defer wsAdmin.Close() // Connect regular user via WebSocket wsRegular := connectWebSocketToServer(t, setup.Server.URL, setup.RegularToken) defer wsRegular.Close() // Admin adds book to collection addReq := map[string]interface{}{ "book_ids": []string{bookID}, } addBody, _ := json.Marshal(addReq) addHTTP, _ := http.NewRequest("POST", setup.Server.URL+"/api/collections/"+collectionID+"/books", bytes.NewBuffer(addBody)) addHTTP.Header.Set("Content-Type", "application/json") addHTTP.Header.Set("Authorization", "Bearer "+setup.Token) addResp, err := client.Do(addHTTP) require.NoError(t, err) defer addResp.Body.Close() require.Equal(t, http.StatusNoContent, addResp.StatusCode) // Admin should receive collection_updated message msgAdmin := readWebSocketMessage(t, wsAdmin, 2*time.Second) if msgAdmin["type"] != "collection_updated" { t.Errorf("Admin should receive collection_updated, got %s", msgAdmin["type"]) } // Regular user should NOT receive collection_updated message wsRegular.SetReadDeadline(time.Now().Add(500 * time.Millisecond)) _, _, err = wsRegular.ReadMessage() if err == nil { t.Errorf("Regular user should not receive collection_updated message") } } // Helper: connectWebSocketToServer establishes WebSocket connection with auth token func connectWebSocketToServer(t *testing.T, serverURL string, token string) *websocket.Conn { wsURL := "ws" + strings.TrimPrefix(serverURL, "http") + "/ws/sync?token=" + token ws, resp, err := websocket.DefaultDialer.Dial(wsURL, nil) require.NoError(t, err, "WebSocket connection should succeed") if resp != nil { resp.Body.Close() } require.NotNil(t, ws, "WebSocket connection should be established") // Wait for connection to be ready time.Sleep(100 * time.Millisecond) return ws } // Helper: readWebSocketMessage reads a message from WebSocket with timeout func readWebSocketMessage(t *testing.T, ws *websocket.Conn, timeout time.Duration) map[string]interface{} { ws.SetReadDeadline(time.Now().Add(timeout)) _, message, err := ws.ReadMessage() require.NoError(t, err, "Should receive WebSocket message") var msg map[string]interface{} err = json.Unmarshal(message, &msg) require.NoError(t, err, "Message should be valid JSON") return msg }