from __future__ import annotations import asyncio from contextlib import suppress from datetime import UTC, datetime, timedelta from cloud.internal_api.models import AssignmentModel from device.manager import DeviceManager from host_agent.app import HostAgentApplication, create_application from host_agent.config import HostAgentConfig from storage.device_config import DeviceConfigStore def _config() -> HostAgentConfig: return HostAgentConfig( control_plane_url="https://control.example", host_id="host-a", token="secret", ) def _assignment() -> AssignmentModel: return AssignmentModel( task_id="task-a", attempt=1, lease_id="lease-a", lease_expires_at=datetime.now(UTC) + timedelta(seconds=30), host_id="host-a", device_id="device-a", goal="open settings", ) def test_create_application_composes_host_agent_services(tmp_path, monkeypatch) -> None: monkeypatch.chdir(tmp_path) application = create_application(config=_config(), manager=DeviceManager()) assert isinstance(application, HostAgentApplication) asyncio.run(application.client.aclose()) def test_create_application_loads_persisted_device_configuration( tmp_path, monkeypatch, ) -> None: monkeypatch.chdir(tmp_path) store = DeviceConfigStore(tmp_path / "devices.sqlite3") store.add( device_id="device-a", name="Lab iPhone", driver_type="wda", connection_info={"url": "http://wda.local"}, ) application = create_application( config=_config(), device_config_store=store, ) devices = application.heartbeat.manager.list_devices() assert [(device.id, device.name, device.driver_type) for device in devices] == [ ("device-a", "Lab iPhone", "wda") ] assert devices[0].connection_info == {"url": "http://wda.local"} asyncio.run(application.client.aclose()) def test_shutdown_cancels_long_poll_and_sends_final_heartbeat() -> None: async def scenario() -> None: claim_started = asyncio.Event() claim_cancelled = asyncio.Event() events: list[str] = [] class BlockingClient: async def claim(self): claim_started.set() try: await asyncio.Event().wait() except asyncio.CancelledError: claim_cancelled.set() raise async def aclose(self): events.append("closed") class RecordingHeartbeat: async def run(self, stop): await stop.wait() async def sync_once(self): events.append("final-heartbeat") class IdleProcessor: async def process(self, assignment): raise AssertionError("no assignment expected") def request_stop(self): events.append("stop-work") stop = asyncio.Event() application = HostAgentApplication( client=BlockingClient(), # type: ignore[arg-type] heartbeat=RecordingHeartbeat(), # type: ignore[arg-type] processor=IdleProcessor(), # type: ignore[arg-type] ) running = asyncio.create_task(application.run_async(stop)) await claim_started.wait() stop.set() await asyncio.wait_for(running, timeout=1) assert claim_cancelled.is_set() assert events == ["stop-work", "final-heartbeat", "closed"] asyncio.run(scenario()) def test_shutdown_interrupts_active_work_before_final_heartbeat() -> None: async def scenario() -> None: processing_started = asyncio.Event() processing_stopped = asyncio.Event() events: list[str] = [] claims = 0 class AssignedClient: async def claim(self): nonlocal claims claims += 1 if claims == 1: return _assignment() await asyncio.Event().wait() async def aclose(self): events.append("closed") class RecordingHeartbeat: async def run(self, stop): await stop.wait() async def sync_once(self): events.append("final-heartbeat") class CooperativeProcessor: async def process(self, assignment): processing_started.set() await processing_stopped.wait() events.append("work-finished") def request_stop(self): events.append("stop-work") processing_stopped.set() stop = asyncio.Event() application = HostAgentApplication( client=AssignedClient(), # type: ignore[arg-type] heartbeat=RecordingHeartbeat(), # type: ignore[arg-type] processor=CooperativeProcessor(), # type: ignore[arg-type] ) running = asyncio.create_task(application.run_async(stop)) await processing_started.wait() stop.set() await asyncio.wait_for(running, timeout=1) assert events.index("work-finished") < events.index("final-heartbeat") assert events[-1] == "closed" asyncio.run(scenario()) def test_final_heartbeat_failure_does_not_prevent_client_close() -> None: async def scenario() -> None: closed = False class StoppedClient: async def claim(self): raise AssertionError("polling must not start") async def aclose(self): nonlocal closed closed = True class FailingHeartbeat: async def run(self, stop): await stop.wait() async def sync_once(self): raise OSError("control plane unavailable") class IdleProcessor: def request_stop(self): return None stop = asyncio.Event() stop.set() application = HostAgentApplication( client=StoppedClient(), # type: ignore[arg-type] heartbeat=FailingHeartbeat(), # type: ignore[arg-type] processor=IdleProcessor(), # type: ignore[arg-type] ) await application.run_async(stop) assert closed asyncio.run(scenario()) def test_main_task_cancellation_waits_for_active_work_shutdown() -> None: async def scenario() -> None: processing_started = asyncio.Event() processing_stopped = asyncio.Event() events: list[str] = [] class AssignedClient: async def claim(self): return _assignment() async def aclose(self): events.append("closed") class RecordingHeartbeat: async def run(self, stop): await stop.wait() async def sync_once(self): events.append("final-heartbeat") class CooperativeProcessor: async def process(self, assignment): processing_started.set() await processing_stopped.wait() events.append("work-finished") def request_stop(self): processing_stopped.set() application = HostAgentApplication( client=AssignedClient(), # type: ignore[arg-type] heartbeat=RecordingHeartbeat(), # type: ignore[arg-type] processor=CooperativeProcessor(), # type: ignore[arg-type] ) running = asyncio.create_task(application.run_async()) await processing_started.wait() running.cancel() with suppress(asyncio.CancelledError): await running assert events == ["work-finished", "final-heartbeat", "closed"] asyncio.run(scenario())