# retoor <retoor@molodetz.nl>
import logging
from collections import OrderedDict
from typing import Any
from fastapi import WebSocket
logger = logging.getLogger("messaging.hub")
DELIVERED_CAP = 4000
class ConnectionManager:
def __init__(self) -> None:
self._connections: dict[str, set[WebSocket]] = {}
self._delivered: "OrderedDict[str, bool]" = OrderedDict()
def register(self, user_uid: str, websocket: WebSocket) -> bool:
sockets = self._connections.setdefault(user_uid, set())
was_offline = len(sockets) == 0
sockets.add(websocket)
logger.info(
"messaging socket registered for %s (sockets=%d)", user_uid, len(sockets)
)
return was_offline
def unregister(self, user_uid: str, websocket: WebSocket) -> bool:
sockets = self._connections.get(user_uid)
if not sockets:
return False
sockets.discard(websocket)
if not sockets:
self._connections.pop(user_uid, None)
logger.info("messaging socket removed for %s (now offline)", user_uid)
return True
logger.debug(
"messaging socket removed for %s (sockets=%d)", user_uid, len(sockets)
)
return False
def has_connections(self) -> bool:
return bool(self._connections)
def connected_user_uids(self) -> set[str]:
return set(self._connections.keys())
def mark_delivered(self, message_uid: str) -> None:
self._delivered[message_uid] = True
self._delivered.move_to_end(message_uid)
while len(self._delivered) > DELIVERED_CAP:
self._delivered.popitem(last=False)
def was_delivered(self, message_uid: str) -> bool:
return message_uid in self._delivered
def sockets_for(self, user_uid: str) -> list[WebSocket]:
return list(self._connections.get(user_uid, ()))
async def send_to_user(self, user_uid: str, frame: dict[str, Any]) -> None:
for websocket in self.sockets_for(user_uid):
try:
await websocket.send_json(frame)
except Exception: # noqa: BLE001
logger.debug("dropping frame to stale socket for %s", user_uid)
async def send_to_users(
self, user_uids: list[str], frame: dict[str, Any]
) -> None:
for user_uid in dict.fromkeys(user_uids):
await self.send_to_user(user_uid, frame)
message_hub = ConnectionManager()