diff --git a/backend/services/network_share.py b/backend/services/network_share.py index 028b71e8..109b4ecc 100644 --- a/backend/services/network_share.py +++ b/backend/services/network_share.py @@ -63,9 +63,14 @@ class ShareState: lan_addresses: list = field(default_factory=list) -_state = ShareState() -_server: Optional["uvicorn.Server"] = None -_task: Optional["asyncio.Task"] = None +@dataclass +class _ShareRuntime: + state: ShareState = field(default_factory=ShareState) + server: Optional["uvicorn.Server"] = None + task: Optional["asyncio.Task"] = None + + +_runtime = _ShareRuntime() def lan_ipv4_addresses() -> list: @@ -98,19 +103,18 @@ def _find_free_port(base: int, tries: int = 20) -> int: def get_state() -> ShareState: - return _state + return _runtime.state async def enable(app) -> ShareState: - global _server, _task, _state - if _state.enabled: - return _state + if _runtime.state.enabled: + return _runtime.state port = _find_free_port(share_port_base()) pin = _gen_pin() config = uvicorn.Config(app, host="0.0.0.0", port=port, log_level="warning") server = uvicorn.Server(config) server.install_signal_handlers = lambda: None # never hijack signals in-process - _task = asyncio.create_task(server.serve()) + _runtime.task = asyncio.create_task(server.serve()) for _ in range(100): # ~5s for the socket to bind if getattr(server, "started", False): break @@ -121,46 +125,45 @@ async def enable(app) -> ShareState: # with a listener that isn't actually up (spec ยง7). server.should_exit = True try: - await asyncio.wait_for(asyncio.shield(_task), timeout=2) + await asyncio.wait_for(asyncio.shield(_runtime.task), timeout=2) except asyncio.CancelledError: - _server = server - _state = ShareState(True, port, pin, lan_ipv4_addresses()) - app.state.network_share = _state + _runtime.server = server + _runtime.state = ShareState(True, port, pin, lan_ipv4_addresses()) + app.state.network_share = _runtime.state raise except Exception as exc: - if _task.done(): - _server = _task = None - _state = ShareState() - app.state.network_share = _state + if _runtime.task.done(): + _runtime.server = _runtime.task = None + _runtime.state = ShareState() + app.state.network_share = _runtime.state raise RuntimeError("share listener failed to start") from exc - _server = server - _state = ShareState(True, port, pin, lan_ipv4_addresses()) - app.state.network_share = _state + _runtime.server = server + _runtime.state = ShareState(True, port, pin, lan_ipv4_addresses()) + app.state.network_share = _runtime.state logger.warning("Failed LAN listener startup could not be cleaned up") raise RuntimeError( "LAN share listener could not be stopped. Retry Disable before enabling again." ) from exc - _server = _task = None + _runtime.server = _runtime.task = None raise RuntimeError("share listener failed to start") - _server = server - _state = ShareState(True, port, pin, lan_ipv4_addresses()) - app.state.network_share = _state - return _state + _runtime.server = server + _runtime.state = ShareState(True, port, pin, lan_ipv4_addresses()) + app.state.network_share = _runtime.state + return _runtime.state async def disable(app) -> ShareState: - global _server, _task, _state - if _server is not None: - _server.should_exit = True - if _task is not None: + if _runtime.server is not None: + _runtime.server.should_exit = True + if _runtime.task is not None: try: - await asyncio.wait_for(asyncio.shield(_task), timeout=5) + await asyncio.wait_for(asyncio.shield(_runtime.task), timeout=5) except Exception as exc: logger.warning("LAN share listener did not stop; retaining enabled state") raise RuntimeError( "LAN sharing could not be disabled. Retry after active connections close." ) from exc - _server = _task = None - _state = ShareState() - app.state.network_share = _state - return _state + _runtime.server = _runtime.task = None + _runtime.state = ShareState() + app.state.network_share = _runtime.state + return _runtime.state diff --git a/tests/test_network_share.py b/tests/test_network_share.py index a2f05184..30fe101b 100644 --- a/tests/test_network_share.py +++ b/tests/test_network_share.py @@ -87,11 +87,8 @@ def test_pin_only_remote_discovery_never_returns_share_pin(monkeypatch): monkeypatch.setenv("OMNIVOICE_SERVER_MODE", "1") monkeypatch.delenv("OMNIVOICE_API_KEY", raising=False) - monkeypatch.setattr( - live_network_share, - "_state", - live_network_share.ShareState(True, 3901, "123456", ["192.168.1.10"]), - ) + monkeypatch.setattr(live_network_share._runtime, "state", + live_network_share.ShareState(True, 3901, "123456", ["192.168.1.10"])) # Keep the consumption middleware inert: this endpoint is testing the # intentional admin read-only exception itself, before a PIN is supplied. monkeypatch.setattr(app.state, "network_share", None, raising=False) diff --git a/tests/test_truthful_degraded_state.py b/tests/test_truthful_degraded_state.py index 8c11dbfb..d9d22f9b 100644 --- a/tests/test_truthful_degraded_state.py +++ b/tests/test_truthful_degraded_state.py @@ -134,9 +134,9 @@ async def test_network_disable_retains_enabled_state_when_listener_does_not_stop server = SimpleNamespace(should_exit=False) state = network_share.ShareState(True, 3901, "123456", ["192.0.2.1"]) app = SimpleNamespace(state=SimpleNamespace(network_share=state)) - monkeypatch.setattr(network_share, "_server", server) - monkeypatch.setattr(network_share, "_task", task) - monkeypatch.setattr(network_share, "_state", state) + monkeypatch.setattr(network_share._runtime, "server", server) + monkeypatch.setattr(network_share._runtime, "task", task) + monkeypatch.setattr(network_share._runtime, "state", state) with pytest.raises(RuntimeError) as caught: await network_share.disable(app) @@ -239,7 +239,36 @@ async def test_terminal_network_start_failure_resets_state(monkeypatch): await network_share.enable(app) assert "secret" not in str(caught.value) assert network_share.get_state().enabled is False - assert network_share._task is None + assert network_share._runtime.task is None + + +@pytest.mark.asyncio +async def test_live_network_start_cleanup_can_be_retried_by_disable(monkeypatch): + network_share = importlib.import_module("services.network_share") + task = asyncio.get_running_loop().create_future() + server = SimpleNamespace(started=False, should_exit=False, serve=lambda: None) + monkeypatch.setattr(network_share, "_find_free_port", lambda _base: 3901) + monkeypatch.setattr(network_share, "_gen_pin", lambda: "123456") + monkeypatch.setattr(network_share, "lan_ipv4_addresses", lambda: []) + monkeypatch.setattr(network_share.uvicorn, "Server", lambda _config: server) + monkeypatch.setattr(network_share.asyncio, "create_task", lambda _coro: task) + async def no_sleep(_seconds): + return None + monkeypatch.setattr(network_share.asyncio, "sleep", no_sleep) + calls = 0 + async def retryable_wait(_awaitable, timeout): + nonlocal calls + calls += 1 + if calls == 1: + raise asyncio.TimeoutError + return None + monkeypatch.setattr(network_share.asyncio, "wait_for", retryable_wait) + app = SimpleNamespace(state=SimpleNamespace()) + with pytest.raises(RuntimeError): + await network_share.enable(app) + assert network_share.get_state().enabled is True + await network_share.disable(app) + assert network_share.get_state().enabled is False def test_dub_abort_false_result_stays_retryable(monkeypatch):