Coverage for server / utilities / websocket_manager.py: 100%
48 statements
« prev ^ index » next coverage.py v7.13.4, created at 2026-10-04 09:33 +0000
« prev ^ index » next coverage.py v7.13.4, created at 2026-10-04 09:33 +0000
1"""
2EI-2966 — Lightweight WebSocket connection manager.
4Maintains per-user connection registries so a service can push events
5to all of a given user's open dashboard tabs without iterating the
6whole global pool. Used by the student-dashboard WebSocket route to
7deliver real-time refresh signals.
9Thread-safety: a single async lock guards the registry. Each method
10is small enough that the lock is held only for the dict mutation,
11not the network I/O.
13Connect lifecycle (caller's responsibility):
14 - On connect: `await manager.connect(user_id, websocket)`
15 - On disconnect (any reason): `await manager.disconnect(user_id, websocket)`
16 - To push: `await manager.broadcast_to_user(user_id, payload)`
17"""
19from __future__ import annotations
21import asyncio
22import json
23from typing import Any, Dict, Set
25from fastapi import WebSocket
28class StudentDashboardWebSocketManager:
29 def __init__(self) -> None:
30 self._connections: Dict[str, Set[WebSocket]] = {}
31 self._lock = asyncio.Lock()
33 async def connect(self, user_id: str, websocket: WebSocket) -> None:
34 async with self._lock:
35 if user_id not in self._connections:
36 self._connections[user_id] = set()
37 self._connections[user_id].add(websocket)
39 async def disconnect(self, user_id: str, websocket: WebSocket) -> None:
40 async with self._lock:
41 if user_id in self._connections:
42 self._connections[user_id].discard(websocket)
43 if not self._connections[user_id]:
44 del self._connections[user_id]
46 async def broadcast_to_user(self, user_id: str, payload: Dict[str, Any]) -> int:
47 """Send `payload` (JSON-serialized) to every connected socket for `user_id`.
49 Returns the number of sockets the payload was successfully sent to.
50 Sockets that error out are silently removed (caller's `disconnect`
51 cleanup is the primary path; this is just defense in depth).
52 """
53 async with self._lock:
54 sockets = list(self._connections.get(user_id, set()))
56 if not sockets:
57 return 0
59 delivered = 0
60 body = json.dumps(payload)
61 dead: list[WebSocket] = []
62 for ws in sockets:
63 try:
64 await ws.send_text(body)
65 delivered += 1
66 except Exception: # noqa: BLE001 — closed/broken sockets are silently culled
67 dead.append(ws)
69 if dead:
70 async with self._lock:
71 pool = self._connections.get(user_id)
72 if pool is not None:
73 for d in dead:
74 pool.discard(d)
75 if not pool:
76 self._connections.pop(user_id, None)
78 return delivered
80 def count_connections(self, user_id: str) -> int:
81 """Diagnostic helper: how many tabs/devices does this user have open?"""
82 return len(self._connections.get(user_id, set()))
84 def total_connections(self) -> int:
85 """Diagnostic helper: how many WS connections are open in total?"""
86 return sum(len(s) for s in self._connections.values())
89# Module-level singleton so the route + publish sites share one registry.
90student_dashboard_ws_manager = StudentDashboardWebSocketManager()