| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301 |
- import asyncio
- import json
- import logging
- from typing import Any
- from fastapi import WebSocket
- logger = logging.getLogger(__name__)
- def _message_printer_id(message: dict[str, Any]) -> int | None:
- """The printer a broadcast is about, if any: top-level or inside ``data``."""
- printer_id = message.get("printer_id")
- if printer_id is None:
- data = message.get("data")
- if isinstance(data, dict):
- printer_id = data.get("printer_id")
- return printer_id if isinstance(printer_id, int) else None
- def _may_receive(connection: WebSocket, printer_id: int | None) -> bool:
- """Whether ``connection``'s printer scope (#1727) covers ``printer_id``.
- The scope is stamped on the socket at connect (``routes/websocket.py``).
- A socket without one is refused anything printer-bound, so a connection
- that slipped past the stamping can't receive every printer's events.
- """
- if printer_id is None:
- return True
- scope = getattr(connection.state, "bambuddy_printer_scope", None)
- return scope is not None and scope.allows(printer_id)
- class ConnectionManager:
- """Manages WebSocket connections and broadcasts."""
- def __init__(self):
- self.active_connections: list[WebSocket] = []
- self._lock = asyncio.Lock()
- async def connect(self, websocket: WebSocket):
- """Accept a new WebSocket connection."""
- await websocket.accept()
- async with self._lock:
- self.active_connections.append(websocket)
- async def disconnect(self, websocket: WebSocket):
- """Remove a WebSocket connection."""
- async with self._lock:
- if websocket in self.active_connections:
- self.active_connections.remove(websocket)
- async def broadcast(self, message: dict[str, Any]):
- """Broadcast a message to all connected clients."""
- if not self.active_connections:
- return
- data = json.dumps(message)
- printer_id = _message_printer_id(message)
- async with self._lock:
- disconnected = []
- for connection in self.active_connections:
- if not _may_receive(connection, printer_id):
- continue
- try:
- await connection.send_text(data)
- except Exception:
- disconnected.append(connection)
- # Clean up disconnected clients
- for conn in disconnected:
- if conn in self.active_connections:
- self.active_connections.remove(conn)
- async def broadcast_to_user(self, user_id: int | None, message: dict[str, Any]):
- """Send a message to every connection authenticated as the given user.
- When ``user_id`` is None the message fans out to all connections —
- this is the auth-disabled single-user path, where neither the queue
- item's ``created_by_id`` nor the WS principal is set, and the
- existing fan-out semantics are exactly what the user wants.
- Per-user routing reads ``websocket.state.bambuddy_principal_user_id``
- stamped at connect time (``routes/websocket.py``). Connections
- without a stamped id are skipped on the targeted path so an
- anonymous reader never receives another user's dispatch toast.
- """
- if user_id is None:
- await self.broadcast(message)
- return
- if not self.active_connections:
- return
- data = json.dumps(message)
- printer_id = _message_printer_id(message)
- async with self._lock:
- disconnected = []
- for connection in self.active_connections:
- conn_uid = getattr(connection.state, "bambuddy_principal_user_id", None)
- if conn_uid != user_id or not _may_receive(connection, printer_id):
- continue
- try:
- await connection.send_text(data)
- except Exception:
- disconnected.append(connection)
- for conn in disconnected:
- if conn in self.active_connections:
- self.active_connections.remove(conn)
- async def refresh_printer_scopes(self):
- """Recompute every connection's printer scope (#1727).
- Called after an admin changes which printers a group may see, or who
- is in a group, so open dashboards stop (or start) receiving those
- printers' events without a reconnect.
- """
- from backend.app.core.auth import is_auth_enabled, principal_printer_scope
- from backend.app.core.database import async_session
- from backend.app.core.printer_scope import ALL_PRINTERS, PrinterScope
- async with self._lock:
- connections = list(self.active_connections)
- if not connections:
- return
- try:
- async with async_session() as db:
- auth_enabled = await is_auth_enabled(db)
- for connection in connections:
- if not auth_enabled:
- connection.state.bambuddy_printer_scope = ALL_PRINTERS
- continue
- username, api_key_id = getattr(connection.state, "bambuddy_scope_principal", (None, None))
- connection.state.bambuddy_printer_scope = await principal_printer_scope(db, username, api_key_id)
- except Exception: # SEC-AUTH-EXC: refresh failed → fail closed (empty scope, then disconnect to re-auth)
- # The old scopes may be wider than what was just granted, so they
- # can't be kept. Drop every socket to no printers and close it with
- # the "unauthorised" code: the SPA mints a new token and reconnects,
- # and its scope is worked out afresh at connect.
- logger.warning("WebSocket printer scope refresh failed; disconnecting clients", exc_info=True)
- for connection in connections:
- connection.state.bambuddy_printer_scope = PrinterScope(frozenset())
- try:
- await connection.close(code=4401)
- except Exception: # noqa: BLE001 -- already gone; disconnect() cleans it up
- pass
- async def send_printer_status(self, printer_id: int, status: dict):
- """Send printer status update to all clients."""
- await self.broadcast(
- {
- "type": "printer_status",
- "printer_id": printer_id,
- "data": status,
- }
- )
- async def send_print_start(self, printer_id: int, data: dict):
- """Notify clients that a print has started."""
- await self.broadcast(
- {
- "type": "print_start",
- "printer_id": printer_id,
- "data": data,
- }
- )
- async def send_print_complete(self, printer_id: int, data: dict):
- """Notify clients that a print has completed."""
- await self.broadcast(
- {
- "type": "print_complete",
- "printer_id": printer_id,
- "data": data,
- }
- )
- async def send_print_confirm_request(self, printer_id: int, data: dict):
- """Ask connected clients for a post-print outcome verdict (#1898)."""
- await self.broadcast(
- {
- "type": "print_confirm_request",
- "printer_id": printer_id,
- "data": data,
- }
- )
- async def send_archive_created(self, archive: dict):
- """Notify clients that a new archive was created."""
- await self.broadcast(
- {
- "type": "archive_created",
- "data": archive,
- }
- )
- async def send_archive_updated(self, archive: dict):
- """Notify clients that an archive was updated."""
- await self.broadcast(
- {
- "type": "archive_updated",
- "data": archive,
- }
- )
- async def send_queue_item_uploading(
- self,
- user_id: int | None,
- queue_item_id: int,
- printer_id: int,
- printer_name: str | None,
- file_name: str,
- total_bytes: int,
- ):
- """Toast trigger: scheduler picked the item up, FTP upload starts."""
- await self.broadcast_to_user(
- user_id,
- {
- "type": "queue_item_uploading",
- "queue_item_id": queue_item_id,
- "printer_id": printer_id,
- "printer_name": printer_name,
- "file_name": file_name,
- "total_bytes": total_bytes,
- },
- )
- async def send_queue_item_upload_progress(
- self,
- user_id: int | None,
- queue_item_id: int,
- bytes_transferred: int,
- total_bytes: int,
- ):
- """Toast update: throttled byte-level progress during the FTP upload."""
- pct = int(round(100 * bytes_transferred / total_bytes)) if total_bytes else 0
- await self.broadcast_to_user(
- user_id,
- {
- "type": "queue_item_upload_progress",
- "queue_item_id": queue_item_id,
- "bytes_transferred": bytes_transferred,
- "total_bytes": total_bytes,
- "pct": pct,
- },
- )
- async def send_queue_item_acked(
- self,
- user_id: int | None,
- queue_item_id: int,
- printer_id: int,
- ):
- """Toast trigger: watchdog confirmed the printer transitioned out of pre_state."""
- await self.broadcast_to_user(
- user_id,
- {
- "type": "queue_item_acked",
- "queue_item_id": queue_item_id,
- "printer_id": printer_id,
- },
- )
- async def send_queue_item_failed(
- self,
- user_id: int | None,
- queue_item_id: int,
- printer_id: int | None,
- reason: str,
- ):
- """Toast trigger: dispatch failed at any stage. Toast turns red, auto-dismisses."""
- await self.broadcast_to_user(
- user_id,
- {
- "type": "queue_item_failed",
- "queue_item_id": queue_item_id,
- "printer_id": printer_id,
- "reason": reason,
- },
- )
- async def send_missing_spool_assignment(
- self,
- printer_id: int,
- printer_name: str,
- missing_slots: list[dict[str, str]],
- ):
- """Notify clients that a print started with missing spool assignments."""
- await self.broadcast(
- {
- "type": "missing_spool_assignment",
- "printer_id": printer_id,
- "printer_name": printer_name,
- "missing_slots": missing_slots,
- }
- )
- # Global connection manager
- ws_manager = ConnectionManager()
|