from __future__ import annotations import asyncio from unittest.mock import AsyncMock, MagicMock, patch import pytest from websockets.exceptions import ConnectionClosed from yuxi.channel.transport.protocol import TransportState from yuxi.channel.transport.websocket import WebSocketTransport class TestWebSocketTransport: def test_initial_state(self): transport = WebSocketTransport("ws://example.com", headers={"X-Token": "t"}) assert transport.url == "ws://example.com" assert transport.headers == {"X-Token": "t"} assert transport.heartbeat_interval == 30.0 assert transport.heartbeat_message == b"ping" assert transport.state == TransportState.DISCONNECTED assert transport._ws is None assert transport._closed is True assert transport._tasks == set() def test_default_headers(self): transport = WebSocketTransport("ws://example.com") assert transport.headers == {} async def test_spawn_task_tracks_and_cleans_up(self): transport = WebSocketTransport("ws://example.com") async def coro(): pass task = transport._spawn_task(coro()) assert task in transport._tasks await task assert task not in transport._tasks async def test_start_connects_and_spawns_loops(self): mock_ws = MagicMock() mock_ws.close = AsyncMock() mock_ws.send = AsyncMock() mock_ws.recv = AsyncMock(return_value=b"raw") with patch( "yuxi.channel.transport.websocket.websockets.connect", new=AsyncMock(return_value=mock_ws) ) as mock_connect: transport = WebSocketTransport("ws://example.com", heartbeat_interval=0.01) await transport.start() mock_connect.assert_awaited_once_with("ws://example.com", additional_headers={}) assert transport.state == TransportState.CONNECTED assert transport._ws is mock_ws assert len(transport._tasks) == 2 # receive + heartbeat await transport.stop() async def test_start_is_idempotent(self): mock_ws = MagicMock() mock_ws.close = AsyncMock() mock_ws.recv = AsyncMock(return_value=b"raw") with patch( "yuxi.channel.transport.websocket.websockets.connect", new=AsyncMock(return_value=mock_ws) ) as mock_connect: transport = WebSocketTransport("ws://example.com") await transport.start() await transport.start() mock_connect.assert_awaited_once() await transport.stop() async def test_start_failure_marks_disconnected(self): with patch( "yuxi.channel.transport.websocket.websockets.connect", new=AsyncMock(side_effect=ConnectionError("failed")), ): transport = WebSocketTransport("ws://example.com") with pytest.raises(ConnectionError): await transport.start() assert transport.state == TransportState.DISCONNECTED assert transport._closed is True assert transport._ws is None async def test_stop_closes_websocket_and_cancels_tasks(self): mock_ws = MagicMock() mock_ws.close = AsyncMock() mock_ws.send = AsyncMock() mock_ws.recv = AsyncMock(side_effect=asyncio.CancelledError) with patch("yuxi.channel.transport.websocket.websockets.connect", new=AsyncMock(return_value=mock_ws)): transport = WebSocketTransport("ws://example.com", heartbeat_interval=0.01) await transport.start() await asyncio.sleep(0) await transport.stop() assert transport.state == TransportState.STOPPED assert transport._ws is None assert transport._tasks == set() mock_ws.close.assert_awaited_once() async def test_stop_is_safe_when_not_started(self): transport = WebSocketTransport("ws://example.com") await transport.stop() assert transport.state == TransportState.STOPPED def test_on_message_sets_handler(self): transport = WebSocketTransport("ws://example.com") handler = AsyncMock() transport.on_message(handler) assert transport._message_handler is handler async def test_send_success(self): mock_ws = MagicMock() mock_ws.send = AsyncMock() transport = WebSocketTransport("ws://example.com") transport._ws = mock_ws transport._state = TransportState.CONNECTED await transport.send(b"hello") mock_ws.send.assert_awaited_once_with(b"hello") async def test_send_raises_when_disconnected(self): transport = WebSocketTransport("ws://example.com") transport._state = TransportState.DISCONNECTED with pytest.raises(RuntimeError, match="WebSocket is not connected"): await transport.send(b"hello") async def test_send_raises_when_ws_is_none(self): transport = WebSocketTransport("ws://example.com") transport._state = TransportState.CONNECTED transport._ws = None with pytest.raises(RuntimeError, match="WebSocket is not connected"): await transport.send(b"hello") async def test_receive_loop_delivers_bytes(self): mock_ws = MagicMock() mock_ws.recv = AsyncMock(side_effect=[b"raw", asyncio.CancelledError]) mock_ws.close = AsyncMock() handler = AsyncMock() transport = WebSocketTransport("ws://example.com") transport._ws = mock_ws transport._closed = False transport.on_message(handler) receive_task = asyncio.create_task(transport._receive_loop()) await asyncio.sleep(0.02) transport._closed = True receive_task.cancel() try: await receive_task except asyncio.CancelledError: pass handler.assert_awaited_once_with(b"raw") async def test_receive_loop_encodes_str(self): mock_ws = MagicMock() mock_ws.recv = AsyncMock(side_effect=["text", asyncio.CancelledError]) mock_ws.close = AsyncMock() handler = AsyncMock() transport = WebSocketTransport("ws://example.com") transport._ws = mock_ws transport._closed = False transport.on_message(handler) receive_task = asyncio.create_task(transport._receive_loop()) await asyncio.sleep(0.02) transport._closed = True receive_task.cancel() try: await receive_task except asyncio.CancelledError: pass handler.assert_awaited_once_with(b"text") async def test_receive_loop_connection_closed_marks_disconnected(self): mock_ws = MagicMock() mock_ws.recv = AsyncMock(side_effect=ConnectionClosed(None, None)) mock_ws.close = AsyncMock() transport = WebSocketTransport("ws://example.com") transport._ws = mock_ws transport._closed = False await transport._receive_loop() assert transport.state == TransportState.DISCONNECTED async def test_receive_loop_handler_error_is_logged(self): mock_ws = MagicMock() mock_ws.recv = AsyncMock(side_effect=[b"raw", asyncio.CancelledError]) mock_ws.close = AsyncMock() handler = AsyncMock(side_effect=RuntimeError("boom")) transport = WebSocketTransport("ws://example.com") transport._ws = mock_ws transport._closed = False transport.on_message(handler) with patch("yuxi.channel.transport.websocket.logger.exception") as mock_log: receive_task = asyncio.create_task(transport._receive_loop()) await asyncio.sleep(0.02) transport._closed = True receive_task.cancel() try: await receive_task except asyncio.CancelledError: pass handler.assert_awaited_once() mock_log.assert_called_once() async def test_heartbeat_loop_sends_message(self): mock_ws = MagicMock() mock_ws.send = AsyncMock() mock_ws.close = AsyncMock() transport = WebSocketTransport("ws://example.com", heartbeat_interval=0.01, heartbeat_message="ping") transport._ws = mock_ws transport._closed = False heartbeat_task = asyncio.create_task(transport._heartbeat_loop()) await asyncio.sleep(0.03) transport._closed = True heartbeat_task.cancel() try: await heartbeat_task except asyncio.CancelledError: pass assert mock_ws.send.await_count >= 1 mock_ws.send.assert_any_await("ping") async def test_heartbeat_loop_connection_closed_marks_disconnected(self): mock_ws = MagicMock() mock_ws.send = AsyncMock(side_effect=ConnectionClosed(None, None)) mock_ws.close = AsyncMock() transport = WebSocketTransport("ws://example.com", heartbeat_interval=0.01) transport._ws = mock_ws transport._closed = False await transport._heartbeat_loop() assert transport.state == TransportState.DISCONNECTED async def test_heartbeat_loop_generic_error_marks_disconnected(self): mock_ws = MagicMock() mock_ws.send = AsyncMock(side_effect=RuntimeError("boom")) mock_ws.close = AsyncMock() transport = WebSocketTransport("ws://example.com", heartbeat_interval=0.01) transport._ws = mock_ws transport._closed = False await transport._heartbeat_loop() assert transport.state == TransportState.DISCONNECTED async def test_heartbeat_loop_skips_when_interval_zero(self): mock_ws = MagicMock() mock_ws.send = AsyncMock() mock_ws.close = AsyncMock() with patch("yuxi.channel.transport.websocket.websockets.connect", new=AsyncMock(return_value=mock_ws)): transport = WebSocketTransport("ws://example.com", heartbeat_interval=0) await transport.start() assert len(transport._tasks) == 1 # only receive loop await transport.stop()