from __future__ import annotations import asyncio from unittest.mock import AsyncMock, patch import pytest from yuxi.channel.transport.polling import PollingTransport from yuxi.channel.transport.protocol import TransportState class TestPollingTransport: def test_initial_state(self): poll_func = AsyncMock() transport = PollingTransport(poll_func, interval=2.0) assert transport.poll_func is poll_func assert transport.interval == 2.0 assert transport.state == TransportState.DISCONNECTED assert transport._closed is True assert transport._message_handler is None assert transport._task is None async def test_start_creates_poll_task(self): poll_func = AsyncMock(return_value=[]) transport = PollingTransport(poll_func) await transport.start() assert transport.state == TransportState.CONNECTED assert transport._closed is False assert transport._task is not None await transport.stop() async def test_start_is_idempotent(self): poll_func = AsyncMock(return_value=[]) transport = PollingTransport(poll_func) await transport.start() first_task = transport._task await transport.start() assert transport._task is first_task await transport.stop() async def test_stop_cancels_task(self): poll_func = AsyncMock(return_value=[]) transport = PollingTransport(poll_func) await transport.start() assert transport._task is not None await transport.stop() assert transport.state == TransportState.STOPPED assert transport._task is None assert transport._closed is True async def test_stop_is_safe_when_not_started(self): transport = PollingTransport(AsyncMock()) await transport.stop() assert transport.state == TransportState.STOPPED def test_on_message_sets_handler(self): transport = PollingTransport(AsyncMock()) handler = AsyncMock() transport.on_message(handler) assert transport._message_handler is handler async def test_send_raises_not_implemented(self): transport = PollingTransport(AsyncMock()) with pytest.raises(NotImplementedError, match="Polling transport does not support direct send"): await transport.send(b"data") async def test_poll_loop_delivers_messages(self): handler = AsyncMock() poll_func = AsyncMock(return_value=[b"msg1", b"msg2"]) transport = PollingTransport(poll_func, interval=0.01) transport.on_message(handler) await transport.start() await asyncio.sleep(0.05) await transport.stop() handler.assert_any_await(b"msg1") handler.assert_any_await(b"msg2") assert handler.await_count == 2 async def test_poll_loop_handler_error_is_logged(self): handler = AsyncMock(side_effect=RuntimeError("boom")) poll_func = AsyncMock(return_value=[b"msg1"]) transport = PollingTransport(poll_func, interval=0.01) transport.on_message(handler) with patch("yuxi.channel.transport.polling.logger.exception") as mock_log: await transport.start() await asyncio.sleep(0.05) await transport.stop() handler.assert_awaited_once_with(b"msg1") mock_log.assert_called_once() async def test_poll_loop_poll_error_is_logged(self): poll_func = AsyncMock(side_effect=RuntimeError("poll failed")) transport = PollingTransport(poll_func, interval=0.01) with patch("yuxi.channel.transport.polling.logger.exception") as mock_log: await transport.start() await asyncio.sleep(0.05) await transport.stop() assert poll_func.await_count >= 1 mock_log.assert_called_once() async def test_poll_loop_three_failures_marks_disconnected(self): poll_func = AsyncMock(side_effect=RuntimeError("poll failed")) transport = PollingTransport(poll_func, interval=0.01) await transport.start() # Wait for at least 3 poll attempts await asyncio.sleep(0.1) assert transport.state == TransportState.DISCONNECTED await transport.stop() async def test_poll_loop_respects_closed_flag(self): poll_func = AsyncMock(return_value=[]) transport = PollingTransport(poll_func, interval=0.01) await transport.start() await transport.stop() # Give the loop a moment to process cancellation await asyncio.sleep(0.02) assert transport._task is None assert transport.state == TransportState.STOPPED