"""Tests for refresh_state / refresh_status WS broadcasting (daemon.server). refresh_state() and refresh_status() re-collect state and broadcast a data-carrying versions message for every (requested) subsystem so all viewers stay in sync. """ import asyncio import json from contextlib import suppress from unittest.mock import AsyncMock, MagicMock, patch import daemon.server as server def _run_and_drain(fn): """Run *fn* inside a running event loop (required for broadcast tasks), then drain the fire-and-forget broadcast tasks.""" async def drive(): fn() for _ in range(20): await asyncio.sleep(0) if not server._ws_tasks: break for task in list(server._ws_tasks): with suppress(Exception): await task asyncio.run(drive()) def _new_ws(): ws = AsyncMock() ws.send_str = AsyncMock() server._ws_subscribers.add(ws) return ws def _messages(ws): return [json.loads(c[0][0]) for c in ws.send_str.call_args_list] class TestRefreshStateBroadcast: def test_broadcasts_each_requested_subsystem(self): """refresh_state(["firewall","dnsmasq"]) broadcasts both, and bumps.""" store = MagicMock() store.get.side_effect = lambda name: {"s": name} ws = _new_ws() try: with patch.object(server, "state_store", store): _run_and_drain(lambda: server.refresh_state(["firewall", "dnsmasq"])) store.populate.assert_called_once_with(["firewall", "dnsmasq"]) store.bump.assert_any_call("firewall") store.bump.assert_any_call("dnsmasq") msgs = _messages(ws) assert sorted(m["subsystem"] for m in msgs) == ["dnsmasq", "firewall"] assert all(m["type"] == "versions" for m in msgs) for m in msgs: assert "updated" not in m assert m["data"] == {"s": m["subsystem"]} finally: server._ws_subscribers.discard(ws) def test_failed_subsystem_skipped_others_buzz(self): """A subsystem whose collection failed (None) is not broadcast.""" store = MagicMock() store.get.side_effect = lambda name: {"s": name} if name != "acme" else None ws = _new_ws() try: with patch.object(server, "state_store", store): _run_and_drain(lambda: server.refresh_state(["firewall", "acme"])) subs = sorted(m["subsystem"] for m in _messages(ws)) assert subs == ["firewall"] finally: server._ws_subscribers.discard(ws) def test_no_subsystems_arg_broadcasts_all(self): """refresh_state() with no filter targets every subsystem.""" from lib.state import State store = State() for name in State.SUBSYSTEMS: store.set(name, {"k": name}) ws = _new_ws() try: with patch.object(server, "state_store", store): _run_and_drain(lambda: server.refresh_state()) subs = sorted(m["subsystem"] for m in _messages(ws)) assert subs == sorted(State.SUBSYSTEMS) finally: server._ws_subscribers.discard(ws) class TestRefreshStatusBroadcast: def test_filtered_response_and_broadcast(self): """POST /status/refresh replies only with the requested subsystems and broadcasts each of them.""" store = MagicMock() store.get.side_effect = lambda name: {"s": name} ws = _new_ws() try: with patch.object(server, "state_store", store): async def drive(): request = MagicMock() request.json = AsyncMock(return_value={"subsystems": ["firewall"]}) response = await server.refresh_status(request) await asyncio.sleep(0.01) return response response = asyncio.run(drive()) body = json.loads(response.body) assert body["ok"] is True assert set(body["data"]) == {"firewall"} store.bump.assert_not_called() subs = sorted(m["subsystem"] for m in _messages(ws)) assert subs == ["firewall"] finally: server._ws_subscribers.discard(ws) def test_no_body_returns_all_subsystems(self): store = MagicMock() store.get.side_effect = lambda name: {"s": name} store.SUBSYSTEMS = ["firewall", "dnsmasq"] ws = _new_ws() try: with patch.object(server, "state_store", store): async def drive(): request = MagicMock() request.json = AsyncMock(return_value=None) response = await server.refresh_status(request) await asyncio.sleep(0.01) return response response = asyncio.run(drive()) body = json.loads(response.body) assert body["ok"] is True assert set(body["data"]) == {"firewall", "dnsmasq"} finally: server._ws_subscribers.discard(ws)