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())
@@ -56,7 +56,7 @@
- [x] 7.2 Build complete device snapshots from the local `DeviceManager` and synchronize them at the configured interval.
- [x] 7.3 Compose local `TaskRunner` and `WorkflowRunner` factories without importing cloud concerns into Runtime-owned packages.
- [x] 7.4 Execute goal assignments through the configured Runtime Planner/Executor and workflow assignments through the existing workflow runner.
- [ ] 7.5 Run lease renewal alongside active execution and stop further interruptible actions after confirmed lease loss.
- [x] 7.5 Run lease renewal alongside active execution and stop further interruptible actions after confirmed lease loss.
- [ ] 7.6 Normalize and report successful/failed terminal outcomes, including Runtime failure reasons, with idempotent retries after response loss.
- [ ] 7.7 Implement graceful shutdown that stops polling, finishes or interrupts current work according to lease policy, and performs a final heartbeat when possible.
- [ ] 7.8 Add fake-driver end-to-end tests for one host, multiple hosts, NAT-style outbound-only operation, control-plane restart, Host Agent restart, and lease loss.
+20 -1
View File
@@ -39,6 +39,7 @@ class TaskRunnerConfig:
Observer = Callable[[str], Scene]
ScreenshotProvider = Callable[[str], bytes]
TaskSucceededHook = Callable[[str, str, Timeline], None]
StopRequested = Callable[[], bool]
class TaskRunner:
@@ -89,7 +90,12 @@ class TaskRunner:
else:
self.on_task_succeeded = None
def run(self, task: Task) -> Task:
def run(
self,
task: Task,
*,
should_stop: StopRequested | None = None,
) -> Task:
context = TaskContext(task_id=task.id, goal=task.goal)
world_handle = self._start_world_view(task.id)
if world_handle is not None:
@@ -97,6 +103,8 @@ class TaskRunner:
self._update_task(task, status="running")
for _ in range(self.config.max_steps):
if should_stop is not None and should_stop():
return self._interrupt_task(task)
try:
scene = self.observer(task.device_id)
context.add_scene(scene)
@@ -119,6 +127,8 @@ class TaskRunner:
return self._complete_task(task)
for step in steps:
if should_stop is not None and should_stop():
return self._interrupt_task(task)
executable_step = self._step_for_device(step, task.device_id)
result = self.executor.execute(
executable_step,
@@ -143,6 +153,15 @@ class TaskRunner:
)
return task
def _interrupt_task(self, task: Task) -> Task:
self._update_task(
task,
status="failed",
completed=True,
failure_reason="execution interrupted",
)
return task
def _start_world_view(self, task_id: str) -> TaskWorldView | None:
"""Start (or resume) this task's isolated `WorldModel` handle, if enabled.
+36
View File
@@ -76,3 +76,39 @@ def test_task_runner_executes_loop_and_writes_timeline(tmp_path) -> None:
assert len(timeline.read(task.id)) == 2
assert metadata.get_task(task.id)["status"] == "completed"
def test_task_runner_stops_before_the_next_planned_action() -> None:
scene = Scene(width=10, height=20, elements=[])
stop_requested = False
actions: list[str] = []
def record_action(**kwargs):
nonlocal stop_requested
actions.append("tap")
stop_requested = True
return {"ok": True}
runner = TaskRunner(
planner=ScriptedPlanner(
[
PlannedStep(action="tap", description="first", args={}),
PlannedStep(action="tap", description="second", args={}),
]
),
executor=Executor(
tools={"tap": record_action},
config=ExecutorConfig(max_retries=1, backoff_seconds=0),
),
config=TaskRunnerConfig(max_steps=5),
observer=lambda device_id: scene,
screenshot_provider=lambda device_id: PNG_10X20,
)
result = runner.run(
Task(goal="perform two actions", device_id="phone"),
should_stop=lambda: stop_requested,
)
assert result.status == "failed"
assert result.failure_reason == "execution interrupted"
assert actions == ["tap"]
+35
View File
@@ -96,6 +96,41 @@ def test_workflow_runner_linear_planned_goal_completes(tmp_path) -> None:
assert [result.step_id for result in run.step_results] == ["first", "second"]
def test_workflow_runner_stops_before_the_next_step(tmp_path) -> None:
stop_requested = False
calls: list[str] = []
class StoppingTaskRunner:
def run(self, task: Task, *, should_stop=None) -> Task:
nonlocal stop_requested
calls.append(task.goal)
stop_requested = True
task.status = "completed"
return task
definition = WorkflowDefinition(
name="interruptible",
entry_step_id="first",
steps=[
PlannedGoalStep("first", "first", next_step_id="second"),
PlannedGoalStep("second", "second"),
],
)
runner = WorkflowRunner(
_store(tmp_path),
task_runner_factory=lambda: StoppingTaskRunner(), # type: ignore[arg-type]
)
run = runner.run(
definition,
"phone",
should_stop=lambda: stop_requested,
)
assert run.status == "cancelled"
assert calls == ["first"]
def test_workflow_runner_failing_planned_goal_marks_run_failed(tmp_path) -> None:
definition = WorkflowDefinition(
name="fail",
+51 -8
View File
@@ -1,7 +1,6 @@
from __future__ import annotations
from collections.abc import Callable
from dataclasses import replace
from datetime import datetime
from types import SimpleNamespace
from typing import Any
@@ -33,6 +32,7 @@ TaskRunnerFactory = Callable[[], TaskRunner]
SceneProvider = Callable[[], Scene | None]
WorldStateProvider = Callable[[], object | None]
SleepFunc = Callable[[float], None]
StopRequested = Callable[[], bool]
TERMINAL_STATUSES = {"completed", "failed", "cancelled"}
@@ -68,6 +68,8 @@ class WorkflowRunner:
definition: WorkflowDefinition,
device_id: str,
initial_variables: dict[str, Any] | None = None,
*,
should_stop: StopRequested | None = None,
) -> WorkflowRun:
self.store.save_definition(definition)
run = self.store.create_run(
@@ -75,9 +77,14 @@ class WorkflowRunner:
initial_variables or {},
device_id=device_id,
)
return self._drive(definition, run)
return self._drive(definition, run, should_stop=should_stop)
def resume(self, run_id: str) -> WorkflowRun:
def resume(
self,
run_id: str,
*,
should_stop: StopRequested | None = None,
) -> WorkflowRun:
run = self.store.get_run(run_id)
if run is None:
raise KeyError(f"unknown workflow run {run_id}")
@@ -86,15 +93,19 @@ class WorkflowRunner:
definition = self.store.get_definition(run.definition_id)
if definition is None:
raise KeyError(f"unknown workflow definition {run.definition_id}")
return self._drive(definition, run)
return self._drive(definition, run, should_stop=should_stop)
def _drive(
self,
definition: WorkflowDefinition,
run: WorkflowRun,
*,
should_stop: StopRequested | None = None,
) -> WorkflowRun:
executed = 0
while run.status not in TERMINAL_STATUSES and run.current_step_id:
if should_stop is not None and should_stop():
return self._checkpoint(run, run.current_step_id, "cancelled")
if self.step_limit is not None and executed >= self.step_limit:
return run
step = definition.step_by_id(run.current_step_id)
@@ -109,7 +120,15 @@ class WorkflowRunner:
run = self._checkpoint(run, next_step_id, next_status)
continue
result, branch_next_step_id = self._execute_step(definition, run, step)
result, branch_next_step_id = self._execute_step(
definition,
run,
step,
should_stop=should_stop,
)
if should_stop is not None and should_stop():
self.store.append_step_result(run.id, result)
return self._checkpoint(run, run.current_step_id, "cancelled")
next_status, next_step_id = self._resolve_outcome(
definition, step, result.success, branch_next_step_id
)
@@ -164,13 +183,23 @@ class WorkflowRunner:
definition: WorkflowDefinition,
run: WorkflowRun,
step: WorkflowStep,
*,
should_stop: StopRequested | None = None,
) -> tuple[WorkflowStepResult, str | None]:
if isinstance(step, PlannedGoalStep):
return self._execute_planned_goal_step(run, step), None
return self._execute_planned_goal_step(
run,
step,
should_stop=should_stop,
), None
if isinstance(step, SkillInvocationStep):
return self._execute_skill_invocation_step(step), None
if isinstance(step, WaitForConditionStep):
return self._execute_wait_step(run, step), None
return self._execute_wait_step(
run,
step,
should_stop=should_stop,
), None
if isinstance(step, BranchStep):
return self._execute_branch_step(run, step)
return (
@@ -187,10 +216,15 @@ class WorkflowRunner:
self,
run: WorkflowRun,
step: PlannedGoalStep,
*,
should_stop: StopRequested | None = None,
) -> WorkflowStepResult:
task = Task(goal=step.goal, device_id=run.device_id or "")
task_runner = self.task_runner_factory()
result_task = task_runner.run(task)
if should_stop is None:
result_task = task_runner.run(task)
else:
result_task = task_runner.run(task, should_stop=should_stop)
success = result_task.status == "completed"
return WorkflowStepResult(
step_id=step.step_id,
@@ -238,11 +272,20 @@ class WorkflowRunner:
self,
run: WorkflowRun,
step: WaitForConditionStep,
*,
should_stop: StopRequested | None = None,
) -> WorkflowStepResult:
started_at = datetime.now().astimezone()
timeout_seconds = step.timeout_seconds
poll_interval_seconds = step.poll_interval_seconds
while True:
if should_stop is not None and should_stop():
return WorkflowStepResult(
step_id=step.step_id,
kind=step.kind,
success=False,
detail={"reason": "execution interrupted"},
)
try:
if self._condition_is_true(run, step.condition, started_at=started_at):
return WorkflowStepResult(