277 lines
9.9 KiB
Python
277 lines
9.9 KiB
Python
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()
|