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.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)
|
||||
|
||||
Reference in New Issue
Block a user