websocket.py 8.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181
  1. """GHSA-r2qv follow-up — WebSocket auth gate.
  2. Previously ``/api/v1/ws`` accepted *any* network client and immediately
  3. streamed every ``printer_status`` / ``print_start`` / ``print_complete``
  4. / ``archive_*`` / ``inventory_changed`` broadcast back to it. That is
  5. the GHSA-gc24 shape on a different protocol — anyone who could reach
  6. the HTTP port could subscribe to every printer event in the system.
  7. This endpoint now validates a short-lived token (minted by
  8. ``POST /api/v1/auth/ws-token`` behind ``Permission.WEBSOCKET_CONNECT``)
  9. *before* ``websocket.accept()``. When auth is disabled, no token is
  10. required (the legacy SPA-friendly path). The token is reused across
  11. reconnects within its 60-minute window so a brief network blip does
  12. not require a round-trip to the auth router.
  13. """
  14. from __future__ import annotations
  15. import logging
  16. from fastapi import APIRouter, Query, WebSocket, WebSocketDisconnect
  17. from sqlalchemy import select
  18. from backend.app.core.auth import is_auth_enabled, principal_printer_scope, verify_websocket_token_principal
  19. from backend.app.core.database import async_session
  20. from backend.app.core.printer_scope import ALL_PRINTERS, PrinterScope
  21. from backend.app.core.websocket import ws_manager
  22. from backend.app.models.user import User
  23. from backend.app.services.printer_manager import printer_manager, printer_state_to_dict
  24. logger = logging.getLogger(__name__)
  25. router = APIRouter()
  26. # 4401 mirrors the WebSocket "unauthorised" application close code
  27. # convention used by Sec-WebSocket-Protocol authors (private-use range
  28. # is 4000-4999 per RFC 6455). The SPA distinguishes 4401 from network
  29. # drops and refetches a token instead of retrying with the old one.
  30. _WS_CLOSE_UNAUTHORIZED = 4401
  31. @router.websocket("/ws")
  32. async def websocket_endpoint(websocket: WebSocket, token: str | None = Query(default=None)) -> None:
  33. """WebSocket endpoint for real-time updates.
  34. Connection auth (GHSA-r2qv follow-up):
  35. - Auth disabled → connect without a token, identical to the prior
  36. behaviour (single-user / local-network deployments).
  37. - Auth enabled → ``?token=<value>`` query param must hold an
  38. unexpired token minted via ``POST /api/v1/auth/ws-token``.
  39. Missing / invalid / expired token → ``close(code=4401)`` *before*
  40. ``accept()`` so no ``ws_manager.broadcast`` ever reaches the
  41. caller (broadcasts walk ``active_connections`` blindly — letting
  42. an unauthenticated socket into that list is a fan-out leak).
  43. The auth check is fail-closed at every error path: a DB exception
  44. while reading the ``auth_enabled`` setting closes the connection
  45. rather than admitting the caller.
  46. """
  47. # Authenticate before accept() so an unauth caller never lands in
  48. # ws_manager.active_connections (where broadcasts blindly fan out).
  49. try:
  50. async with async_session() as db:
  51. auth_required = await is_auth_enabled(db)
  52. except Exception: # SEC-AUTH-EXC: DB failure on auth probe → fail-closed (refuse connect), matches is_auth_enabled itself which returns True on error
  53. logger.error("WebSocket auth probe failed; refusing connection", exc_info=True)
  54. await websocket.close(code=_WS_CLOSE_UNAUTHORIZED)
  55. return
  56. principal: str | None = None
  57. api_key_id: int | None = None
  58. printer_scope: PrinterScope = ALL_PRINTERS
  59. if auth_required:
  60. if not token:
  61. logger.info("WebSocket connect refused: no token (auth enabled)")
  62. await websocket.close(code=_WS_CLOSE_UNAUTHORIZED)
  63. return
  64. token_principal = await verify_websocket_token_principal(token)
  65. if token_principal is None:
  66. logger.info("WebSocket connect refused: invalid or expired token")
  67. await websocket.close(code=_WS_CLOSE_UNAUTHORIZED)
  68. return
  69. principal, api_key_id = token_principal
  70. # Which printers this socket may hear about (#1727). Fail-closed: if
  71. # it can't be worked out the socket is refused rather than admitted
  72. # with every printer.
  73. try:
  74. async with async_session() as db:
  75. printer_scope = await principal_printer_scope(db, principal, api_key_id)
  76. except Exception: # SEC-AUTH-EXC: scope lookup failure → refuse connect (fail-closed)
  77. logger.error("WebSocket printer scope lookup failed; refusing connection", exc_info=True)
  78. await websocket.close(code=_WS_CLOSE_UNAUTHORIZED)
  79. return
  80. # Token verified (or auth disabled); now safe to admit the connection.
  81. logger.info("WebSocket client connecting (principal=%s)", principal if principal else "<anonymous>")
  82. # Stamped before connect() puts the socket in the broadcast list, so no
  83. # broadcast can reach it unfiltered (ws_manager refuses printer-bound
  84. # messages to a socket without a scope anyway).
  85. websocket.state.bambuddy_printer_scope = printer_scope
  86. websocket.state.bambuddy_scope_principal = (principal, api_key_id)
  87. await ws_manager.connect(websocket)
  88. # Stash on connection state for any future per-message permission
  89. # logic; today the message handlers are read-only and only respond
  90. # to the requesting socket, so the stash is informational. The
  91. # explicit attribute (rather than a side dict) means a future
  92. # ``broadcast_to_principal()`` helper can filter on it without
  93. # touching every call site.
  94. websocket.state.bambuddy_principal = principal
  95. # Resolve principal username → User.id once at connect so
  96. # ``ws_manager.broadcast_to_user()`` can filter without re-querying
  97. # per message. Auth-disabled path keeps None (broadcast_to_user fans
  98. # out to all when target is None — matches the legacy single-user
  99. # toast behaviour). API-keyed principal is empty string → None.
  100. principal_user_id: int | None = None
  101. if principal:
  102. try:
  103. async with async_session() as db:
  104. row = await db.execute(select(User.id).where(User.username == principal))
  105. principal_user_id = row.scalar_one_or_none()
  106. except Exception: # SEC-AUTH-EXC: resolution failure is non-fatal — degrades to no per-user routing
  107. logger.warning("WebSocket principal resolve failed for %s", principal, exc_info=True)
  108. websocket.state.bambuddy_principal_user_id = principal_user_id
  109. logger.info("WebSocket client connected")
  110. try:
  111. # Send initial status of all printers.
  112. statuses = {
  113. pid: state
  114. for pid, state in printer_manager.get_all_statuses().items()
  115. if websocket.state.bambuddy_printer_scope.allows(pid)
  116. }
  117. for printer_id, state in statuses.items():
  118. await websocket.send_json(
  119. {
  120. "type": "printer_status",
  121. "printer_id": printer_id,
  122. "data": printer_state_to_dict(
  123. state,
  124. printer_id,
  125. printer_manager.get_model(printer_id),
  126. printer_manager.get_drying_targets(printer_id),
  127. ),
  128. }
  129. )
  130. logger.info("Sent initial status for %s printers", len(statuses))
  131. # Keep connection alive and handle incoming messages.
  132. while True:
  133. data = await websocket.receive_json()
  134. # Handle ping/pong for keepalive
  135. if data.get("type") == "ping":
  136. await websocket.send_json({"type": "pong"})
  137. # Handle status request
  138. elif data.get("type") == "get_status":
  139. printer_id = data.get("printer_id")
  140. if printer_id and websocket.state.bambuddy_printer_scope.allows(printer_id):
  141. state = printer_manager.get_status(printer_id)
  142. if state:
  143. await websocket.send_json(
  144. {
  145. "type": "printer_status",
  146. "printer_id": printer_id,
  147. "data": printer_state_to_dict(
  148. state,
  149. printer_id,
  150. printer_manager.get_model(printer_id),
  151. printer_manager.get_drying_targets(printer_id),
  152. ),
  153. }
  154. )
  155. except WebSocketDisconnect:
  156. logger.info("WebSocket client disconnected normally")
  157. await ws_manager.disconnect(websocket)
  158. except Exception as e:
  159. logger.error("WebSocket error: %s", e, exc_info=True)
  160. await ws_manager.disconnect(websocket)