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 __future__ import annotations
from collections.abc import Callable
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import Any from typing import Any
@@ -19,11 +20,21 @@ class AssignmentExecutor:
def __init__(self, factories: ExecutionFactories) -> None: def __init__(self, factories: ExecutionFactories) -> None:
self.factories = factories 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: 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: if assignment.goal is not None:
return self._execute_goal(assignment) return self._execute_goal(assignment, should_stop=should_stop)
return AssignmentExecutionResult( return AssignmentExecutionResult(
status="failed", status="failed",
failure_reason="assignment has neither goal nor workflow definition", failure_reason="assignment has neither goal nor workflow definition",
@@ -32,9 +43,15 @@ class AssignmentExecutor:
def _execute_goal( def _execute_goal(
self, self,
assignment: AssignmentModel, assignment: AssignmentModel,
*,
should_stop: Callable[[], bool] | None,
) -> AssignmentExecutionResult: ) -> AssignmentExecutionResult:
task = Task(goal=assignment.goal or "", device_id=assignment.device_id) 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( return AssignmentExecutionResult(
status="done" if completed.status == "completed" else "failed", status="done" if completed.status == "completed" else "failed",
failure_reason=completed.failure_reason, failure_reason=completed.failure_reason,
@@ -47,6 +64,8 @@ class AssignmentExecutor:
def _execute_workflow( def _execute_workflow(
self, self,
assignment: AssignmentModel, assignment: AssignmentModel,
*,
should_stop: Callable[[], bool] | None,
) -> AssignmentExecutionResult: ) -> AssignmentExecutionResult:
definition_id = assignment.workflow_definition_id or "" definition_id = assignment.workflow_definition_id or ""
definition = self.factories.workflow_store.get_definition(definition_id) definition = self.factories.workflow_store.get_definition(definition_id)
@@ -55,10 +74,15 @@ class AssignmentExecutor:
status="failed", status="failed",
failure_reason=f"unknown workflow definition {definition_id!r}", failure_reason=f"unknown workflow definition {definition_id!r}",
) )
run = self.factories.workflow_runner_factory().run( runner = self.factories.workflow_runner_factory()
definition, if should_stop is None:
device_id=assignment.device_id, 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( return AssignmentExecutionResult(
status="done" if run.status == "completed" else "failed", status="done" if run.status == "completed" else "failed",
failure_reason=( 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.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.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. - [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.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.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. - [ ] 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] Observer = Callable[[str], Scene]
ScreenshotProvider = Callable[[str], bytes] ScreenshotProvider = Callable[[str], bytes]
TaskSucceededHook = Callable[[str, str, Timeline], None] TaskSucceededHook = Callable[[str, str, Timeline], None]
StopRequested = Callable[[], bool]
class TaskRunner: class TaskRunner:
@@ -89,7 +90,12 @@ class TaskRunner:
else: else:
self.on_task_succeeded = None 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) context = TaskContext(task_id=task.id, goal=task.goal)
world_handle = self._start_world_view(task.id) world_handle = self._start_world_view(task.id)
if world_handle is not None: if world_handle is not None:
@@ -97,6 +103,8 @@ class TaskRunner:
self._update_task(task, status="running") self._update_task(task, status="running")
for _ in range(self.config.max_steps): for _ in range(self.config.max_steps):
if should_stop is not None and should_stop():
return self._interrupt_task(task)
try: try:
scene = self.observer(task.device_id) scene = self.observer(task.device_id)
context.add_scene(scene) context.add_scene(scene)
@@ -119,6 +127,8 @@ class TaskRunner:
return self._complete_task(task) return self._complete_task(task)
for step in steps: 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) executable_step = self._step_for_device(step, task.device_id)
result = self.executor.execute( result = self.executor.execute(
executable_step, executable_step,
@@ -143,6 +153,15 @@ class TaskRunner:
) )
return task 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: def _start_world_view(self, task_id: str) -> TaskWorldView | None:
"""Start (or resume) this task's isolated `WorldModel` handle, if enabled. """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 len(timeline.read(task.id)) == 2
assert metadata.get_task(task.id)["status"] == "completed" 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"] 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: def test_workflow_runner_failing_planned_goal_marks_run_failed(tmp_path) -> None:
definition = WorkflowDefinition( definition = WorkflowDefinition(
name="fail", name="fail",
+51 -8
View File
@@ -1,7 +1,6 @@
from __future__ import annotations from __future__ import annotations
from collections.abc import Callable from collections.abc import Callable
from dataclasses import replace
from datetime import datetime from datetime import datetime
from types import SimpleNamespace from types import SimpleNamespace
from typing import Any from typing import Any
@@ -33,6 +32,7 @@ TaskRunnerFactory = Callable[[], TaskRunner]
SceneProvider = Callable[[], Scene | None] SceneProvider = Callable[[], Scene | None]
WorldStateProvider = Callable[[], object | None] WorldStateProvider = Callable[[], object | None]
SleepFunc = Callable[[float], None] SleepFunc = Callable[[float], None]
StopRequested = Callable[[], bool]
TERMINAL_STATUSES = {"completed", "failed", "cancelled"} TERMINAL_STATUSES = {"completed", "failed", "cancelled"}
@@ -68,6 +68,8 @@ class WorkflowRunner:
definition: WorkflowDefinition, definition: WorkflowDefinition,
device_id: str, device_id: str,
initial_variables: dict[str, Any] | None = None, initial_variables: dict[str, Any] | None = None,
*,
should_stop: StopRequested | None = None,
) -> WorkflowRun: ) -> WorkflowRun:
self.store.save_definition(definition) self.store.save_definition(definition)
run = self.store.create_run( run = self.store.create_run(
@@ -75,9 +77,14 @@ class WorkflowRunner:
initial_variables or {}, initial_variables or {},
device_id=device_id, 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) run = self.store.get_run(run_id)
if run is None: if run is None:
raise KeyError(f"unknown workflow run {run_id}") raise KeyError(f"unknown workflow run {run_id}")
@@ -86,15 +93,19 @@ class WorkflowRunner:
definition = self.store.get_definition(run.definition_id) definition = self.store.get_definition(run.definition_id)
if definition is None: if definition is None:
raise KeyError(f"unknown workflow definition {run.definition_id}") 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( def _drive(
self, self,
definition: WorkflowDefinition, definition: WorkflowDefinition,
run: WorkflowRun, run: WorkflowRun,
*,
should_stop: StopRequested | None = None,
) -> WorkflowRun: ) -> WorkflowRun:
executed = 0 executed = 0
while run.status not in TERMINAL_STATUSES and run.current_step_id: 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: if self.step_limit is not None and executed >= self.step_limit:
return run return run
step = definition.step_by_id(run.current_step_id) step = definition.step_by_id(run.current_step_id)
@@ -109,7 +120,15 @@ class WorkflowRunner:
run = self._checkpoint(run, next_step_id, next_status) run = self._checkpoint(run, next_step_id, next_status)
continue 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( next_status, next_step_id = self._resolve_outcome(
definition, step, result.success, branch_next_step_id definition, step, result.success, branch_next_step_id
) )
@@ -164,13 +183,23 @@ class WorkflowRunner:
definition: WorkflowDefinition, definition: WorkflowDefinition,
run: WorkflowRun, run: WorkflowRun,
step: WorkflowStep, step: WorkflowStep,
*,
should_stop: StopRequested | None = None,
) -> tuple[WorkflowStepResult, str | None]: ) -> tuple[WorkflowStepResult, str | None]:
if isinstance(step, PlannedGoalStep): 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): if isinstance(step, SkillInvocationStep):
return self._execute_skill_invocation_step(step), None return self._execute_skill_invocation_step(step), None
if isinstance(step, WaitForConditionStep): 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): if isinstance(step, BranchStep):
return self._execute_branch_step(run, step) return self._execute_branch_step(run, step)
return ( return (
@@ -187,10 +216,15 @@ class WorkflowRunner:
self, self,
run: WorkflowRun, run: WorkflowRun,
step: PlannedGoalStep, step: PlannedGoalStep,
*,
should_stop: StopRequested | None = None,
) -> WorkflowStepResult: ) -> WorkflowStepResult:
task = Task(goal=step.goal, device_id=run.device_id or "") task = Task(goal=step.goal, device_id=run.device_id or "")
task_runner = self.task_runner_factory() 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" success = result_task.status == "completed"
return WorkflowStepResult( return WorkflowStepResult(
step_id=step.step_id, step_id=step.step_id,
@@ -238,11 +272,20 @@ class WorkflowRunner:
self, self,
run: WorkflowRun, run: WorkflowRun,
step: WaitForConditionStep, step: WaitForConditionStep,
*,
should_stop: StopRequested | None = None,
) -> WorkflowStepResult: ) -> WorkflowStepResult:
started_at = datetime.now().astimezone() started_at = datetime.now().astimezone()
timeout_seconds = step.timeout_seconds timeout_seconds = step.timeout_seconds
poll_interval_seconds = step.poll_interval_seconds poll_interval_seconds = step.poll_interval_seconds
while True: 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: try:
if self._condition_is_true(run, step.condition, started_at=started_at): if self._condition_is_true(run, step.condition, started_at=started_at):
return WorkflowStepResult( return WorkflowStepResult(