diff --git a/backend/src/api/packet_api.py b/backend/src/api/packet_api.py index cad84a4..457f193 100644 --- a/backend/src/api/packet_api.py +++ b/backend/src/api/packet_api.py @@ -9,6 +9,7 @@ from typing import Any, Dict, List, Optional, Union from fastapi import APIRouter, HTTPException, Query, WebSocket, WebSocketDisconnect from fastapi.responses import JSONResponse from pydantic import BaseModel +from starlette.websockets import WebSocketState import src.shared_objects as shared from src.Models.packets import PacketDBModel @@ -120,7 +121,41 @@ async def websocket_packets(ws: WebSocket) -> None: logger.info("WebSocket subscribed client %s (queue maxsize=%d)", ws.client, queue.maxsize) while True: - message = await queue.get() + queue_task = asyncio.create_task(queue.get()) + receive_task = asyncio.create_task(ws.receive()) + done, pending = await asyncio.wait( + {queue_task, receive_task}, + return_when=asyncio.FIRST_COMPLETED, + ) + + for task in pending: + task.cancel() + if pending: + await asyncio.gather(*pending, return_exceptions=True) + + if receive_task in done: + try: + inbound = receive_task.result() + except WebSocketDisconnect: + logger.info("WebSocket client disconnected: %s", ws.client) + break + except Exception: + logger.info("WebSocket receive failed for client %s; unsubscribing", ws.client) + break + + if inbound.get("type") == "websocket.disconnect": + logger.info("WebSocket disconnect received for client %s", ws.client) + break + + if queue_task not in done: + if ws.client_state is not WebSocketState.CONNECTED: + break + continue + + message = queue_task.result() + if isinstance(message, dict) and message.get("type") == "__broadcaster_shutdown__": + logger.info("Broadcaster shutdown delivered to client %s", ws.client) + break if isinstance(message, dict): payload: Any = _serialize_row_for_json(message) diff --git a/backend/src/utilities/packet_broadcaster.py b/backend/src/utilities/packet_broadcaster.py index d4d7db0..eb8771e 100644 --- a/backend/src/utilities/packet_broadcaster.py +++ b/backend/src/utilities/packet_broadcaster.py @@ -6,6 +6,8 @@ from typing import Any, Dict, List, Optional logger = logging.getLogger("packet_broadcaster") +_SHUTDOWN_SENTINEL: Dict[str, Any] = {"type": "__broadcaster_shutdown__"} + class PacketBroadcaster: """Manage subscriber queues and publish packet events.""" @@ -96,5 +98,6 @@ class PacketBroadcaster: try: while not queue.empty(): queue.get_nowait() + queue.put_nowait(_SHUTDOWN_SENTINEL) except Exception: pass