fix: centralize LAN share runtime state
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user