try to fix websocket shutdown
All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 1m39s
All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 1m39s
This commit is contained in:
@@ -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)
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user