websocket.py 11 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301
  1. import asyncio
  2. import json
  3. import logging
  4. from typing import Any
  5. from fastapi import WebSocket
  6. logger = logging.getLogger(__name__)
  7. def _message_printer_id(message: dict[str, Any]) -> int | None:
  8. """The printer a broadcast is about, if any: top-level or inside ``data``."""
  9. printer_id = message.get("printer_id")
  10. if printer_id is None:
  11. data = message.get("data")
  12. if isinstance(data, dict):
  13. printer_id = data.get("printer_id")
  14. return printer_id if isinstance(printer_id, int) else None
  15. def _may_receive(connection: WebSocket, printer_id: int | None) -> bool:
  16. """Whether ``connection``'s printer scope (#1727) covers ``printer_id``.
  17. The scope is stamped on the socket at connect (``routes/websocket.py``).
  18. A socket without one is refused anything printer-bound, so a connection
  19. that slipped past the stamping can't receive every printer's events.
  20. """
  21. if printer_id is None:
  22. return True
  23. scope = getattr(connection.state, "bambuddy_printer_scope", None)
  24. return scope is not None and scope.allows(printer_id)
  25. class ConnectionManager:
  26. """Manages WebSocket connections and broadcasts."""
  27. def __init__(self):
  28. self.active_connections: list[WebSocket] = []
  29. self._lock = asyncio.Lock()
  30. async def connect(self, websocket: WebSocket):
  31. """Accept a new WebSocket connection."""
  32. await websocket.accept()
  33. async with self._lock:
  34. self.active_connections.append(websocket)
  35. async def disconnect(self, websocket: WebSocket):
  36. """Remove a WebSocket connection."""
  37. async with self._lock:
  38. if websocket in self.active_connections:
  39. self.active_connections.remove(websocket)
  40. async def broadcast(self, message: dict[str, Any]):
  41. """Broadcast a message to all connected clients."""
  42. if not self.active_connections:
  43. return
  44. data = json.dumps(message)
  45. printer_id = _message_printer_id(message)
  46. async with self._lock:
  47. disconnected = []
  48. for connection in self.active_connections:
  49. if not _may_receive(connection, printer_id):
  50. continue
  51. try:
  52. await connection.send_text(data)
  53. except Exception:
  54. disconnected.append(connection)
  55. # Clean up disconnected clients
  56. for conn in disconnected:
  57. if conn in self.active_connections:
  58. self.active_connections.remove(conn)
  59. async def broadcast_to_user(self, user_id: int | None, message: dict[str, Any]):
  60. """Send a message to every connection authenticated as the given user.
  61. When ``user_id`` is None the message fans out to all connections —
  62. this is the auth-disabled single-user path, where neither the queue
  63. item's ``created_by_id`` nor the WS principal is set, and the
  64. existing fan-out semantics are exactly what the user wants.
  65. Per-user routing reads ``websocket.state.bambuddy_principal_user_id``
  66. stamped at connect time (``routes/websocket.py``). Connections
  67. without a stamped id are skipped on the targeted path so an
  68. anonymous reader never receives another user's dispatch toast.
  69. """
  70. if user_id is None:
  71. await self.broadcast(message)
  72. return
  73. if not self.active_connections:
  74. return
  75. data = json.dumps(message)
  76. printer_id = _message_printer_id(message)
  77. async with self._lock:
  78. disconnected = []
  79. for connection in self.active_connections:
  80. conn_uid = getattr(connection.state, "bambuddy_principal_user_id", None)
  81. if conn_uid != user_id or not _may_receive(connection, printer_id):
  82. continue
  83. try:
  84. await connection.send_text(data)
  85. except Exception:
  86. disconnected.append(connection)
  87. for conn in disconnected:
  88. if conn in self.active_connections:
  89. self.active_connections.remove(conn)
  90. async def refresh_printer_scopes(self):
  91. """Recompute every connection's printer scope (#1727).
  92. Called after an admin changes which printers a group may see, or who
  93. is in a group, so open dashboards stop (or start) receiving those
  94. printers' events without a reconnect.
  95. """
  96. from backend.app.core.auth import is_auth_enabled, principal_printer_scope
  97. from backend.app.core.database import async_session
  98. from backend.app.core.printer_scope import ALL_PRINTERS, PrinterScope
  99. async with self._lock:
  100. connections = list(self.active_connections)
  101. if not connections:
  102. return
  103. try:
  104. async with async_session() as db:
  105. auth_enabled = await is_auth_enabled(db)
  106. for connection in connections:
  107. if not auth_enabled:
  108. connection.state.bambuddy_printer_scope = ALL_PRINTERS
  109. continue
  110. username, api_key_id = getattr(connection.state, "bambuddy_scope_principal", (None, None))
  111. connection.state.bambuddy_printer_scope = await principal_printer_scope(db, username, api_key_id)
  112. except Exception: # SEC-AUTH-EXC: refresh failed → fail closed (empty scope, then disconnect to re-auth)
  113. # The old scopes may be wider than what was just granted, so they
  114. # can't be kept. Drop every socket to no printers and close it with
  115. # the "unauthorised" code: the SPA mints a new token and reconnects,
  116. # and its scope is worked out afresh at connect.
  117. logger.warning("WebSocket printer scope refresh failed; disconnecting clients", exc_info=True)
  118. for connection in connections:
  119. connection.state.bambuddy_printer_scope = PrinterScope(frozenset())
  120. try:
  121. await connection.close(code=4401)
  122. except Exception: # noqa: BLE001 -- already gone; disconnect() cleans it up
  123. pass
  124. async def send_printer_status(self, printer_id: int, status: dict):
  125. """Send printer status update to all clients."""
  126. await self.broadcast(
  127. {
  128. "type": "printer_status",
  129. "printer_id": printer_id,
  130. "data": status,
  131. }
  132. )
  133. async def send_print_start(self, printer_id: int, data: dict):
  134. """Notify clients that a print has started."""
  135. await self.broadcast(
  136. {
  137. "type": "print_start",
  138. "printer_id": printer_id,
  139. "data": data,
  140. }
  141. )
  142. async def send_print_complete(self, printer_id: int, data: dict):
  143. """Notify clients that a print has completed."""
  144. await self.broadcast(
  145. {
  146. "type": "print_complete",
  147. "printer_id": printer_id,
  148. "data": data,
  149. }
  150. )
  151. async def send_print_confirm_request(self, printer_id: int, data: dict):
  152. """Ask connected clients for a post-print outcome verdict (#1898)."""
  153. await self.broadcast(
  154. {
  155. "type": "print_confirm_request",
  156. "printer_id": printer_id,
  157. "data": data,
  158. }
  159. )
  160. async def send_archive_created(self, archive: dict):
  161. """Notify clients that a new archive was created."""
  162. await self.broadcast(
  163. {
  164. "type": "archive_created",
  165. "data": archive,
  166. }
  167. )
  168. async def send_archive_updated(self, archive: dict):
  169. """Notify clients that an archive was updated."""
  170. await self.broadcast(
  171. {
  172. "type": "archive_updated",
  173. "data": archive,
  174. }
  175. )
  176. async def send_queue_item_uploading(
  177. self,
  178. user_id: int | None,
  179. queue_item_id: int,
  180. printer_id: int,
  181. printer_name: str | None,
  182. file_name: str,
  183. total_bytes: int,
  184. ):
  185. """Toast trigger: scheduler picked the item up, FTP upload starts."""
  186. await self.broadcast_to_user(
  187. user_id,
  188. {
  189. "type": "queue_item_uploading",
  190. "queue_item_id": queue_item_id,
  191. "printer_id": printer_id,
  192. "printer_name": printer_name,
  193. "file_name": file_name,
  194. "total_bytes": total_bytes,
  195. },
  196. )
  197. async def send_queue_item_upload_progress(
  198. self,
  199. user_id: int | None,
  200. queue_item_id: int,
  201. bytes_transferred: int,
  202. total_bytes: int,
  203. ):
  204. """Toast update: throttled byte-level progress during the FTP upload."""
  205. pct = int(round(100 * bytes_transferred / total_bytes)) if total_bytes else 0
  206. await self.broadcast_to_user(
  207. user_id,
  208. {
  209. "type": "queue_item_upload_progress",
  210. "queue_item_id": queue_item_id,
  211. "bytes_transferred": bytes_transferred,
  212. "total_bytes": total_bytes,
  213. "pct": pct,
  214. },
  215. )
  216. async def send_queue_item_acked(
  217. self,
  218. user_id: int | None,
  219. queue_item_id: int,
  220. printer_id: int,
  221. ):
  222. """Toast trigger: watchdog confirmed the printer transitioned out of pre_state."""
  223. await self.broadcast_to_user(
  224. user_id,
  225. {
  226. "type": "queue_item_acked",
  227. "queue_item_id": queue_item_id,
  228. "printer_id": printer_id,
  229. },
  230. )
  231. async def send_queue_item_failed(
  232. self,
  233. user_id: int | None,
  234. queue_item_id: int,
  235. printer_id: int | None,
  236. reason: str,
  237. ):
  238. """Toast trigger: dispatch failed at any stage. Toast turns red, auto-dismisses."""
  239. await self.broadcast_to_user(
  240. user_id,
  241. {
  242. "type": "queue_item_failed",
  243. "queue_item_id": queue_item_id,
  244. "printer_id": printer_id,
  245. "reason": reason,
  246. },
  247. )
  248. async def send_missing_spool_assignment(
  249. self,
  250. printer_id: int,
  251. printer_name: str,
  252. missing_slots: list[dict[str, str]],
  253. ):
  254. """Notify clients that a print started with missing spool assignments."""
  255. await self.broadcast(
  256. {
  257. "type": "missing_spool_assignment",
  258. "printer_id": printer_id,
  259. "printer_name": printer_name,
  260. "missing_slots": missing_slots,
  261. }
  262. )
  263. # Global connection manager
  264. ws_manager = ConnectionManager()