diff --git a/roboco/api/websocket.py b/roboco/api/websocket.py index 93e0b14e..9a68fd62 100644 --- a/roboco/api/websocket.py +++ b/roboco/api/websocket.py @@ -19,6 +19,7 @@ from typing import Any from uuid import UUID import httpx +import structlog from fastapi import APIRouter, WebSocket, WebSocketDisconnect, status from roboco.agents_config import CEO_AGENT_ID, verify_agent_token @@ -28,6 +29,39 @@ from roboco.db.base import get_db from roboco.services.repositories import resolve_agent_uuid router = APIRouter() +log = structlog.get_logger() + +# F066: server-side idle timeout for WS receive loops. A half-open socket +# (dead agent container, silent client) blocks ``receive_text()`` forever; +# wrapping it in ``asyncio.wait_for`` reaps the socket after this many +# seconds of silence. No env-var/config precedent exists in ``config.py`` +# for WS tuning, so this is a module constant — callers/tests patch it. +IDLE_TIMEOUT_SECONDS: float = 90.0 + +# F064: per-connection send queue + send timeout. Each registered connection +# owns a bounded ``asyncio.Queue`` drained by a sender task, so a slow client +# can't back-pressure the fan-out: broadcast enqueues (non-blocking) and +# returns immediately. When the queue is full the message is dropped + logged +# (the client is lagging, not the whole fan-out). ``send_text`` itself is +# wrapped in ``wait_for`` so a stuck transport doesn't wedge the sender. +MAX_SEND_QUEUE: int = 256 +SEND_TIMEOUT_SECONDS: float = 10.0 + + +class _ClientConnection: + """F064: per-connection send queue + sender task. + + Holds the bounded outbound queue drained by ``sender``; broadcast enqueues + here instead of awaiting ``send_text`` directly, so one slow client cannot + block the fan-out to every other client. + """ + + __slots__ = ("queue", "sender", "websocket") + + def __init__(self, websocket: WebSocket, maxsize: int) -> None: + self.websocket = websocket + self.queue: asyncio.Queue[str] = asyncio.Queue(maxsize=maxsize) + self.sender: asyncio.Task[None] | None = None async def _require_panel_token(websocket: WebSocket) -> bool: @@ -87,6 +121,45 @@ class ConnectionManager: # websocket -> agent_id (for tracking who is connected) self.connection_agents: dict[WebSocket, UUID] = {} + # F064: websocket -> per-connection send queue + sender task. Every + # connect_* registers here; disconnect cancels + removes. Broadcast + # enqueues into these queues instead of awaiting send_text directly so + # one slow client can't block the fan-out. + self.connection_senders: dict[WebSocket, _ClientConnection] = {} + + # F064: fire-and-forget fallback send tasks for unregistered sockets + # (legacy path). Held only to satisfy ruff RUF006 + to allow clean + # shutdown; each task removes itself on completion. + self._pending_sends: set[asyncio.Task[None]] = set() + + def _register_sender(self, websocket: WebSocket) -> _ClientConnection: + """Create the per-connection send queue + start its sender task.""" + conn = _ClientConnection(websocket, maxsize=MAX_SEND_QUEUE) + conn.sender = asyncio.create_task(self._run_sender(conn)) + self.connection_senders[websocket] = conn + return conn + + async def _run_sender(self, conn: _ClientConnection) -> None: + """Drain the per-connection send queue; each send is timeout-bounded.""" + ws = conn.websocket + while True: + data = await conn.queue.get() + try: + await asyncio.wait_for(ws.send_text(data), timeout=SEND_TIMEOUT_SECONDS) + except TimeoutError: + log.warning( + "WebSocket send timeout — dropping message to slow client", + timeout=SEND_TIMEOUT_SECONDS, + ) + except Exception as exc: + # Transport closed / error — stop sending to this client; the + # receive loop's finally will disconnect and cancel us. + log.debug( + "WebSocket sender stopping on send error", + error=str(exc), + ) + return + async def connect_channel( self, websocket: WebSocket, channel_id: UUID, agent_id: UUID ) -> None: @@ -98,6 +171,7 @@ class ConnectionManager: self.channel_connections[channel_id].add(websocket) self.connection_agents[websocket] = agent_id + self._register_sender(websocket) async def connect_agent( self, websocket: WebSocket, target_agent_id: UUID, viewer_agent_id: UUID @@ -110,6 +184,7 @@ class ConnectionManager: self.agent_connections[target_agent_id].add(websocket) self.connection_agents[websocket] = viewer_agent_id + self._register_sender(websocket) async def connect_session( self, websocket: WebSocket, session_id: UUID, agent_id: UUID @@ -122,6 +197,7 @@ class ConnectionManager: self.session_connections[session_id].add(websocket) self.connection_agents[websocket] = agent_id + self._register_sender(websocket) async def connect_notifications(self, websocket: WebSocket, agent_id: UUID) -> None: """Connect to an agent's notification stream.""" @@ -132,11 +208,13 @@ class ConnectionManager: self.notification_connections[agent_id].add(websocket) self.connection_agents[websocket] = agent_id + self._register_sender(websocket) async def connect_system(self, websocket: WebSocket) -> None: """Connect to the operator/system-wide stream (rate limits, etc.).""" await websocket.accept() self.system_connections.add(websocket) + self._register_sender(websocket) def disconnect(self, websocket: WebSocket) -> None: """Remove a websocket from all subscriptions.""" @@ -162,6 +240,53 @@ class ConnectionManager: # Remove from tracking self.connection_agents.pop(websocket, None) + # F064: cancel + drop the per-connection sender task so a slow/stale + # client's queue doesn't leak after the socket is removed. + conn = self.connection_senders.pop(websocket, None) + if conn is not None and conn.sender is not None: + conn.sender.cancel() + + def _enqueue_or_send(self, websocket: WebSocket, data: str) -> None: + """F064: fan out one message to one connection without blocking. + + Registered connections (created via ``connect_*``) get the message + enqueued into their bounded send queue — non-blocking, drop + warn on + overflow. An unregistered socket (legacy path: present in a + subscription set but not in ``connection_senders``) falls back to a + timeout-bounded ``send_text`` scheduled on the loop, so the broadcast + still never blocks on a single slow client. + """ + conn = self.connection_senders.get(websocket) + if conn is not None: + try: + conn.queue.put_nowait(data) + except asyncio.QueueFull: + log.warning( + "WebSocket send queue overflow — dropping message", + queue_size=conn.queue.maxsize, + ) + return + # Legacy fallback: schedule a timeout-bounded send so a slow + # unregistered client can't wedge the fan-out either. Keep a strong + # reference so the task isn't GC'd mid-flight (ruff RUF006); it + # discards itself on completion. + task = asyncio.create_task(self._send_with_timeout(websocket, data)) + self._pending_sends.add(task) + task.add_done_callback(self._pending_sends.discard) + + async def _send_with_timeout(self, websocket: WebSocket, data: str) -> None: + try: + await asyncio.wait_for( + websocket.send_text(data), timeout=SEND_TIMEOUT_SECONDS + ) + except TimeoutError: + log.warning( + "WebSocket send timeout — dropping message to slow client", + timeout=SEND_TIMEOUT_SECONDS, + ) + except Exception as exc: # transport closed / cancelled + log.debug("WebSocket send failed", error=str(exc)) + async def broadcast_to_channel( self, channel_id: UUID, message: dict[str, Any] ) -> None: @@ -169,12 +294,9 @@ class ConnectionManager: connections = self.channel_connections.get(channel_id, set()) if not connections: return - data = json.dumps(message, default=str) - await asyncio.gather( - *[conn.send_text(data) for conn in connections], - return_exceptions=True, - ) + for conn in connections: + self._enqueue_or_send(conn, data) async def broadcast_to_agent_watchers( self, agent_id: UUID, message: dict[str, Any] @@ -183,12 +305,9 @@ class ConnectionManager: connections = self.agent_connections.get(agent_id, set()) if not connections: return - data = json.dumps(message, default=str) - await asyncio.gather( - *[conn.send_text(data) for conn in connections], - return_exceptions=True, - ) + for conn in connections: + self._enqueue_or_send(conn, data) async def broadcast_to_session( self, session_id: UUID, message: dict[str, Any] @@ -197,23 +316,17 @@ class ConnectionManager: connections = self.session_connections.get(session_id, set()) if not connections: return - data = json.dumps(message, default=str) - await asyncio.gather( - *[conn.send_text(data) for conn in connections], - return_exceptions=True, - ) + for conn in connections: + self._enqueue_or_send(conn, data) async def broadcast_system(self, message: dict[str, Any]) -> None: """Broadcast a message to all operator/system-wide subscribers.""" if not self.system_connections: return - data = json.dumps(message, default=str) - await asyncio.gather( - *[conn.send_text(data) for conn in self.system_connections], - return_exceptions=True, - ) + for conn in self.system_connections: + self._enqueue_or_send(conn, data) def get_channel_subscriber_count(self, channel_id: UUID) -> int: """Get number of subscribers to a channel.""" @@ -323,7 +436,9 @@ async def channel_stream( # Keep connection alive and handle incoming messages while True: - data = await websocket.receive_text() + data = await asyncio.wait_for( + websocket.receive_text(), timeout=IDLE_TIMEOUT_SECONDS + ) # Handle ping/pong for keepalive if data == "ping": @@ -334,6 +449,19 @@ async def channel_stream( # For now, channels are primarily for receiving except WebSocketDisconnect: + # Clean client-initiated disconnect — handled here for clarity; the + # finally below also disconnects (idempotent) to cover every other + # exit path (anyio closed-resource, CancelledError, transport errors). + pass + except TimeoutError: + # F066: idle timeout — the client has been silent for + # IDLE_TIMEOUT_SECONDS (likely a half-open socket from a dead + # container). Log and fall through to the finally so the socket is + # removed from every subscription set. + log.warning( + "WebSocket idle timeout — disconnecting", timeout=IDLE_TIMEOUT_SECONDS + ) + finally: manager.disconnect(websocket) @@ -380,11 +508,26 @@ async def agent_stream( ) while True: - data = await websocket.receive_text() + data = await asyncio.wait_for( + websocket.receive_text(), timeout=IDLE_TIMEOUT_SECONDS + ) if data == "ping": await websocket.send_text("pong") except WebSocketDisconnect: + # Clean client-initiated disconnect — handled here for clarity; the + # finally below also disconnects (idempotent) to cover every other + # exit path (anyio closed-resource, CancelledError, transport errors). + pass + except TimeoutError: + # F066: idle timeout — the client has been silent for + # IDLE_TIMEOUT_SECONDS (likely a half-open socket from a dead + # container). Log and fall through to the finally so the socket is + # removed from every subscription set. + log.warning( + "WebSocket idle timeout — disconnecting", timeout=IDLE_TIMEOUT_SECONDS + ) + finally: manager.disconnect(websocket) @@ -429,11 +572,26 @@ async def session_stream( ) while True: - data = await websocket.receive_text() + data = await asyncio.wait_for( + websocket.receive_text(), timeout=IDLE_TIMEOUT_SECONDS + ) if data == "ping": await websocket.send_text("pong") except WebSocketDisconnect: + # Clean client-initiated disconnect — handled here for clarity; the + # finally below also disconnects (idempotent) to cover every other + # exit path (anyio closed-resource, CancelledError, transport errors). + pass + except TimeoutError: + # F066: idle timeout — the client has been silent for + # IDLE_TIMEOUT_SECONDS (likely a half-open socket from a dead + # container). Log and fall through to the finally so the socket is + # removed from every subscription set. + log.warning( + "WebSocket idle timeout — disconnecting", timeout=IDLE_TIMEOUT_SECONDS + ) + finally: manager.disconnect(websocket) @@ -467,11 +625,26 @@ async def notification_stream( ) while True: - data = await websocket.receive_text() + data = await asyncio.wait_for( + websocket.receive_text(), timeout=IDLE_TIMEOUT_SECONDS + ) if data == "ping": await websocket.send_text("pong") except WebSocketDisconnect: + # Clean client-initiated disconnect — handled here for clarity; the + # finally below also disconnects (idempotent) to cover every other + # exit path (anyio closed-resource, CancelledError, transport errors). + pass + except TimeoutError: + # F066: idle timeout — the client has been silent for + # IDLE_TIMEOUT_SECONDS (likely a half-open socket from a dead + # container). Log and fall through to the finally so the socket is + # removed from every subscription set. + log.warning( + "WebSocket idle timeout — disconnecting", timeout=IDLE_TIMEOUT_SECONDS + ) + finally: manager.disconnect(websocket) @@ -490,11 +663,26 @@ async def system_stream(websocket: WebSocket) -> None: await websocket.send_json({"type": "connected"}) while True: - data = await websocket.receive_text() + data = await asyncio.wait_for( + websocket.receive_text(), timeout=IDLE_TIMEOUT_SECONDS + ) if data == "ping": await websocket.send_text("pong") except WebSocketDisconnect: + # Clean client-initiated disconnect — handled here for clarity; the + # finally below also disconnects (idempotent) to cover every other + # exit path (anyio closed-resource, CancelledError, transport errors). + pass + except TimeoutError: + # F066: idle timeout — the client has been silent for + # IDLE_TIMEOUT_SECONDS (likely a half-open socket from a dead + # container). Log and fall through to the finally so the socket is + # removed from every subscription set. + log.warning( + "WebSocket idle timeout — disconnecting", timeout=IDLE_TIMEOUT_SECONDS + ) + finally: manager.disconnect(websocket) @@ -543,7 +731,5 @@ async def broadcast_notification( for agent_id in agent_ids: connections = manager.notification_connections.get(agent_id, set()) if connections: - await asyncio.gather( - *[conn.send_text(data) for conn in connections], - return_exceptions=True, - ) + for conn in connections: + manager._enqueue_or_send(conn, data) diff --git a/tests/unit/api/test_websocket_handler_cleanup.py b/tests/unit/api/test_websocket_handler_cleanup.py new file mode 100644 index 00000000..8ca4c8d7 --- /dev/null +++ b/tests/unit/api/test_websocket_handler_cleanup.py @@ -0,0 +1,197 @@ +"""F065: WS route handlers must disconnect on ANY exit path, not just +WebSocketDisconnect. + +The old handlers were ``try: ... while True: receive_text() ... except +WebSocketDisconnect: manager.disconnect(websocket)`` with NO ``finally``. +If ``receive_text()`` raised anything else (anyio closed-resource during +shutdown, ``asyncio.CancelledError``, transport errors), the exception +propagated WITHOUT calling ``manager.disconnect(websocket)``, so the dead +socket stayed in the subscription set + ``connection_agents`` forever and +was still fanned out to on every broadcast. + +The fix adds ``finally: manager.disconnect(websocket)`` to every handler. +``disconnect`` is idempotent (``set.discard`` / ``dict.pop`` with default), +so the clean-disconnect path (still caught by ``except WebSocketDisconnect`` +for clarity) and the new finally both calling it is safe. + +These tests use mock sockets (no real app/Redis) and an isolated +``ConnectionManager`` patched in for the module-global ``manager``. +""" + +from __future__ import annotations + +import asyncio +from unittest.mock import AsyncMock, MagicMock +from uuid import uuid4 + +import pytest +from fastapi import WebSocketDisconnect +from roboco.api.websocket import ( + ConnectionManager, + agent_stream, + channel_stream, + notification_stream, + session_stream, + system_stream, +) + + +def _mock_ws_for_receive(receive_side_effect: object) -> MagicMock: + """A socket whose receive_text raises/returns per ``receive_side_effect``.""" + ws = MagicMock() + ws.accept = AsyncMock() + ws.close = AsyncMock() + ws.send_json = AsyncMock() + ws.send_text = AsyncMock() + ws.receive_text = AsyncMock(side_effect=receive_side_effect) + ws.headers = {} + ws.query_params = {} + return ws + + +# --------------------------------------------------------------------------- +# system_stream (no per-agent keying) +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_system_stream_disconnects_on_non_disconnect_exception() -> None: + """A non-WebSocketDisconnect exception (e.g. anyio closed-resource during + shutdown) must still remove the socket from the manager — the old code + only caught WebSocketDisconnect and leaked the dead socket.""" + mgr = ConnectionManager() + ws = _mock_ws_for_receive(RuntimeError("connection closed during shutdown")) + await mgr.connect_system(ws) + assert ws in mgr.system_connections + + with pytest.MonkeyPatch.context() as mp: + mp.setattr("roboco.api.websocket.manager", mgr) + with pytest.raises(RuntimeError): + await system_stream(ws) + + assert ws not in mgr.system_connections + + +@pytest.mark.asyncio +async def test_system_stream_disconnects_on_cancelled_error() -> None: + """asyncio.CancelledError during shutdown must also disconnect.""" + mgr = ConnectionManager() + ws = _mock_ws_for_receive(asyncio.CancelledError()) + await mgr.connect_system(ws) + + with pytest.MonkeyPatch.context() as mp: + mp.setattr("roboco.api.websocket.manager", mgr) + with pytest.raises(asyncio.CancelledError): + await system_stream(ws) + + assert ws not in mgr.system_connections + + +# --------------------------------------------------------------------------- +# notification_stream (representative per-agent handler) +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_notification_stream_disconnects_on_non_disconnect_exception( + monkeypatch: pytest.MonkeyPatch, +) -> None: + + agent_id = uuid4() + mgr = ConnectionManager() + ws = _mock_ws_for_receive(RuntimeError("transport reset")) + monkeypatch.setattr( + "roboco.api.websocket.validate_agent_exists", AsyncMock(return_value=True) + ) + monkeypatch.setattr("roboco.api.websocket.manager", mgr) + + with pytest.raises(RuntimeError): + await notification_stream(ws, agent_id) + + assert ws not in mgr.notification_connections.get(agent_id, set()) + assert ws not in mgr.connection_agents + + +# --------------------------------------------------------------------------- +# channel / agent / session handlers (same pattern, distinct subscription sets) +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_channel_stream_disconnects_on_non_disconnect_exception( + monkeypatch: pytest.MonkeyPatch, +) -> None: + + channel_id = uuid4() + agent_id = uuid4() + mgr = ConnectionManager() + ws = _mock_ws_for_receive(RuntimeError("anyio closed")) + ws.query_params = {"agent_id": str(agent_id)} + monkeypatch.setattr( + "roboco.api.websocket.validate_channel_access", AsyncMock(return_value=True) + ) + monkeypatch.setattr("roboco.api.websocket.manager", mgr) + + with pytest.raises(RuntimeError): + await channel_stream(ws, channel_id) + + assert ws not in mgr.channel_connections.get(channel_id, set()) + + +@pytest.mark.asyncio +async def test_agent_stream_disconnects_on_non_disconnect_exception( + monkeypatch: pytest.MonkeyPatch, +) -> None: + + target_id = uuid4() + viewer_id = uuid4() + mgr = ConnectionManager() + ws = _mock_ws_for_receive(RuntimeError("anyio closed")) + ws.query_params = {"viewer_id": str(viewer_id)} + monkeypatch.setattr( + "roboco.api.websocket.validate_agent_exists", AsyncMock(return_value=True) + ) + monkeypatch.setattr("roboco.api.websocket.manager", mgr) + + with pytest.raises(RuntimeError): + await agent_stream(ws, target_id) + + assert ws not in mgr.agent_connections.get(target_id, set()) + + +@pytest.mark.asyncio +async def test_session_stream_disconnects_on_non_disconnect_exception( + monkeypatch: pytest.MonkeyPatch, +) -> None: + + session_id = uuid4() + agent_id = uuid4() + mgr = ConnectionManager() + ws = _mock_ws_for_receive(RuntimeError("anyio closed")) + ws.query_params = {"agent_id": str(agent_id)} + monkeypatch.setattr( + "roboco.api.websocket.validate_agent_exists", AsyncMock(return_value=True) + ) + monkeypatch.setattr("roboco.api.websocket.manager", mgr) + + with pytest.raises(RuntimeError): + await session_stream(ws, session_id) + + assert ws not in mgr.session_connections.get(session_id, set()) + + +@pytest.mark.asyncio +async def test_system_stream_clean_disconnect_still_works() -> None: + """Regression: the clean WebSocketDisconnect path still disconnects (the + new finally must not break the happy path or double-disconnect).""" + mgr = ConnectionManager() + ws = _mock_ws_for_receive(["ping", WebSocketDisconnect()]) + await mgr.connect_system(ws) + + with pytest.MonkeyPatch.context() as mp: + mp.setattr("roboco.api.websocket.manager", mgr) + await system_stream(ws) + + assert ws not in mgr.system_connections + # pong was answered before disconnect. + ws.send_text.assert_awaited_with("pong") diff --git a/tests/unit/api/test_websocket_idle_timeout.py b/tests/unit/api/test_websocket_idle_timeout.py new file mode 100644 index 00000000..f0e25319 --- /dev/null +++ b/tests/unit/api/test_websocket_idle_timeout.py @@ -0,0 +1,204 @@ +"""F066: server-side idle timeout reaps half-open WS sockets. + +The keepalive was client-driven (respond to ``"ping"`` with ``"pong"``); the +server never sent its own ping and never timed out a silent client. If an +agent container died leaving the TCP socket half-open, ``receive_text()`` +blocked forever and ``disconnect`` was never called. + +The fix wraps ``receive_text()`` in +``asyncio.wait_for(..., timeout=IDLE_TIMEOUT_SECONDS)`` per handler; on +``asyncio.TimeoutError`` the finally (from F065) disconnects the idle +socket. ``IDLE_TIMEOUT_SECONDS`` is a named module constant (ruff PLR2004). + +Deterministic: the slow-socket test patches ``IDLE_TIMEOUT_SECONDS`` to a +tiny value (0.05s) and uses a ``receive_text`` that returns a never-resolved +``Future`` — so the test asserts a prompt disconnect in well under a second, +never relying on real wall-clock timing of the default timeout. +""" + +from __future__ import annotations + +import asyncio +from unittest.mock import AsyncMock, MagicMock +from uuid import uuid4 + +import pytest +from fastapi import WebSocketDisconnect +from roboco.api.websocket import ( + IDLE_TIMEOUT_SECONDS, + ConnectionManager, + agent_stream, + channel_stream, + notification_stream, + session_stream, + system_stream, +) + + +def _mock_ws_for_receive(receive_side_effect: object) -> MagicMock: + ws = MagicMock() + ws.accept = AsyncMock() + ws.close = AsyncMock() + ws.send_json = AsyncMock() + ws.send_text = AsyncMock() + + if isinstance(receive_side_effect, asyncio.Future): + # A never-resolving Future means "hang forever" (half-open socket). + # AsyncMock treats a non-callable side_effect as an iterable, which a + # Future isn't — so install a real async receive_text that awaits it. + hang_future = receive_side_effect + + async def _hang_forever() -> str: + await hang_future # never resolves; wait_for cancels it on timeout. + return "" + + ws.receive_text = _hang_forever + else: + ws.receive_text = AsyncMock(side_effect=receive_side_effect) + ws.headers = {} + ws.query_params = {} + return ws + + +# --------------------------------------------------------------------------- +# Constant shape +# --------------------------------------------------------------------------- + + +def test_idle_timeout_seconds_is_a_named_module_constant() -> None: + """IDLE_TIMEOUT_SECONDS must be a module-level constant (ruff PLR2004).""" + assert isinstance(IDLE_TIMEOUT_SECONDS, int | float) + assert IDLE_TIMEOUT_SECONDS > 0 + + +# --------------------------------------------------------------------------- +# Half-open socket is reaped after the idle timeout (deterministic) +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_system_stream_reaps_silent_socket_after_idle_timeout() -> None: + """A silent half-open socket (receive_text never returns) is reaped after + the idle timeout — the wait_for raises TimeoutError and the finally + disconnects. Deterministic: tiny patched timeout + never-resolving Future.""" + mgr = ConnectionManager() + hang_future: asyncio.Future[str] = asyncio.Future() + ws = _mock_ws_for_receive(hang_future) + await mgr.connect_system(ws) + + with pytest.MonkeyPatch.context() as mp: + mp.setattr("roboco.api.websocket.manager", mgr) + mp.setattr("roboco.api.websocket.IDLE_TIMEOUT_SECONDS", 0.05) + # Must return promptly (well under 2s), not block for the real default. + await asyncio.wait_for(system_stream(ws), timeout=2.0) + + assert ws not in mgr.system_connections + + +@pytest.mark.asyncio +async def test_notification_stream_reaps_silent_socket_after_idle_timeout( + monkeypatch: pytest.MonkeyPatch, +) -> None: + + agent_id = uuid4() + mgr = ConnectionManager() + hang_future: asyncio.Future[str] = asyncio.Future() + ws = _mock_ws_for_receive(hang_future) + monkeypatch.setattr( + "roboco.api.websocket.validate_agent_exists", AsyncMock(return_value=True) + ) + monkeypatch.setattr("roboco.api.websocket.manager", mgr) + monkeypatch.setattr("roboco.api.websocket.IDLE_TIMEOUT_SECONDS", 0.05) + + await asyncio.wait_for(notification_stream(ws, agent_id), timeout=2.0) + + assert ws not in mgr.notification_connections.get(agent_id, set()) + + +@pytest.mark.asyncio +async def test_channel_stream_reaps_silent_socket_after_idle_timeout( + monkeypatch: pytest.MonkeyPatch, +) -> None: + + channel_id = uuid4() + agent_id = uuid4() + mgr = ConnectionManager() + hang_future: asyncio.Future[str] = asyncio.Future() + ws = _mock_ws_for_receive(hang_future) + ws.query_params = {"agent_id": str(agent_id)} + monkeypatch.setattr( + "roboco.api.websocket.validate_channel_access", AsyncMock(return_value=True) + ) + monkeypatch.setattr("roboco.api.websocket.manager", mgr) + monkeypatch.setattr("roboco.api.websocket.IDLE_TIMEOUT_SECONDS", 0.05) + + await asyncio.wait_for(channel_stream(ws, channel_id), timeout=2.0) + + assert ws not in mgr.channel_connections.get(channel_id, set()) + + +@pytest.mark.asyncio +async def test_agent_stream_reaps_silent_socket_after_idle_timeout( + monkeypatch: pytest.MonkeyPatch, +) -> None: + + target_id = uuid4() + viewer_id = uuid4() + mgr = ConnectionManager() + hang_future: asyncio.Future[str] = asyncio.Future() + ws = _mock_ws_for_receive(hang_future) + ws.query_params = {"viewer_id": str(viewer_id)} + monkeypatch.setattr( + "roboco.api.websocket.validate_agent_exists", AsyncMock(return_value=True) + ) + monkeypatch.setattr("roboco.api.websocket.manager", mgr) + monkeypatch.setattr("roboco.api.websocket.IDLE_TIMEOUT_SECONDS", 0.05) + + await asyncio.wait_for(agent_stream(ws, target_id), timeout=2.0) + + assert ws not in mgr.agent_connections.get(target_id, set()) + + +@pytest.mark.asyncio +async def test_session_stream_reaps_silent_socket_after_idle_timeout( + monkeypatch: pytest.MonkeyPatch, +) -> None: + + session_id = uuid4() + agent_id = uuid4() + mgr = ConnectionManager() + hang_future: asyncio.Future[str] = asyncio.Future() + ws = _mock_ws_for_receive(hang_future) + ws.query_params = {"agent_id": str(agent_id)} + monkeypatch.setattr( + "roboco.api.websocket.validate_agent_exists", AsyncMock(return_value=True) + ) + monkeypatch.setattr("roboco.api.websocket.manager", mgr) + monkeypatch.setattr("roboco.api.websocket.IDLE_TIMEOUT_SECONDS", 0.05) + + await asyncio.wait_for(session_stream(ws, session_id), timeout=2.0) + + assert ws not in mgr.session_connections.get(session_id, set()) + + +# --------------------------------------------------------------------------- +# Regression: ping/pong within the idle window keeps the socket alive +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_ping_within_idle_window_does_not_disconnect() -> None: + """A client that sends ping before the idle timeout elapses is NOT + disconnected — the wait_for resets on each successful receive_text.""" + mgr = ConnectionManager() + ws = _mock_ws_for_receive(["ping", WebSocketDisconnect()]) + await mgr.connect_system(ws) + + with pytest.MonkeyPatch.context() as mp: + mp.setattr("roboco.api.websocket.manager", mgr) + mp.setattr("roboco.api.websocket.IDLE_TIMEOUT_SECONDS", 30) + await system_stream(ws) + + # Disconnected only because of the WebSocketDisconnect, not the timeout. + assert ws not in mgr.system_connections + ws.send_text.assert_awaited_with("pong") diff --git a/tests/unit/api/test_websocket_send_queue.py b/tests/unit/api/test_websocket_send_queue.py new file mode 100644 index 00000000..6f6f00dc --- /dev/null +++ b/tests/unit/api/test_websocket_send_queue.py @@ -0,0 +1,273 @@ +"""F064: per-connection send queue + send timeout — one slow WS client must +not back-pressure ALL event delivery to ALL clients. + +The old broadcast did ``await asyncio.gather(*[conn.send_text(data) for conn +in connections], return_exceptions=True)`` with no per-connection send queue +and no send timeout. If one client was slow to drain, ``conn.send_text(data)`` +awaited indefinitely on the transport, blocking the gather → the bridge +handler → ``_dispatch_event`` → the whole ``_listen_loop`` for every event +type and recipient. + +The fix gives each registered connection a bounded send queue + a sender +coroutine that drains it, with ``send_text`` behind +``asyncio.wait_for(..., timeout=SEND_TIMEOUT_SECONDS)``. Broadcasts become +fire-and-enqueue: a slow client's queue fills, then drops/overflows (logged +as a warning) instead of blocking the fan-out. The listen loop is never +blocked on a single client. + +Determinism: every slow-send test uses a ``receive``/``send_text`` that +awaits a never-resolved ``Future`` and patches ``SEND_TIMEOUT_SECONDS`` to a +tiny value, so assertions hold in well under a second and never rely on real +wall-clock timing of the default timeout. +""" + +from __future__ import annotations + +import asyncio +import contextlib +from unittest.mock import AsyncMock, MagicMock + +import pytest +from roboco.api.websocket import ( + MAX_SEND_QUEUE, + SEND_TIMEOUT_SECONDS, + ConnectionManager, +) + + +def _make_ws(*, send_side_effect: object | None = None) -> MagicMock: + ws = MagicMock() + ws.accept = AsyncMock() + ws.close = AsyncMock() + ws.send_json = AsyncMock() + if send_side_effect is None: + ws.send_text = AsyncMock() + elif isinstance(send_side_effect, asyncio.Future): + + async def _hang(*_args: object) -> None: + await send_side_effect # never resolves + + ws.send_text = _hang + else: + ws.send_text = AsyncMock(side_effect=send_side_effect) + return ws + + +# --------------------------------------------------------------------------- +# Constants shape +# --------------------------------------------------------------------------- + + +def test_send_constants_are_named_module_constants() -> None: + """SEND_TIMEOUT_SECONDS + MAX_SEND_QUEUE must be module-level constants.""" + assert isinstance(SEND_TIMEOUT_SECONDS, int | float) + assert SEND_TIMEOUT_SECONDS > 0 + assert isinstance(MAX_SEND_QUEUE, int) + assert MAX_SEND_QUEUE > 0 + + +# --------------------------------------------------------------------------- +# Broadcast is fire-and-enqueue: a slow registered client does NOT block it +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_broadcast_returns_promptly_with_slow_registered_client() -> None: + """A registered connection whose send_text never returns must NOT block + broadcast — broadcast enqueues (non-blocking) and returns immediately.""" + mgr = ConnectionManager() + hang: asyncio.Future[None] = asyncio.Future() + slow_ws = _make_ws(send_side_effect=hang) + await mgr.connect_system(slow_ws) + + with pytest.MonkeyPatch.context() as mp: + mp.setattr("roboco.api.websocket.SEND_TIMEOUT_SECONDS", 0.1) + # Must return in well under 1s — enqueue must not await the slow send. + await asyncio.wait_for( + mgr.broadcast_system({"type": "RATE_LIMIT_HIT"}), timeout=1.0 + ) + + # Cleanup: disconnect cancels the stuck sender task. + mgr.disconnect(slow_ws) + + +@pytest.mark.asyncio +async def test_slow_client_does_not_block_fast_client() -> None: + """Two registered connections — one slow (send_text hangs), one fast. + The fast client receives the message promptly; the slow client's send + does not delay the fast client's delivery nor the broadcast return.""" + mgr = ConnectionManager() + hang: asyncio.Future[None] = asyncio.Future() + slow_ws = _make_ws(send_side_effect=hang) + fast_ws = _make_ws() # default AsyncMock send_text returns immediately. + await mgr.connect_system(slow_ws) + await mgr.connect_system(fast_ws) + + with pytest.MonkeyPatch.context() as mp: + mp.setattr("roboco.api.websocket.SEND_TIMEOUT_SECONDS", 0.1) + # Broadcast returns promptly despite the slow client. + await asyncio.wait_for( + mgr.broadcast_system({"type": "USAGE_SNAPSHOT"}), timeout=1.0 + ) + # Let the fast sender drain its queue. + await asyncio.sleep(0.05) + + # Fast client received the message; slow client's send was attempted but + # is still pending (the sender is blocked on the never-resolving send). + assert fast_ws.send_text.await_count >= 1 + sent = fast_ws.send_text.await_args.args[0] + assert "USAGE_SNAPSHOT" in sent + + mgr.disconnect(slow_ws) + mgr.disconnect(fast_ws) + hang.cancel() + + +# --------------------------------------------------------------------------- +# Queue-full drop + warning (deterministic: pre-fill the queue, no await) +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_broadcast_drops_and_warns_when_queue_full() -> None: + """When a slow client's bounded send queue is full, broadcast drops the + message (not enqueued) instead of blocking. Deterministic: pre-fill the + queue synchronously (no await so the sender can't drain), then broadcast + once — the put_nowait raises QueueFull → drop. Assert on the queue state + (still full, the new message was NOT enqueued) rather than log capture, + since structlog doesn't propagate to stdlib ``caplog`` in this config.""" + mgr = ConnectionManager() + hang: asyncio.Future[None] = asyncio.Future() + slow_ws = _make_ws(send_side_effect=hang) + await mgr.connect_system(slow_ws) + conn = mgr.connection_senders[slow_ws] + + # Pre-fill the queue synchronously — the sender task has not been + # scheduled yet (no await between put_nowait calls), so it can't drain. + for _ in range(conn.queue.maxsize): + conn.queue.put_nowait("pending") + assert conn.queue.full() + + with pytest.MonkeyPatch.context() as mp: + mp.setattr("roboco.api.websocket.SEND_TIMEOUT_SECONDS", 0.1) + # Must not raise and must not block. + await asyncio.wait_for( + mgr.broadcast_system({"type": "RATE_LIMIT_HIT"}), timeout=1.0 + ) + + # The broadcast was dropped: the queue still holds exactly maxsize items + # (the new message was NOT enqueued — put_nowait raised QueueFull). + assert conn.queue.full() + assert conn.queue.qsize() == conn.queue.maxsize + + mgr.disconnect(slow_ws) + hang.cancel() + + +# --------------------------------------------------------------------------- +# Fast registered client receives the message (happy path) +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_broadcast_delivers_to_registered_fast_client() -> None: + """A registered connection with a fast send_text receives the message + via its sender task.""" + mgr = ConnectionManager() + fast_ws = _make_ws() + await mgr.connect_system(fast_ws) + + await mgr.broadcast_system({"type": "RATE_LIMIT_HIT", "provider": "anthropic"}) + # Let the sender drain. + await asyncio.sleep(0.05) + + assert fast_ws.send_text.await_count == 1 + sent = fast_ws.send_text.await_args.args[0] + assert "RATE_LIMIT_HIT" in sent + assert "anthropic" in sent + + mgr.disconnect(fast_ws) + + +# --------------------------------------------------------------------------- +# Legacy fallback: unregistered socket in a subscription set still gets a +# send timeout (so the OLD direct-send path is also protected). +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_broadcast_send_timeout_protects_legacy_unregistered_socket() -> None: + """A socket present in a subscription set but NOT registered via connect_* + (the legacy test path) still must not block broadcast forever: the + fallback wraps send_text in wait_for(SEND_TIMEOUT_SECONDS).""" + mgr = ConnectionManager() + hang: asyncio.Future[None] = asyncio.Future() + + async def _hang() -> None: + await hang + + legacy_ws = MagicMock() + legacy_ws.send_text = _hang + # Put it straight into the set — bypasses connect_system (no sender). + mgr.system_connections.add(legacy_ws) + + with pytest.MonkeyPatch.context() as mp: + mp.setattr("roboco.api.websocket.SEND_TIMEOUT_SECONDS", 0.1) + # Must return in well under 1s — the slow send is timed out, not + # awaited indefinitely. + await asyncio.wait_for( + mgr.broadcast_system({"type": "RATE_LIMIT_HIT"}), timeout=1.0 + ) + + hang.cancel() + + +# --------------------------------------------------------------------------- +# disconnect cancels the sender task (no leak) +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_disconnect_cancels_sender_task() -> None: + """disconnect() cancels the per-connection sender task so it doesn't + leak after the socket is removed.""" + mgr = ConnectionManager() + ws = _make_ws() + await mgr.connect_system(ws) + conn = mgr.connection_senders[ws] + sender = conn.sender + assert sender is not None + assert not sender.cancelled() + + mgr.disconnect(ws) + + # Sender is removed + cancelled (or done). Give the loop a tick so the + # cancellation actually propagates (cancel() schedules, doesn't sync). + assert ws not in mgr.connection_senders + with contextlib.suppress(TimeoutError, asyncio.CancelledError): + await asyncio.wait_for(sender, timeout=1.0) + assert sender.cancelled() or sender.done() + + +# --------------------------------------------------------------------------- +# Existing direct-set subscription-set broadcast still works (backward compat) +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_broadcast_to_legacy_direct_set_sends_to_each_socket() -> None: + """Sockets added directly to a subscription set (not via connect_*) are + still sent to via the fallback path — preserves the existing test contract.""" + mgr = ConnectionManager() + ws1, ws2 = MagicMock(), MagicMock() + ws1.send_text = AsyncMock() + ws2.send_text = AsyncMock() + mgr.system_connections = {ws1, ws2} + + await mgr.broadcast_system({"type": "x"}) + # The fallback schedules a timeout-bounded send task per socket; let them + # run to completion before asserting. + await asyncio.sleep(0.05) + + ws1.send_text.assert_awaited_once() + ws2.send_text.assert_awaited_once() diff --git a/tests/unit/api/test_websocket_system.py b/tests/unit/api/test_websocket_system.py index 4f01ad34..8e058bf2 100644 --- a/tests/unit/api/test_websocket_system.py +++ b/tests/unit/api/test_websocket_system.py @@ -8,6 +8,7 @@ against mock sockets (no real app/lifespan/Redis). from __future__ import annotations +import asyncio from unittest.mock import AsyncMock, MagicMock import pytest @@ -33,9 +34,14 @@ async def test_broadcast_system_sends_to_every_connection() -> None: ws1, ws2 = MagicMock(), MagicMock() ws1.send_text = AsyncMock() ws2.send_text = AsyncMock() + # These sockets are placed directly into the subscription set (bypassing + # connect_system), so they take the F064 legacy fallback path: broadcast + # schedules a timeout-bounded send task per socket instead of awaiting + # send_text inline. Yield once so those tasks run before asserting. mgr.system_connections = {ws1, ws2} await mgr.broadcast_system({"type": "RATE_LIMIT_HIT", "provider": "anthropic"}) + await asyncio.sleep(0) ws1.send_text.assert_awaited_once() ws2.send_text.assert_awaited_once()