feat(host-agent): renew active assignment leases
This commit is contained in:
@@ -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=(
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -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
@@ -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.
|
||||||
|
|
||||||
|
|||||||
@@ -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"]
|
||||||
|
|||||||
@@ -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
@@ -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(
|
||||||
|
|||||||
Reference in New Issue
Block a user