feat(host-agent): renew active assignment leases

This commit is contained in:
2026-07-12 19:06:22 +08:00
parent 6dc2803ccc
commit ec1e8c20d8
8 changed files with 405 additions and 18 deletions
@@ -1,5 +1,6 @@
from __future__ import annotations
from collections.abc import Callable
from dataclasses import dataclass, field
from typing import Any
@@ -19,11 +20,21 @@ class AssignmentExecutor:
def __init__(self, factories: ExecutionFactories) -> None:
self.factories = factories
def execute(self, assignment: AssignmentModel) -> AssignmentExecutionResult:
def execute(
self,
assignment: AssignmentModel,
*,
should_stop: Callable[[], bool] | None = None,
) -> AssignmentExecutionResult:
if should_stop is not None and should_stop():
return AssignmentExecutionResult(
status="failed",
failure_reason="execution interrupted",
)
if assignment.workflow_definition_id is not None:
return self._execute_workflow(assignment)
return self._execute_workflow(assignment, should_stop=should_stop)
if assignment.goal is not None:
return self._execute_goal(assignment)
return self._execute_goal(assignment, should_stop=should_stop)
return AssignmentExecutionResult(
status="failed",
failure_reason="assignment has neither goal nor workflow definition",
@@ -32,9 +43,15 @@ class AssignmentExecutor:
def _execute_goal(
self,
assignment: AssignmentModel,
*,
should_stop: Callable[[], bool] | None,
) -> AssignmentExecutionResult:
task = Task(goal=assignment.goal or "", device_id=assignment.device_id)
completed = self.factories.task_runner_factory().run(task)
runner = self.factories.task_runner_factory()
if should_stop is None:
completed = runner.run(task)
else:
completed = runner.run(task, should_stop=should_stop)
return AssignmentExecutionResult(
status="done" if completed.status == "completed" else "failed",
failure_reason=completed.failure_reason,
@@ -47,6 +64,8 @@ class AssignmentExecutor:
def _execute_workflow(
self,
assignment: AssignmentModel,
*,
should_stop: Callable[[], bool] | None,
) -> AssignmentExecutionResult:
definition_id = assignment.workflow_definition_id or ""
definition = self.factories.workflow_store.get_definition(definition_id)
@@ -55,10 +74,15 @@ class AssignmentExecutor:
status="failed",
failure_reason=f"unknown workflow definition {definition_id!r}",
)
run = self.factories.workflow_runner_factory().run(
definition,
device_id=assignment.device_id,
)
runner = self.factories.workflow_runner_factory()
if should_stop is None:
run = runner.run(definition, device_id=assignment.device_id)
else:
run = runner.run(
definition,
device_id=assignment.device_id,
should_stop=should_stop,
)
return AssignmentExecutionResult(
status="done" if run.status == "completed" else "failed",
failure_reason=(
+102
View File
@@ -0,0 +1,102 @@
from __future__ import annotations
import asyncio
from collections.abc import Callable
from datetime import UTC, datetime
from threading import Event, Lock
from typing import Protocol
import httpx
from cloud.internal_api.models import AssignmentModel
from host_agent.assignment import AssignmentExecutionResult
from host_agent.client import HostAgentAPIError, HostAgentClient, StaleLeaseError
class InterruptibleAssignmentExecutor(Protocol):
def execute(
self,
assignment: AssignmentModel,
*,
should_stop: Callable[[], bool] | None = None,
) -> AssignmentExecutionResult: ...
class LeaseGuard:
def __init__(self) -> None:
self._lost = Event()
self._lock = Lock()
self._reason: str | None = None
@property
def reason(self) -> str | None:
with self._lock:
return self._reason
def is_lost(self) -> bool:
return self._lost.is_set()
def mark_lost(self, reason: str) -> None:
with self._lock:
if self._reason is None:
self._reason = reason
self._lost.set()
class ActiveAssignmentRunner:
def __init__(
self,
client: HostAgentClient,
executor: InterruptibleAssignmentExecutor,
*,
now: Callable[[], datetime] | None = None,
) -> None:
self.client = client
self.executor = executor
self._now = now or (lambda: datetime.now(UTC))
async def run(self, assignment: AssignmentModel) -> AssignmentExecutionResult:
guard = LeaseGuard()
execution = asyncio.create_task(
asyncio.to_thread(
self.executor.execute,
assignment,
should_stop=guard.is_lost,
)
)
renewal = asyncio.create_task(
self._renew_while_running(assignment, execution, guard)
)
try:
return await execution
finally:
await renewal
async def _renew_while_running(
self,
assignment: AssignmentModel,
execution: asyncio.Task[AssignmentExecutionResult],
guard: LeaseGuard,
) -> None:
lease_expires_at = assignment.lease_expires_at
while not execution.done():
delay = max(
0.0,
(lease_expires_at - self._now()).total_seconds() / 3,
)
done, _ = await asyncio.wait({execution}, timeout=delay)
if done:
return
try:
response = await self.client.renew(assignment)
except StaleLeaseError:
guard.mark_lost("lease rejected by control plane")
return
except HostAgentAPIError as exc:
guard.mark_lost(f"lease renewal rejected: {exc.detail}")
return
except httpx.TransportError:
guard.mark_lost("lease renewal failed after transport retries")
return
else:
lease_expires_at = response.lease_expires_at
+128
View File
@@ -0,0 +1,128 @@
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())