|
# 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()
|