test_printer_scope.py 4.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128
  1. """PrinterScope semantics and WebSocket fan-out filtering (#1727)."""
  2. from __future__ import annotations
  3. from types import SimpleNamespace
  4. from unittest.mock import AsyncMock
  5. import pytest
  6. from fastapi import HTTPException
  7. from backend.app.core.printer_scope import ALL_PRINTERS, PrinterScope
  8. from backend.app.core.websocket import ConnectionManager
  9. class TestPrinterScope:
  10. def test_unrestricted_allows_everything(self):
  11. assert ALL_PRINTERS.is_unrestricted
  12. assert ALL_PRINTERS.allows(1)
  13. assert ALL_PRINTERS.where(None) is None
  14. def test_restricted_allows_only_its_printers(self):
  15. scope = PrinterScope(frozenset({1, 2}))
  16. assert scope.allows(1)
  17. assert not scope.allows(3)
  18. assert scope.filter_ids([3, 2, 1]) == [2, 1]
  19. def test_no_printer_is_always_in_scope(self):
  20. # Rows not bound to a printer (orphaned archives, model-based jobs)
  21. assert PrinterScope(frozenset()).allows(None)
  22. def test_ensure_reports_missing_not_forbidden(self):
  23. with pytest.raises(HTTPException) as exc:
  24. PrinterScope(frozenset({1})).ensure(2)
  25. assert exc.value.status_code == 404
  26. assert exc.value.detail == "Printer not found"
  27. def test_intersect(self):
  28. team = PrinterScope(frozenset({1, 2}))
  29. assert ALL_PRINTERS.intersect(team) == team
  30. assert team.intersect(ALL_PRINTERS) == team
  31. assert team.intersect(PrinterScope(frozenset({2, 3}))) == PrinterScope(frozenset({2}))
  32. def _socket(scope: PrinterScope | None):
  33. state = SimpleNamespace()
  34. if scope is not None:
  35. state.bambuddy_printer_scope = scope
  36. return SimpleNamespace(state=state, send_text=AsyncMock())
  37. class TestBroadcastFiltering:
  38. @pytest.mark.asyncio
  39. async def test_printer_events_reach_only_sockets_that_may_see_the_printer(self):
  40. mgr = ConnectionManager()
  41. team = _socket(PrinterScope(frozenset({1})))
  42. everyone = _socket(ALL_PRINTERS)
  43. mgr.active_connections = [team, everyone]
  44. await mgr.send_printer_status(2, {})
  45. team.send_text.assert_not_awaited()
  46. everyone.send_text.assert_awaited_once()
  47. @pytest.mark.asyncio
  48. async def test_printer_id_inside_data_is_honoured(self):
  49. mgr = ConnectionManager()
  50. team = _socket(PrinterScope(frozenset({1})))
  51. mgr.active_connections = [team]
  52. await mgr.send_archive_created({"id": 9, "printer_id": 2})
  53. team.send_text.assert_not_awaited()
  54. await mgr.send_archive_created({"id": 10, "printer_id": 1})
  55. team.send_text.assert_awaited_once()
  56. @pytest.mark.asyncio
  57. async def test_messages_about_no_printer_reach_everyone(self):
  58. mgr = ConnectionManager()
  59. team = _socket(PrinterScope(frozenset()))
  60. mgr.active_connections = [team]
  61. await mgr.broadcast({"type": "inventory_changed"})
  62. team.send_text.assert_awaited_once()
  63. @pytest.mark.asyncio
  64. async def test_socket_without_a_scope_gets_no_printer_events(self):
  65. """Fail closed: a socket that missed the connect-time stamp hears nothing printer-bound."""
  66. mgr = ConnectionManager()
  67. unstamped = _socket(None)
  68. mgr.active_connections = [unstamped]
  69. await mgr.send_printer_status(1, {})
  70. unstamped.send_text.assert_not_awaited()
  71. await mgr.broadcast({"type": "inventory_changed"})
  72. unstamped.send_text.assert_awaited_once()
  73. @pytest.mark.asyncio
  74. async def test_targeted_broadcast_is_filtered_too(self):
  75. mgr = ConnectionManager()
  76. own = _socket(PrinterScope(frozenset({1})))
  77. own.state.bambuddy_principal_user_id = 7
  78. mgr.active_connections = [own]
  79. await mgr.send_queue_item_acked(7, queue_item_id=1, printer_id=2)
  80. own.send_text.assert_not_awaited()
  81. class TestScopeRefresh:
  82. @pytest.mark.asyncio
  83. async def test_a_failed_refresh_fails_closed(self):
  84. """Old scopes may be wider than what was just granted, so they can't be kept."""
  85. from unittest.mock import patch
  86. mgr = ConnectionManager()
  87. socket = _socket(ALL_PRINTERS)
  88. socket.state.bambuddy_scope_principal = ("member", None)
  89. socket.close = AsyncMock()
  90. mgr.active_connections = [socket]
  91. with (
  92. patch("backend.app.core.auth.is_auth_enabled", AsyncMock(return_value=True)),
  93. patch("backend.app.core.auth.principal_printer_scope", AsyncMock(side_effect=RuntimeError("db down"))),
  94. ):
  95. await mgr.refresh_printer_scopes()
  96. assert socket.state.bambuddy_printer_scope == PrinterScope(frozenset())
  97. socket.close.assert_awaited_once_with(code=4401)