from __future__ import annotations import asyncio import socket from contextlib import suppress from datetime import UTC, datetime, timedelta import httpx import pytest from starlette.testclient import TestClient 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.assignment import AssignmentExecutor from host_agent.config import HostAgentConfig from host_agent.execution import create_execution_factories from host_agent.history import ConsoleHistoryStore from host_agent.identity import HostIdentityStore from host_agent.instance_lock import InstanceAlreadyRunningError from host_agent.local_account import LocalAccountStore from host_agent.mcp_lock import McpBusyTracker from host_agent.mcp_token import McpTokenStore from host_agent.status import AgentStatusTracker from host_agent.web.app import create_console_app from host_agent.web.auth import SessionManager from host_agent.web.mcp import build_mcp_server from storage.artifact_store import ArtifactStore from storage.device_config import DeviceConfigStore from storage.task_metadata import TaskMetadataStore from storage.timeline import Timeline def _free_loopback_port() -> int: with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as probe: probe.bind(("127.0.0.1", 0)) return probe.getsockname()[1] 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_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()) def test_console_serves_http_and_shuts_down_cleanly(tmp_path) -> None: async def scenario() -> None: port = _free_loopback_port() config = HostAgentConfig( control_plane_url="https://control.example", host_id="host-a", token="secret", identity_path=tmp_path / "host_identity.json", local_account_path=tmp_path / "host_local_account.json", console_bind_host="127.0.0.1", console_port=port, ) application = create_application(config=config, manager=DeviceManager()) assert application.console_server is not 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") application.client = BlockingClient() # type: ignore[assignment] application.heartbeat = RecordingHeartbeat() # type: ignore[assignment] application.processor = IdleProcessor() # type: ignore[assignment] stop = asyncio.Event() running = asyncio.create_task(application.run_async(stop)) await claim_started.wait() response: httpx.Response | None = None async with httpx.AsyncClient() as http_client: loop = asyncio.get_running_loop() deadline = loop.time() + 5 while loop.time() < deadline: try: response = await http_client.get( f"http://127.0.0.1:{port}/login", timeout=0.5 ) except httpx.TransportError: await asyncio.sleep(0.05) continue break assert response is not None assert response.status_code == 200 assert "Login" in response.text stop.set() await asyncio.wait_for(running, timeout=5) assert claim_cancelled.is_set() assert events == ["stop-work", "final-heartbeat", "closed"] assert application.console_server.should_exit is True with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as probe: probe.settimeout(0.5) with suppress(ConnectionRefusedError, OSError): probe.connect(("127.0.0.1", port)) raise AssertionError("console socket should be closed after shutdown") asyncio.run(scenario()) def test_console_is_created_by_default(tmp_path, monkeypatch) -> None: monkeypatch.chdir(tmp_path) application = create_application(config=_config(), manager=DeviceManager()) assert application.console_server is not None asyncio.run(application.client.aclose()) def test_dependency_supervisor_is_none_when_disabled(tmp_path, monkeypatch) -> None: monkeypatch.chdir(tmp_path) application = create_application(config=_config(), manager=DeviceManager()) assert application.dependency_supervisor is None asyncio.run(application.client.aclose()) def test_dependency_supervisor_constructed_when_enabled_with_no_deps( tmp_path, monkeypatch ) -> None: monkeypatch.chdir(tmp_path) config = HostAgentConfig( control_plane_url="https://control.example", host_id="host-a", token="secret", dependency_supervisor_enabled=True, ) application = create_application(config=config, manager=DeviceManager()) assert application.dependency_supervisor is not None assert application.dependency_supervisor.dependencies == [] asyncio.run(application.client.aclose()) def test_dependency_supervisor_constructed_when_enabled_with_appium_only( tmp_path, monkeypatch ) -> None: monkeypatch.chdir(tmp_path) config = HostAgentConfig( control_plane_url="https://control.example", host_id="host-a", token="secret", dependency_supervisor_enabled=True, appium_supervised=True, appium_host="127.0.0.1", appium_port=4723, ) application = create_application(config=config, manager=DeviceManager()) assert application.dependency_supervisor is not None deps = application.dependency_supervisor.dependencies assert [dep.name for dep in deps] == ["appium"] asyncio.run(application.client.aclose()) def test_run_async_starts_supervisor_before_first_heartbeat_connect() -> None: async def scenario() -> None: events: list[str] = [] class SupervisedNoOp: def __init__(self) -> None: self.started = False self.stopped = False async def start(self) -> None: self.started = True events.append("supervisor-start") async def run(self, stop: asyncio.Event) -> None: events.append("supervisor-run-entered") await stop.wait() async def stop(self) -> None: self.stopped = True events.append("supervisor-stop") class BlockingClient: async def claim(self): await asyncio.Event().wait() async def aclose(self): events.append("closed") class RecordingHeartbeat: def __init__(self) -> None: self.connect_called = False def connect_devices(self) -> None: self.connect_called = True events.append("connect-devices") async def run(self, stop: asyncio.Event) -> None: # Mirror HeartbeatSynchronizer.run which calls connect_devices() # at the very top — supervisor must have started already. self.connect_devices() 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") supervisor = SupervisedNoOp() application = HostAgentApplication( client=BlockingClient(), # type: ignore[arg-type] heartbeat=RecordingHeartbeat(), # type: ignore[arg-type] processor=IdleProcessor(), # type: ignore[arg-type] dependency_supervisor=supervisor, # type: ignore[arg-type] ) stop = asyncio.Event() running = asyncio.create_task(application.run_async(stop)) # Yield long enough for startup sequencing to land. await asyncio.sleep(0.05) stop.set() await asyncio.wait_for(running, timeout=1.0) assert supervisor.started is True assert supervisor.stopped is True # Supervisor startup must precede the heartbeat's connect_devices(). assert events.index("supervisor-start") < events.index("connect-devices") # Supervisor stop must run before client close. assert events.index("supervisor-stop") < events.index("closed") asyncio.run(scenario()) def test_second_create_application_against_held_lock_raises_before_enrollment( tmp_path, monkeypatch ) -> None: monkeypatch.chdir(tmp_path) identity_path = tmp_path / "host_identity.json" config = HostAgentConfig( control_plane_url="https://control.example", host_id="host-a", token="secret", identity_path=identity_path, ) class TrackingEnrollmentClient: def __init__(self) -> None: self.config = config self.calls: list[str] = [] def enroll_host(self, **payload): self.calls.append("host") return HostEnrollmentResponse(host_id="host-a") def enroll_device(self, **payload): self.calls.append("device") return DeviceEnrollmentResponse(device_id="device-a") def close(self) -> None: return None first_client = TrackingEnrollmentClient() first_app = create_application( config=config, identity_store=HostIdentityStore(identity_path), enrollment_client=first_client, # type: ignore[arg-type] ) try: second_client = TrackingEnrollmentClient() with pytest.raises(InstanceAlreadyRunningError) as info: create_application( config=config, identity_store=HostIdentityStore(identity_path), enrollment_client=second_client, # type: ignore[arg-type] ) assert info.value.lock_path == identity_path.parent / "host_agent.lock" assert second_client.calls == [] finally: asyncio.run(first_app.client.aclose()) def test_create_application_with_independent_identity_paths_coexist( tmp_path, monkeypatch ) -> None: monkeypatch.chdir(tmp_path) config_a = HostAgentConfig( control_plane_url="https://control.example", host_id="host-a", token="secret", identity_path=tmp_path / "identity-a" / "host_identity.json", ) config_b = HostAgentConfig( control_plane_url="https://control.example", host_id="host-b", token="secret", identity_path=tmp_path / "identity-b" / "host_identity.json", ) app_a = create_application(config=config_a, manager=DeviceManager()) try: app_b = create_application(config=config_b, manager=DeviceManager()) asyncio.run(app_b.client.aclose()) finally: asyncio.run(app_a.client.aclose()) def test_create_application_wires_mcp_components(tmp_path, monkeypatch) -> None: """create_application produces a console app with /mcp mounted (auth-protected) and persists the host_mcp_token.json file alongside the identity.""" monkeypatch.chdir(tmp_path) config = _config() config_store = DeviceConfigStore(tmp_path / "devices.sqlite3") identity_store = HostIdentityStore(config.identity_path) history_store = ConsoleHistoryStore( tmp_path / "host_console_history.sqlite3", limit=config.console_history_limit, ) metadata_store = TaskMetadataStore(db_path=config.task_progress_db_path) timeline = Timeline(ArtifactStore(root=config.task_artifact_dir)) status_tracker = AgentStatusTracker() application = create_application( config=config, device_config_store=config_store, identity_store=identity_store, manager=DeviceManager(), ) # Token file must exist after create_application. assert (config.identity_path.parent / "host_mcp_token.json").exists() # Heartbeat must hold the in-process McpBusyTracker. assert application.heartbeat.mcp_busy_tracker is not None # Build the same console app the production path builds and verify /mcp # is mounted (responds 401, not 404) without a bearer token. mcp_token_store = McpTokenStore(config.identity_path.parent / "host_mcp_token.json") mcp_token_store.load_or_create() mcp_busy_tracker = McpBusyTracker(ttl_seconds=60.0) mcp_server = build_mcp_server( manager=application.heartbeat.manager, mcp_busy_tracker=mcp_busy_tracker, status_tracker=status_tracker, ) console_app = create_console_app( config=config, manager=application.heartbeat.manager, config_store=config_store, local_account_store=LocalAccountStore(config.local_account_path), identity_store=identity_store, history_store=history_store, status_tracker=status_tracker, session_manager=SessionManager(ttl_seconds=config.console_session_ttl_seconds), enrollment_client=None, host_client=application.client, metadata_store=metadata_store, timeline=timeline, executor=AssignmentExecutor( create_execution_factories( application.heartbeat.manager, metadata_store=metadata_store, timeline=timeline, host_agent_config=config, ), mcp_busy_tracker=mcp_busy_tracker, ), mcp_server=mcp_server, mcp_token_store=mcp_token_store, mcp_busy_tracker=mcp_busy_tracker, ) with TestClient(console_app) as client: resp = client.post("/mcp/") assert resp.status_code == 401 # auth required, not 404 asyncio.run(application.client.aclose()) def test_lock_released_after_run_async_allows_restart(tmp_path, monkeypatch) -> None: monkeypatch.chdir(tmp_path) identity_path = tmp_path / "host_identity.json" config = HostAgentConfig( control_plane_url="https://control.example", host_id="host-a", token="secret", identity_path=identity_path, ) first_app = create_application(config=config, manager=DeviceManager()) class StoppedClient: async def claim(self): raise AssertionError("polling must not start") async def aclose(self): return None class FastHeartbeat: async def run(self, stop): await stop.wait() async def sync_once(self): return None class IdleProcessor: async def process(self, assignment): raise AssertionError("no assignment expected") def request_stop(self): return None first_app.client = StoppedClient() # type: ignore[assignment] first_app.heartbeat = FastHeartbeat() # type: ignore[assignment] first_app.processor = IdleProcessor() # type: ignore[assignment] stop = asyncio.Event() stop.set() asyncio.run(first_app.run_async(stop)) # Lock must be free now; a fresh create_application against the same # identity_path must succeed (simulates a clean restart). second_app = create_application(config=config, manager=DeviceManager()) asyncio.run(second_app.client.aclose())