diff --git a/apps/device-host-agent/host_agent/assignment.py b/apps/device-host-agent/host_agent/assignment.py index 6d5983e..3a3d687 100644 --- a/apps/device-host-agent/host_agent/assignment.py +++ b/apps/device-host-agent/host_agent/assignment.py @@ -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=( diff --git a/apps/device-host-agent/host_agent/lease.py b/apps/device-host-agent/host_agent/lease.py new file mode 100644 index 0000000..beab2b5 --- /dev/null +++ b/apps/device-host-agent/host_agent/lease.py @@ -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 diff --git a/apps/device-host-agent/tests/test_lease.py b/apps/device-host-agent/tests/test_lease.py new file mode 100644 index 0000000..d5f5560 --- /dev/null +++ b/apps/device-host-agent/tests/test_lease.py @@ -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()) diff --git a/openspec/changes/cloud-control-plane-integration/tasks.md b/openspec/changes/cloud-control-plane-integration/tasks.md index 8243036..e3993f2 100644 --- a/openspec/changes/cloud-control-plane-integration/tasks.md +++ b/openspec/changes/cloud-control-plane-integration/tasks.md @@ -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. diff --git a/runtime/task.py b/runtime/task.py index 2269c08..f31bc7f 100644 --- a/runtime/task.py +++ b/runtime/task.py @@ -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. diff --git a/tests/test_task_loop.py b/tests/test_task_loop.py index 756e14a..f04f31a 100644 --- a/tests/test_task_loop.py +++ b/tests/test_task_loop.py @@ -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"] diff --git a/tests/test_workflow_runner.py b/tests/test_workflow_runner.py index da2fccd..0c62306 100644 --- a/tests/test_workflow_runner.py +++ b/tests/test_workflow_runner.py @@ -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", diff --git a/workflow/runner.py b/workflow/runner.py index 146a085..ba04cd5 100644 --- a/workflow/runner.py +++ b/workflow/runner.py @@ -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(