from __future__ import annotations import asyncio from contextlib import suppress from datetime import UTC, datetime, timedelta from cloud.internal_api.models import ( AssignmentModel, DeviceEnrollmentResponse, HostEnrollmentResponse, ) from device.manager import DeviceManager from host_agent.app import HostAgentApplication, create_application from host_agent.config import HostAgentConfig from host_agent.identity import HostIdentityStore 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_create_application_enrolls_host_and_devices_before_managed_startup( tmp_path, monkeypatch, ) -> None: monkeypatch.chdir(tmp_path) store = DeviceConfigStore(tmp_path / "devices.sqlite3") store.add( device_id="local-device-a", name="Lab iPhone", driver_type="wda", connection_info={"server_url": "http://127.0.0.1:4723"}, ) events: list[str] = [] class EnrollmentClient: def __init__(self) -> None: self.config = HostAgentConfig( control_plane_url="https://control.example", enrollment_token="one-time-token", enrollment_managed=True, ) def enroll_host(self, **payload): events.append(f"host:{payload['agent_instance_id']}") return HostEnrollmentResponse(host_id="host-cloud-a") def enroll_device(self, **payload): events.append(f"device:{payload['local_device_id']}") return DeviceEnrollmentResponse(device_id="device-cloud-a") def close(self): raise AssertionError("injected client must not be closed") identity_store = HostIdentityStore(tmp_path / "host_identity.json") enrollment_client = EnrollmentClient() application = create_application( config=enrollment_client.config, device_config_store=store, identity_store=identity_store, enrollment_client=enrollment_client, # type: ignore[arg-type] ) assert events[0].startswith("host:agent-") assert events[1] == "device:local-device-a" assert application.client.config.host_id == "host-cloud-a" assert application.client.config.enrollment_managed is True assert [device.id for device in application.heartbeat.manager.list_devices()] == [ "device-cloud-a" ] assert store.get("local-device-a")["cloud_device_id"] == "device-cloud-a" assert identity_store.load().host_id == "host-cloud-a" asyncio.run(application.client.aclose()) def test_managed_restart_reuses_identity_and_recovers_device_mapping( tmp_path, monkeypatch, ) -> None: monkeypatch.chdir(tmp_path) store = DeviceConfigStore(tmp_path / "devices.sqlite3") store.add( device_id="local-device-a", driver_type="wda", connection_info={}, ) identity_store = HostIdentityStore(tmp_path / "host_identity.json") identity_store.complete(identity_store.load_or_create(), "host-cloud-a") events: list[str] = [] class EnrollmentClient: config = HostAgentConfig( control_plane_url="https://control.example", identity_path=tmp_path / "host_identity.json", enrollment_managed=True, ) def enroll_host(self, **payload): raise AssertionError("completed identity must skip Host enrollment") def enroll_device(self, **payload): events.append(payload["local_device_id"]) return DeviceEnrollmentResponse(device_id="device-cloud-a") def close(self): return None enrollment_client = EnrollmentClient() application = create_application( config=enrollment_client.config, device_config_store=store, identity_store=identity_store, enrollment_client=enrollment_client, # type: ignore[arg-type] ) assert events == ["local-device-a"] assert application.client.config.host_id == "host-cloud-a" assert store.get("local-device-a")["cloud_device_id"] == "device-cloud-a" 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())