ForcePilot/backend/test/unit/channel/transport/test_websocket.py
Kris bab30f2715
Some checks failed
Deploy VitePress site to Pages / build (push) Has been cancelled
Ruff Format Check / Ruff Format & Lint (push) Has been cancelled
Deploy VitePress site to Pages / Deploy (push) Has been cancelled
feat:0715
2026-07-15 12:30:58 +08:00

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()