try to fix websocket shutdown
All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 1m39s

This commit is contained in:
2026-03-07 10:35:29 +01:00
parent 299f2e9978
commit 8e6c759cb8
2 changed files with 39 additions and 1 deletions

View File

@@ -9,6 +9,7 @@ from typing import Any, Dict, List, Optional, Union
from fastapi import APIRouter, HTTPException, Query, WebSocket, WebSocketDisconnect from fastapi import APIRouter, HTTPException, Query, WebSocket, WebSocketDisconnect
from fastapi.responses import JSONResponse from fastapi.responses import JSONResponse
from pydantic import BaseModel from pydantic import BaseModel
from starlette.websockets import WebSocketState
import src.shared_objects as shared import src.shared_objects as shared
from src.Models.packets import PacketDBModel 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) logger.info("WebSocket subscribed client %s (queue maxsize=%d)", ws.client, queue.maxsize)
while True: 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): if isinstance(message, dict):
payload: Any = _serialize_row_for_json(message) payload: Any = _serialize_row_for_json(message)

View File

@@ -6,6 +6,8 @@ from typing import Any, Dict, List, Optional
logger = logging.getLogger("packet_broadcaster") logger = logging.getLogger("packet_broadcaster")
_SHUTDOWN_SENTINEL: Dict[str, Any] = {"type": "__broadcaster_shutdown__"}
class PacketBroadcaster: class PacketBroadcaster:
"""Manage subscriber queues and publish packet events.""" """Manage subscriber queues and publish packet events."""
@@ -96,5 +98,6 @@ class PacketBroadcaster:
try: try:
while not queue.empty(): while not queue.empty():
queue.get_nowait() queue.get_nowait()
queue.put_nowait(_SHUTDOWN_SENTINEL)
except Exception: except Exception:
pass pass