Untitled
unknown
python
10 months ago
2.8 kB
22
Indexable
import asyncio
from typing import Callable, Coroutine, Any
from fastapi import APIRouter, FastAPI, WebSocket, WebSocketDisconnect
from .schemas import MessageSchema, MoveRequestSchema
type Schemas = MessageSchema
type SendSchema = MoveRequestSchema
type CoroutineResult = Coroutine[Any, Any, None]
type BroadcastSchemaFunction = Callable[[SendSchema], CoroutineResult]
type CallbackType = Callable[[BroadcastSchemaFunction, Schemas], CoroutineResult]
class NetworkLayer:
def __init__(self, router: APIRouter | FastAPI) -> None:
self.callback_map: dict[str, list[CallbackType]] = {}
self.register_websocket_router(router)
self.connected_clients: set[WebSocket] = set()
def register_websocket_router(self, router: APIRouter | FastAPI) -> None:
@router.websocket("/ws")
async def websocket_endpoint(websocket: WebSocket):
await websocket.accept()
self.connected_clients.add(websocket)
try:
while True:
data: MessageSchema = await websocket.receive_json()
message_type = data.get("type", "")
if message_type:
await self.trigger_callbacks(message_type, data)
except WebSocketDisconnect:
print("WebSocket disconnected")
except Exception as e:
print(f"WebSocket error: {e}")
def register_callback(self, message_type: str, callback: CallbackType) -> None:
if message_type not in self.callback_map:
self.callback_map[message_type] = []
self.callback_map[message_type].append(callback)
def unregister_callback(self, message_type: str, callback: CallbackType) -> None:
if message_type in self.callback_map:
self.callback_map[message_type].remove(callback)
async def trigger_callbacks(self, message_type: str, data: MessageSchema) -> None:
if message_type in self.callback_map:
for callback in self.callback_map[message_type]:
task = asyncio.create_task(callback(self.broadcast, data))
task.add_done_callback(self._handle_callback_error)
async def broadcast(self, data: SendSchema) -> None:
disconnected_clients = []
for client in self.connected_clients:
try:
await client.send_json(data)
except WebSocketDisconnect:
disconnected_clients.append(client)
except Exception as e:
print(f"Broadcast error: {e}")
for client in disconnected_clients:
self.connected_clients.remove(client)
def _handle_callback_error(self, task: asyncio.Task) -> None:
try:
task.result()
except Exception as e:
print(f"Callback error: {e}")Editor is loading...
Leave a Comment