from __future__ import annotations import asyncio from datetime import UTC, datetime, timedelta from threading import Event from cloud.internal_api.models import AssignmentModel, LeaseRenewalResponse from host_agent.assignment import AssignmentExecutionResult from host_agent.client import StaleLeaseError from host_agent.lease import ActiveAssignmentRunner def _assignment(*, expires_in: float = 0.0) -> AssignmentModel: return AssignmentModel( task_id="cloud-task", attempt=1, lease_id="lease-a", lease_expires_at=datetime.now(UTC) + timedelta(seconds=expires_in), host_id="host-a", device_id="device-a", goal="open settings", ) def test_lease_renews_while_execution_is_active() -> None: async def scenario() -> None: execution_started = Event() release_execution = Event() renewed = asyncio.Event() class BlockingExecutor: def execute(self, assignment, *, should_stop=None): execution_started.set() release_execution.wait(timeout=2) return AssignmentExecutionResult(status="done") class RenewingClient: async def renew(self, assignment): renewed.set() return LeaseRenewalResponse( status="renewed", lease_expires_at=datetime.now(UTC) + timedelta(seconds=30), ) run = asyncio.create_task( ActiveAssignmentRunner( RenewingClient(), # type: ignore[arg-type] BlockingExecutor(), ).run(_assignment()) ) assert await asyncio.to_thread(execution_started.wait, 1) await asyncio.wait_for(renewed.wait(), timeout=1) assert not run.done() release_execution.set() assert (await asyncio.wait_for(run, timeout=1)).status == "done" asyncio.run(scenario()) def test_stale_lease_stops_later_interruptible_actions() -> None: async def scenario() -> None: first_action_started = Event() actions: list[str] = [] class CooperativeExecutor: def execute(self, assignment, *, should_stop=None): assert should_stop is not None actions.append("first") first_action_started.set() assert first_action_started.wait(timeout=1) while not should_stop(): Event().wait(0.001) if not should_stop(): actions.append("second") return AssignmentExecutionResult( status="failed", failure_reason="execution interrupted", ) class StaleClient: async def renew(self, assignment): assert await asyncio.to_thread(first_action_started.wait, 1) raise StaleLeaseError(409, "stale lease") result = await asyncio.wait_for( ActiveAssignmentRunner( StaleClient(), # type: ignore[arg-type] CooperativeExecutor(), ).run(_assignment()), timeout=1, ) assert result.status == "failed" assert actions == ["first"] asyncio.run(scenario()) def test_renewal_loop_exits_when_execution_finishes() -> None: async def scenario() -> None: renew_calls = 0 class ImmediateExecutor: def execute(self, assignment, *, should_stop=None): return AssignmentExecutionResult(status="done") class CountingClient: async def renew(self, assignment): nonlocal renew_calls renew_calls += 1 return LeaseRenewalResponse( status="renewed", lease_expires_at=datetime.now(UTC) + timedelta(seconds=30), ) result = await asyncio.wait_for( ActiveAssignmentRunner( CountingClient(), # type: ignore[arg-type] ImmediateExecutor(), ).run(_assignment(expires_in=30)), timeout=1, ) assert result.status == "done" assert renew_calls == 0 asyncio.run(scenario()) def test_shutdown_request_stops_active_execution_cooperatively() -> None: async def scenario() -> None: execution_started = Event() class CooperativeExecutor: def execute(self, assignment, *, should_stop=None): assert should_stop is not None execution_started.set() while not should_stop(): Event().wait(0.001) return AssignmentExecutionResult( status="failed", failure_reason="execution interrupted", ) class RenewingClient: async def renew(self, assignment): return LeaseRenewalResponse( status="renewed", lease_expires_at=datetime.now(UTC) + timedelta(seconds=30), ) runner = ActiveAssignmentRunner( RenewingClient(), # type: ignore[arg-type] CooperativeExecutor(), ) running = asyncio.create_task(runner.run(_assignment(expires_in=30))) assert await asyncio.to_thread(execution_started.wait, 1) runner.request_stop() result = await asyncio.wait_for(running, timeout=1) assert result.failure_reason == "execution interrupted" asyncio.run(scenario())