| """ |
| WebSocket Manager Edge Case Tests |
| |
| Comprehensive edge case testing for WebSocketConnectionManager. |
| Target: Fill remaining coverage gaps (97% -> 99%+). |
| |
| Tests cover: |
| - Connection lifecycle edge cases (reconnect, multiple disconnect, non-existent connections) |
| - Broadcast failure scenarios (all fail, partial failures, serialization errors) |
| - State transitions (connection states, stream lifecycle, manager state) |
| - Edge cases (empty streams, disconnected connections, race conditions) |
| |
| Uses AsyncMock patterns proven in test_websocket_manager_coverage.py. |
| """ |
|
|
| import pytest |
| import asyncio |
| from unittest.mock import AsyncMock, Mock |
| from fastapi import WebSocket |
| from datetime import datetime |
|
|
| from core.websocket_manager import ( |
| WebSocketConnectionManager, |
| DebuggingWebSocketManager, |
| ) |
|
|
|
|
| |
| |
| |
|
|
| @pytest.fixture |
| def mock_websocket(): |
| """Create mock WebSocket with AsyncMock for async methods.""" |
| mock_ws = Mock(spec=WebSocket) |
| mock_ws.accept = AsyncMock() |
| mock_ws.send_text = AsyncMock() |
| mock_ws.send_json = AsyncMock() |
| mock_ws.close = AsyncMock() |
| return mock_ws |
|
|
|
|
| @pytest.fixture |
| def manager(): |
| """Create fresh WebSocket manager for each test.""" |
| return WebSocketConnectionManager() |
|
|
|
|
| @pytest.fixture |
| async def manager_with_connections(manager): |
| """Create manager with pre-populated connections for state testing.""" |
| |
| for i in range(3): |
| mock_ws = Mock(spec=WebSocket) |
| mock_ws.accept = AsyncMock() |
| mock_ws.send_text = AsyncMock() |
| await manager.connect(mock_ws, "stream1") |
|
|
| |
| for i in range(2): |
| mock_ws = Mock(spec=WebSocket) |
| mock_ws.accept = AsyncMock() |
| mock_ws.send_text = AsyncMock() |
| await manager.connect(mock_ws, "stream2") |
|
|
| return manager |
|
|
|
|
| @pytest.fixture |
| def mock_ws_with_send_failure(): |
| """Create WebSocket that fails on alternate send_text calls.""" |
| mock_ws = Mock(spec=WebSocket) |
| mock_ws.accept = AsyncMock() |
|
|
| |
| call_count = [0] |
| async def failing_send(text): |
| call_count[0] += 1 |
| if call_count[0] % 2 == 0: |
| raise Exception("Alternating send failure") |
| |
|
|
| mock_ws.send_text = AsyncMock(side_effect=failing_send) |
| return mock_ws |
|
|
|
|
| |
| |
| |
|
|
| class TestWebSocketConnectionEdgeCases: |
| """Test WebSocket connection lifecycle edge cases.""" |
|
|
| @pytest.mark.asyncio |
| async def test_connect_after_disconnect(self, manager, mock_websocket): |
| """Test reconnecting same WebSocket after disconnect.""" |
| |
| await manager.connect(mock_websocket, stream_id="stream1") |
| assert mock_websocket in manager.active_connections["stream1"] |
|
|
| |
| manager.disconnect(mock_websocket) |
| assert mock_websocket not in manager.connection_streams |
|
|
| |
| await manager.connect(mock_websocket, stream_id="stream1") |
| assert mock_websocket in manager.active_connections["stream1"] |
| assert manager.connection_streams[mock_websocket] == "stream1" |
|
|
| @pytest.mark.asyncio |
| async def test_multiple_disconnect_calls(self, manager, mock_websocket): |
| """Test that calling disconnect() twice doesn't error.""" |
| await manager.connect(mock_websocket, stream_id="stream1") |
|
|
| |
| manager.disconnect(mock_websocket) |
| assert mock_websocket not in manager.connection_streams |
|
|
| |
| manager.disconnect(mock_websocket) |
| assert mock_websocket not in manager.connection_streams |
|
|
| @pytest.mark.asyncio |
| async def test_disconnect_non_existent_connection(self, manager): |
| """Test disconnecting a WebSocket that was never connected.""" |
| mock_ws = Mock(spec=WebSocket) |
| mock_ws.accept = AsyncMock() |
|
|
| |
| manager.disconnect(mock_ws) |
| assert mock_ws not in manager.connection_streams |
|
|
| @pytest.mark.asyncio |
| async def test_connect_to_multiple_streams(self, manager, mock_websocket): |
| """Test same WebSocket connecting to multiple streams (last wins).""" |
| |
| await manager.connect(mock_websocket, stream_id="stream1") |
| assert manager.connection_streams[mock_websocket] == "stream1" |
| assert mock_websocket in manager.active_connections["stream1"] |
|
|
| |
| |
| |
| await manager.connect(mock_websocket, stream_id="stream2") |
| assert manager.connection_streams[mock_websocket] == "stream2" |
| |
| assert mock_websocket in manager.active_connections["stream1"] |
| assert mock_websocket in manager.active_connections["stream2"] |
|
|
| @pytest.mark.asyncio |
| async def test_empty_stream_cleanup(self, manager, mock_websocket): |
| """Test that stream is removed when last connection disconnects.""" |
| await manager.connect(mock_websocket, stream_id="temp_stream") |
| assert "temp_stream" in manager.active_connections |
|
|
| |
| manager.disconnect(mock_websocket) |
|
|
| |
| assert "temp_stream" not in manager.active_connections |
|
|
| @pytest.mark.asyncio |
| async def test_connection_info_persistence(self, manager, mock_websocket): |
| """Test that connection_info persists across operations.""" |
| await manager.connect(mock_websocket, stream_id="stream1") |
|
|
| |
| original_info = manager.connection_info[mock_websocket] |
| assert "stream_id" in original_info |
| assert "connected_at" in original_info |
|
|
| |
| await manager.send_personal(mock_websocket, {"type": "test"}) |
| await manager.broadcast("stream1", {"type": "broadcast"}) |
|
|
| |
| assert manager.connection_info[mock_websocket] == original_info |
|
|
| @pytest.mark.asyncio |
| async def test_send_to_disconnected_connection(self, manager, mock_websocket): |
| """Test that send_personal() handles closed connection gracefully.""" |
| await manager.connect(mock_websocket, stream_id="stream1") |
|
|
| |
| manager.disconnect(mock_websocket) |
|
|
| |
| mock_websocket.send_text = AsyncMock(side_effect=Exception("Connection closed")) |
| success = await manager.send_personal(mock_websocket, {"type": "test"}) |
|
|
| |
| assert success is False |
| assert mock_websocket not in manager.connection_streams |
|
|
| @pytest.mark.asyncio |
| async def test_broadcast_with_no_connections(self, manager): |
| """Test broadcast on empty stream returns 0.""" |
| sent_count = await manager.broadcast("empty_stream", {"type": "test"}) |
| assert sent_count == 0 |
|
|
|
|
| |
| |
| |
|
|
| class TestWebSocketBroadcastEdgeCases: |
| """Test WebSocket broadcast edge cases and failure scenarios.""" |
|
|
| @pytest.mark.asyncio |
| async def test_broadcast_all_connections_fail(self, manager): |
| """Test broadcast when all connections fail.""" |
| |
| connections = [] |
| for i in range(3): |
| mock_ws = Mock(spec=WebSocket) |
| mock_ws.accept = AsyncMock() |
| mock_ws.send_text = AsyncMock(side_effect=Exception("All failed")) |
| connections.append(mock_ws) |
| await manager.connect(mock_ws, "stream1") |
|
|
| |
| sent_count = await manager.broadcast("stream1", {"type": "test"}) |
|
|
| |
| assert sent_count == 0 |
|
|
| |
| for conn in connections: |
| assert conn not in manager.active_connections.get("stream1", set()) |
|
|
| @pytest.mark.asyncio |
| async def test_broadcast_partial_failures_continue(self, manager): |
| """Test that broadcast continues even when some connections fail.""" |
| |
| mock_ws_success = Mock(spec=WebSocket) |
| mock_ws_success.accept = AsyncMock() |
| mock_ws_success.send_text = AsyncMock() |
|
|
| mock_ws_fail1 = Mock(spec=WebSocket) |
| mock_ws_fail1.accept = AsyncMock() |
| mock_ws_fail1.send_text = AsyncMock(side_effect=Exception("Failed 1")) |
|
|
| mock_ws_fail2 = Mock(spec=WebSocket) |
| mock_ws_fail2.accept = AsyncMock() |
| mock_ws_fail2.send_text = AsyncMock(side_effect=Exception("Failed 2")) |
|
|
| await manager.connect(mock_ws_success, "stream1") |
| await manager.connect(mock_ws_fail1, "stream1") |
| await manager.connect(mock_ws_fail2, "stream1") |
|
|
| |
| sent_count = await manager.broadcast("stream1", {"type": "test"}) |
|
|
| |
| assert sent_count == 1 |
|
|
| |
| assert mock_ws_fail1 not in manager.active_connections.get("stream1", set()) |
| assert mock_ws_fail2 not in manager.active_connections.get("stream1", set()) |
|
|
| |
| assert mock_ws_success in manager.active_connections["stream1"] |
|
|
| @pytest.mark.asyncio |
| async def test_broadcast_with_json_serialization_error(self, manager): |
| """Test broadcast with non-serializable data (handled by json.dumps).""" |
| mock_ws = Mock(spec=WebSocket) |
| mock_ws.accept = AsyncMock() |
| mock_ws.send_text = AsyncMock() |
|
|
| await manager.connect(mock_ws, "stream1") |
|
|
| |
| sent_count = await manager.broadcast("stream1", {"type": "test", "data": {"nested": "value"}}) |
| assert sent_count == 1 |
|
|
| @pytest.mark.asyncio |
| async def test_broadcast_with_mixed_connection_types(self, manager): |
| """Test broadcast with connections in different WebSocket states.""" |
| |
| mock_ws_healthy = Mock(spec=WebSocket) |
| mock_ws_healthy.accept = AsyncMock() |
| mock_ws_healthy.send_text = AsyncMock() |
|
|
| mock_ws_closed = Mock(spec=WebSocket) |
| mock_ws_closed.accept = AsyncMock() |
| mock_ws_closed.send_text = AsyncMock(side_effect=Exception("Connection closed")) |
|
|
| await manager.connect(mock_ws_healthy, "stream1") |
| await manager.connect(mock_ws_closed, "stream1") |
|
|
| |
| sent_count = await manager.broadcast("stream1", {"type": "test"}) |
|
|
| |
| assert sent_count == 1 |
| assert mock_ws_healthy in manager.active_connections["stream1"] |
| assert mock_ws_closed not in manager.active_connections.get("stream1", set()) |
|
|
| @pytest.mark.asyncio |
| async def test_broadcast_race_condition(self, manager): |
| """Test that concurrent broadcasts don't corrupt state.""" |
| |
| connections = [] |
| for i in range(5): |
| mock_ws = Mock(spec=WebSocket) |
| mock_ws.accept = AsyncMock() |
| mock_ws.send_text = AsyncMock() |
| connections.append(mock_ws) |
| await manager.connect(mock_ws, "stream1") |
|
|
| |
| tasks = [ |
| manager.broadcast("stream1", {"type": "broadcast", "id": i}) |
| for i in range(10) |
| ] |
|
|
| results = await asyncio.gather(*tasks) |
|
|
| |
| assert all(r == 5 for r in results) |
|
|
| |
| assert len(manager.active_connections["stream1"]) == 5 |
| for conn in connections: |
| assert conn in manager.active_connections["stream1"] |
|
|
|
|
| |
| |
| |
|
|
| class TestWebSocketStateTransitions: |
| """Test WebSocket state transition edge cases.""" |
|
|
| @pytest.mark.asyncio |
| async def test_connection_state_transitions(self, manager, mock_websocket): |
| """Test connection state: New -> Connected -> Disconnected.""" |
| |
| assert mock_websocket not in manager.connection_streams |
| assert mock_websocket not in manager.connection_info |
|
|
| |
| await manager.connect(mock_websocket, stream_id="stream1") |
| assert mock_websocket in manager.connection_streams |
| assert manager.connection_streams[mock_websocket] == "stream1" |
| assert mock_websocket in manager.connection_info |
|
|
| |
| manager.disconnect(mock_websocket) |
| assert mock_websocket not in manager.connection_streams |
| assert mock_websocket not in manager.connection_info |
|
|
| @pytest.mark.asyncio |
| async def test_stream_state_empty_to_populated(self, manager): |
| """Test stream state created/deleted with connections.""" |
| |
| assert "new_stream" not in manager.active_connections |
|
|
| |
| mock_ws1 = Mock(spec=WebSocket) |
| mock_ws1.accept = AsyncMock() |
| mock_ws1.send_text = AsyncMock() |
| await manager.connect(mock_ws1, "new_stream") |
|
|
| |
| assert "new_stream" in manager.active_connections |
| assert len(manager.active_connections["new_stream"]) == 1 |
|
|
| |
| mock_ws2 = Mock(spec=WebSocket) |
| mock_ws2.accept = AsyncMock() |
| mock_ws2.send_text = AsyncMock() |
| await manager.connect(mock_ws2, "new_stream") |
|
|
| |
| assert len(manager.active_connections["new_stream"]) == 2 |
|
|
| |
| manager.disconnect(mock_ws1) |
| assert len(manager.active_connections["new_stream"]) == 1 |
|
|
| |
| manager.disconnect(mock_ws2) |
|
|
| |
| assert "new_stream" not in manager.active_connections |
|
|
| @pytest.mark.asyncio |
| async def test_manager_state_after_all_disconnect(self): |
| """Test manager returns to initial state after all disconnects.""" |
| |
| manager = WebSocketConnectionManager() |
|
|
| |
| stream1_conns = [] |
| for i in range(3): |
| mock_ws = Mock(spec=WebSocket) |
| mock_ws.accept = AsyncMock() |
| mock_ws.send_text = AsyncMock() |
| await manager.connect(mock_ws, "stream1") |
| stream1_conns.append(mock_ws) |
|
|
| |
| stream2_conns = [] |
| for i in range(2): |
| mock_ws = Mock(spec=WebSocket) |
| mock_ws.accept = AsyncMock() |
| mock_ws.send_text = AsyncMock() |
| await manager.connect(mock_ws, "stream2") |
| stream2_conns.append(mock_ws) |
|
|
| |
| assert len(manager.active_connections) == 2 |
| assert manager.get_connection_count("stream1") == 3 |
| assert manager.get_connection_count("stream2") == 2 |
|
|
| |
| for conn in stream1_conns: |
| manager.disconnect(conn) |
|
|
| |
| assert "stream1" not in manager.active_connections |
|
|
| |
| for conn in stream2_conns: |
| manager.disconnect(conn) |
|
|
| |
| assert "stream2" not in manager.active_connections |
|
|
| |
| assert len(manager.active_connections) == 0 |
| |
| |
| assert len(manager.connection_streams) == 0 |
| assert len(manager.connection_info) == 0 |
|
|
| @pytest.mark.asyncio |
| async def test_connection_metadata_updates(self, manager, mock_websocket): |
| """Test that connection metadata is updated during lifecycle.""" |
| |
| await manager.connect(mock_websocket, stream_id="stream1") |
|
|
| |
| info = manager.connection_info[mock_websocket] |
| assert info["stream_id"] == "stream1" |
| assert "connected_at" in info |
|
|
| |
| datetime.fromisoformat(info["connected_at"]) |
|
|
| @pytest.mark.asyncio |
| async def test_singleton_state_persistence(self): |
| """Test that singleton maintains state across access.""" |
| from core.websocket_manager import get_websocket_manager |
|
|
| |
| mgr1 = get_websocket_manager() |
|
|
| |
| mock_ws = Mock(spec=WebSocket) |
| mock_ws.accept = AsyncMock() |
| mock_ws.send_text = AsyncMock() |
| await mgr1.connect(mock_ws, "singleton_test") |
|
|
| |
| mgr2 = get_websocket_manager() |
|
|
| |
| assert mgr1 is mgr2 |
| assert "singleton_test" in mgr2.active_connections |
| assert mock_ws in mgr2.active_connections["singleton_test"] |
|
|
| |
| mgr2.disconnect(mock_ws) |
|
|