from unittest.mock import AsyncMock, MagicMock import pytest from jam_session.models import QueueState from jam_session.ws import ConnectionManager @pytest.fixture def cm(): return ConnectionManager() @pytest.fixture def mock_ws(): ws = AsyncMock() ws.send_json = AsyncMock() return ws class TestConnectDisconnect: @pytest.mark.asyncio async def test_connect_returns_conn_id(self, cm, mock_ws): conn_id = await cm.connect(mock_ws) assert conn_id.startswith("conn_") assert conn_id in cm.active_connections @pytest.mark.asyncio async def test_disconnect_removes(self, cm, mock_ws): conn_id = await cm.connect(mock_ws) cm.disconnect(conn_id) assert conn_id not in cm.active_connections @pytest.mark.asyncio async def test_disconnect_unknown_id_no_error(self, cm): cm.disconnect("unknown") def test_connected_count(self, cm, mock_ws): assert cm.connected_count == 0 class TestBroadcast: @pytest.mark.asyncio async def test_sends_to_all_active(self, cm): ws1 = AsyncMock() ws2 = AsyncMock() ws1.send_json = AsyncMock() ws2.send_json = AsyncMock() await cm.connect(ws1) await cm.connect(ws2) await cm.broadcast({"type": "test", "payload": {}}) ws1.send_json.assert_called_once_with({"type": "test", "payload": {}}) ws2.send_json.assert_called_once_with({"type": "test", "payload": {}}) @pytest.mark.asyncio async def test_removes_dead_connections(self, cm): ws1 = AsyncMock() ws2 = AsyncMock() ws1.send_json = AsyncMock(side_effect=Exception("dead")) ws2.send_json = AsyncMock() conn1 = await cm.connect(ws1) conn2 = await cm.connect(ws2) await cm.broadcast({"type": "test", "payload": {}}) assert conn1 not in cm.active_connections assert conn2 in cm.active_connections ws2.send_json.assert_called_once() class TestSendTo: @pytest.mark.asyncio async def test_sends_to_specific_connection(self, cm, mock_ws): conn_id = await cm.connect(mock_ws) await cm.send_to(conn_id, {"type": "hello"}) mock_ws.send_json.assert_called_once_with({"type": "hello"}) @pytest.mark.asyncio async def test_missing_connection_no_error(self, cm): await cm.send_to("nobody", {"type": "hello"}) @pytest.mark.asyncio async def test_error_removes_connection(self, cm, mock_ws): mock_ws.send_json = AsyncMock(side_effect=Exception("dead")) conn_id = await cm.connect(mock_ws) await cm.send_to(conn_id, {"type": "hello"}) assert conn_id not in cm.active_connections