| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128 |
- """PrinterScope semantics and WebSocket fan-out filtering (#1727)."""
- from __future__ import annotations
- from types import SimpleNamespace
- from unittest.mock import AsyncMock
- import pytest
- from fastapi import HTTPException
- from backend.app.core.printer_scope import ALL_PRINTERS, PrinterScope
- from backend.app.core.websocket import ConnectionManager
- class TestPrinterScope:
- def test_unrestricted_allows_everything(self):
- assert ALL_PRINTERS.is_unrestricted
- assert ALL_PRINTERS.allows(1)
- assert ALL_PRINTERS.where(None) is None
- def test_restricted_allows_only_its_printers(self):
- scope = PrinterScope(frozenset({1, 2}))
- assert scope.allows(1)
- assert not scope.allows(3)
- assert scope.filter_ids([3, 2, 1]) == [2, 1]
- def test_no_printer_is_always_in_scope(self):
- # Rows not bound to a printer (orphaned archives, model-based jobs)
- assert PrinterScope(frozenset()).allows(None)
- def test_ensure_reports_missing_not_forbidden(self):
- with pytest.raises(HTTPException) as exc:
- PrinterScope(frozenset({1})).ensure(2)
- assert exc.value.status_code == 404
- assert exc.value.detail == "Printer not found"
- def test_intersect(self):
- team = PrinterScope(frozenset({1, 2}))
- assert ALL_PRINTERS.intersect(team) == team
- assert team.intersect(ALL_PRINTERS) == team
- assert team.intersect(PrinterScope(frozenset({2, 3}))) == PrinterScope(frozenset({2}))
- def _socket(scope: PrinterScope | None):
- state = SimpleNamespace()
- if scope is not None:
- state.bambuddy_printer_scope = scope
- return SimpleNamespace(state=state, send_text=AsyncMock())
- class TestBroadcastFiltering:
- @pytest.mark.asyncio
- async def test_printer_events_reach_only_sockets_that_may_see_the_printer(self):
- mgr = ConnectionManager()
- team = _socket(PrinterScope(frozenset({1})))
- everyone = _socket(ALL_PRINTERS)
- mgr.active_connections = [team, everyone]
- await mgr.send_printer_status(2, {})
- team.send_text.assert_not_awaited()
- everyone.send_text.assert_awaited_once()
- @pytest.mark.asyncio
- async def test_printer_id_inside_data_is_honoured(self):
- mgr = ConnectionManager()
- team = _socket(PrinterScope(frozenset({1})))
- mgr.active_connections = [team]
- await mgr.send_archive_created({"id": 9, "printer_id": 2})
- team.send_text.assert_not_awaited()
- await mgr.send_archive_created({"id": 10, "printer_id": 1})
- team.send_text.assert_awaited_once()
- @pytest.mark.asyncio
- async def test_messages_about_no_printer_reach_everyone(self):
- mgr = ConnectionManager()
- team = _socket(PrinterScope(frozenset()))
- mgr.active_connections = [team]
- await mgr.broadcast({"type": "inventory_changed"})
- team.send_text.assert_awaited_once()
- @pytest.mark.asyncio
- async def test_socket_without_a_scope_gets_no_printer_events(self):
- """Fail closed: a socket that missed the connect-time stamp hears nothing printer-bound."""
- mgr = ConnectionManager()
- unstamped = _socket(None)
- mgr.active_connections = [unstamped]
- await mgr.send_printer_status(1, {})
- unstamped.send_text.assert_not_awaited()
- await mgr.broadcast({"type": "inventory_changed"})
- unstamped.send_text.assert_awaited_once()
- @pytest.mark.asyncio
- async def test_targeted_broadcast_is_filtered_too(self):
- mgr = ConnectionManager()
- own = _socket(PrinterScope(frozenset({1})))
- own.state.bambuddy_principal_user_id = 7
- mgr.active_connections = [own]
- await mgr.send_queue_item_acked(7, queue_item_id=1, printer_id=2)
- own.send_text.assert_not_awaited()
- class TestScopeRefresh:
- @pytest.mark.asyncio
- async def test_a_failed_refresh_fails_closed(self):
- """Old scopes may be wider than what was just granted, so they can't be kept."""
- from unittest.mock import patch
- mgr = ConnectionManager()
- socket = _socket(ALL_PRINTERS)
- socket.state.bambuddy_scope_principal = ("member", None)
- socket.close = AsyncMock()
- mgr.active_connections = [socket]
- with (
- patch("backend.app.core.auth.is_auth_enabled", AsyncMock(return_value=True)),
- patch("backend.app.core.auth.principal_printer_scope", AsyncMock(side_effect=RuntimeError("db down"))),
- ):
- await mgr.refresh_printer_scopes()
- assert socket.state.bambuddy_printer_scope == PrinterScope(frozenset())
- socket.close.assert_awaited_once_with(code=4401)
|