| """ |
| WebSocket connection lifecycle tests (Wave 1, Task 1.1). |
| |
| Tests cover: |
| - WebSocket connection establishment |
| - WebSocket authentication |
| - WebSocket with valid token |
| - WebSocket disconnect |
| - WebSocket reconnection |
| """ |
| import asyncio |
| import json |
| import pytest |
| from datetime import datetime, timedelta |
| from unittest.mock import AsyncMock, MagicMock |
|
|
| from core.auth import create_access_token |
| from tests.property_tests.conftest import db_session |
|
|
|
|
| |
| |
| |
|
|
| @pytest.fixture |
| def mock_websocket(): |
| """Create a mock WebSocket for testing.""" |
| ws = MagicMock() |
| ws.accept = AsyncMock() |
| ws.send_json = AsyncMock() |
| ws.send_text = AsyncMock() |
| ws.receive_json = AsyncMock() |
| ws.receive_text = AsyncMock() |
| ws.close = AsyncMock() |
| return ws |
|
|
|
|
| @pytest.fixture |
| def cleanup_websocket_manager(): |
| """Cleanup WebSocket manager state before/after tests.""" |
| from core.websockets import manager |
| |
| original_connections = manager.active_connections.copy() |
| original_user_connections = manager.user_connections.copy() |
|
|
| yield |
|
|
| |
| manager.active_connections.clear() |
| manager.user_connections.clear() |
| manager.active_connections.update(original_connections) |
| manager.user_connections.update(original_user_connections) |
|
|
|
|
| |
| |
| |
|
|
| class TestWebSocketConnectionEstablishment: |
| """Test WebSocket connection can be established successfully.""" |
|
|
| @pytest.mark.asyncio(mode="auto") |
| async def test_websocket_connection_accept(self, mock_websocket, cleanup_websocket_manager): |
| """Test WebSocket connection is accepted.""" |
| from core.websockets import manager |
|
|
| |
| token = "dev-token" |
|
|
| |
| connected_user = await manager.connect(mock_websocket, token) |
|
|
| |
| assert connected_user is not None |
| assert connected_user.id == "dev-user" |
| mock_websocket.accept.assert_called_once() |
|
|
| @pytest.mark.asyncio(mode="auto") |
| async def test_websocket_connection_sends_welcome_message(self, mock_websocket, cleanup_websocket_manager): |
| """Test WebSocket sends welcome message on connection.""" |
| from core.websockets import manager |
|
|
| |
| token = "dev-token" |
|
|
| |
| await manager.connect(mock_websocket, token) |
|
|
| |
| |
| mock_websocket.accept.assert_called_once() |
|
|
| @pytest.mark.asyncio(mode="auto") |
| async def test_websocket_connection_registers_user(self, mock_websocket, cleanup_websocket_manager): |
| """Test WebSocket connection registers user in connections.""" |
| from core.websockets import manager |
|
|
| |
| token = "dev-token" |
|
|
| |
| connected_user = await manager.connect(mock_websocket, token) |
|
|
| |
| assert connected_user.id in manager.user_connections |
| assert mock_websocket in manager.user_connections[connected_user.id] |
|
|
| @pytest.mark.asyncio(mode="auto") |
| async def test_websocket_connection_auto_subscribes_to_user_channel(self, mock_websocket, cleanup_websocket_manager): |
| """Test WebSocket auto-subscribes to user channel on connection.""" |
| from core.websockets import manager |
|
|
| |
| token = "dev-token" |
|
|
| |
| connected_user = await manager.connect(mock_websocket, token) |
|
|
| |
| user_channel = f"user:{connected_user.id}" |
| assert user_channel in manager.active_connections |
| assert mock_websocket in manager.active_connections[user_channel] |
|
|
| @pytest.mark.asyncio(mode="auto") |
| async def test_websocket_connection_auto_subscribes_to_workspace_channel(self, mock_websocket, cleanup_websocket_manager): |
| """Test WebSocket auto-subscribes to workspace channel on connection.""" |
| from core.websockets import manager |
|
|
| |
| token = "dev-token" |
|
|
| |
| connected_user = await manager.connect(mock_websocket, token) |
|
|
| |
| workspace_channel = f"workspace:{connected_user.workspace_id}" |
| assert workspace_channel in manager.active_connections |
| assert mock_websocket in manager.active_connections[workspace_channel] |
|
|
|
|
| |
| |
| |
|
|
| class TestWebSocketAuthentication: |
| """Test WebSocket authentication enforcement.""" |
|
|
| @pytest.mark.asyncio(mode="auto") |
| async def test_websocket_requires_authentication(self, mock_websocket, cleanup_websocket_manager): |
| """Test WebSocket connection requires valid authentication.""" |
| from core.websockets import manager |
|
|
| |
| invalid_token = "invalid.jwt.token" |
|
|
| |
| result = await manager.connect(mock_websocket, invalid_token) |
|
|
| |
| assert result is None |
| mock_websocket.close.assert_called_once_with(code=4001) |
|
|
| @pytest.mark.asyncio(mode="auto") |
| async def test_websocket_rejects_expired_token(self, db_session, mock_websocket, cleanup_websocket_manager): |
| """Test WebSocket rejects expired JWT token.""" |
| from core.websockets import manager |
| from tests.factories.user_factory import UserFactory |
| from core.auth import get_password_hash |
| try: |
| from freezegun import freeze_time |
| except ImportError: |
| pytest.skip("freezegun not available") |
|
|
| |
| user = UserFactory( |
| email="ws_expired@example.com", |
| password_hash=get_password_hash("password123"), |
| _session=db_session |
| ) |
| db_session.add(user) |
| db_session.commit() |
|
|
| |
| with freeze_time("2026-02-01 10:00:00"): |
| token = create_access_token( |
| data={"sub": str(user.id)}, |
| expires_delta=timedelta(minutes=15) |
| ) |
|
|
| |
| with freeze_time("2026-02-01 11:00:00"): |
| result = await manager.connect(mock_websocket, token) |
|
|
| |
| assert result is None |
|
|
| @pytest.mark.asyncio(mode="auto") |
| async def test_websocket_accepts_valid_token(self, db_session, mock_websocket, cleanup_websocket_manager): |
| """Test WebSocket accepts valid JWT token.""" |
| from core.websockets import manager |
| from tests.factories.user_factory import UserFactory |
|
|
| |
| user = UserFactory(email="ws_valid@example.com", _session=db_session) |
| db_session.add(user) |
| db_session.commit() |
|
|
| token = create_access_token(data={"sub": str(user.id)}) |
|
|
| |
| result = await manager.connect(mock_websocket, token) |
|
|
| |
| assert result is not None |
| assert result.id == user.id |
| mock_websocket.accept.assert_called_once() |
|
|
| @pytest.mark.asyncio(mode="auto") |
| async def test_websocket_dev_token_bypass_in_non_production(self, mock_websocket, cleanup_websocket_manager, monkeypatch): |
| """Test WebSocket dev-token bypass in non-production environments.""" |
| from core.websockets import manager |
| import os |
|
|
| |
| monkeypatch.setenv("ENVIRONMENT", "development") |
|
|
| |
| result = await manager.connect(mock_websocket, "dev-token") |
|
|
| |
| assert result is not None |
| assert result.id == "dev-user" |
| mock_websocket.accept.assert_called_once() |
|
|
|
|
| |
| |
| |
|
|
| class TestWebSocketDisconnect: |
| """Test WebSocket disconnect handling.""" |
|
|
| @pytest.mark.asyncio(mode="auto") |
| async def test_websocket_disconnect_removes_from_user_connections(self, mock_websocket, cleanup_websocket_manager): |
| """Test WebSocket disconnect removes user from user_connections.""" |
| from core.websockets import manager |
|
|
| |
| token = "dev-token" |
| connected_user = await manager.connect(mock_websocket, token) |
|
|
| |
| manager.disconnect(mock_websocket, connected_user.id) |
|
|
| |
| assert mock_websocket not in manager.user_connections.get(connected_user.id, []) |
|
|
| @pytest.mark.asyncio(mode="auto") |
| async def test_websocket_disconnect_removes_from_all_channels(self, mock_websocket, cleanup_websocket_manager): |
| """Test WebSocket disconnect removes connection from all channels.""" |
| from core.websockets import manager |
|
|
| |
| token = "dev-token" |
| connected_user = await manager.connect(mock_websocket, token) |
| manager.subscribe(mock_websocket, "extra_channel") |
|
|
| |
| manager.disconnect(mock_websocket, connected_user.id) |
|
|
| |
| for channel_connections in manager.active_connections.values(): |
| assert mock_websocket not in channel_connections |
|
|
| @pytest.mark.asyncio(mode="auto") |
| async def test_websocket_disconnect_handles_nonexistent_user(self, mock_websocket, cleanup_websocket_manager): |
| """Test WebSocket disconnect handles nonexistent user gracefully.""" |
| from core.websockets import manager |
|
|
| |
| nonexistent_user = "nonexistent_user_xyz" |
|
|
| |
| manager.disconnect(mock_websocket, nonexistent_user) |
|
|
| |
| assert True |
|
|
| @pytest.mark.asyncio(mode="auto") |
| async def test_websocket_disconnect_handles_empty_channels(self, mock_websocket, cleanup_websocket_manager): |
| """Test WebSocket disconnect cleans up empty channels.""" |
| from core.websockets import manager |
|
|
| |
| manager.subscribe(mock_websocket, "test_channel") |
| assert "test_channel" in manager.active_connections |
|
|
| |
| manager.disconnect(mock_websocket, "test_user") |
|
|
| |
| |
| |
| if mock_websocket not in manager.active_connections.get("test_channel", []): |
| |
| pass |
|
|
|
|
| |
| |
| |
|
|
| class TestWebSocketReconnection: |
| """Test WebSocket reconnection scenarios.""" |
|
|
| @pytest.mark.asyncio(mode="auto") |
| async def test_websocket_reconnection_after_disconnect(self, mock_websocket, cleanup_websocket_manager): |
| """Test WebSocket can reconnect after disconnect.""" |
| from core.websockets import manager |
|
|
| |
| token = "dev-token" |
|
|
| |
| user1 = await manager.connect(mock_websocket, token) |
| manager.disconnect(mock_websocket, user1.id) |
|
|
| |
| user2 = await manager.connect(mock_websocket, token) |
|
|
| |
| assert user2 is not None |
| assert user2.id == user1.id |
| mock_websocket.accept.assert_called() |
|
|
| @pytest.mark.asyncio(mode="auto") |
| async def test_websocket_multiple_simultaneous_connections(self, cleanup_websocket_manager): |
| """Test multiple WebSocket connections from same user.""" |
| from core.websockets import manager |
|
|
| |
| token = "dev-token" |
| ws1 = MagicMock() |
| ws1.accept = AsyncMock() |
| ws2 = MagicMock() |
| ws2.accept = AsyncMock() |
|
|
| |
| user1 = await manager.connect(ws1, token) |
| user2 = await manager.connect(ws2, token) |
|
|
| |
| assert user1 is not None |
| assert user2 is not None |
| assert user1.id == user2.id |
| assert len(manager.user_connections[user1.id]) == 2 |
|
|
| @pytest.mark.asyncio(mode="auto") |
| async def test_websocket_reconnection_maintains_subscriptions(self, mock_websocket, cleanup_websocket_manager): |
| """Test WebSocket reconnection can restore subscriptions.""" |
| from core.websockets import manager |
|
|
| |
| token = "dev-token" |
| user = await manager.connect(mock_websocket, token) |
| manager.subscribe(mock_websocket, "custom_channel") |
|
|
| |
| manager.disconnect(mock_websocket, user.id) |
| mock_websocket2 = MagicMock() |
| mock_websocket2.accept = AsyncMock() |
| await manager.connect(mock_websocket2, token) |
|
|
| |
| |
| |
| assert mock_websocket2 in manager.user_connections[user.id] |
|
|
|
|
| |
| |
| |
|
|
| class TestWebSocketConnectionLifecycleE2E: |
| """End-to-end WebSocket connection lifecycle tests.""" |
|
|
| @pytest.mark.asyncio(mode="auto") |
| async def test_full_websocket_connection_lifecycle(self, cleanup_websocket_manager): |
| """Test complete WebSocket connection lifecycle: connect -> use -> disconnect.""" |
| from core.websockets import manager |
|
|
| |
| ws = MagicMock() |
| ws.accept = AsyncMock() |
| ws.send_json = AsyncMock() |
| token = "dev-token" |
|
|
| |
| user = await manager.connect(ws, token) |
| assert user is not None |
|
|
| |
| await manager.send_personal_message(user.id, {"type": "test", "content": "Hello"}) |
| ws.send_json.assert_called_once() |
|
|
| |
| manager.disconnect(ws, user.id) |
|
|
| |
| assert ws not in manager.user_connections.get(user.id, []) |
|
|
| @pytest.mark.asyncio(mode="auto") |
| async def test_websocket_connection_error_handling(self, cleanup_websocket_manager): |
| """Test WebSocket connection handles errors gracefully.""" |
| from core.websockets import manager |
|
|
| |
| ws_broken = MagicMock() |
| ws_broken.accept = AsyncMock(side_effect=Exception("Connection failed")) |
| ws_broken.close = AsyncMock() |
|
|
| |
| result = await manager.connect(ws_broken, "dev-token") |
|
|
| |
| |
| assert result is None or ws_broken.close.called |
|
|