177 lines
5.8 KiB
Python
177 lines
5.8 KiB
Python
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")
|
|
|
|
def latest_progress(self):
|
|
return None
|
|
|
|
class RenewingClient:
|
|
async def renew(self, assignment, *, progress=None):
|
|
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",
|
|
)
|
|
|
|
def latest_progress(self):
|
|
return None
|
|
|
|
class StaleClient:
|
|
async def renew(self, assignment, *, progress=None):
|
|
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")
|
|
|
|
def latest_progress(self):
|
|
return None
|
|
|
|
class CountingClient:
|
|
async def renew(self, assignment, *, progress=None):
|
|
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",
|
|
)
|
|
|
|
def latest_progress(self):
|
|
return None
|
|
|
|
class RenewingClient:
|
|
async def renew(self, assignment, *, progress=None):
|
|
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())
|