98 lines
3.5 KiB
Python
98 lines
3.5 KiB
Python
from __future__ import annotations
|
|
|
|
from collections.abc import Callable
|
|
from dataclasses import dataclass, field
|
|
from typing import Any
|
|
|
|
from cloud.internal_api.models import AssignmentModel
|
|
from core.models import Task
|
|
from host_agent.execution import ExecutionFactories
|
|
from host_agent.planner_context import bind_planner_execution_context
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class AssignmentExecutionResult:
|
|
status: str
|
|
failure_reason: str | None = None
|
|
metadata: dict[str, Any] = field(default_factory=dict)
|
|
|
|
|
|
class AssignmentExecutor:
|
|
def __init__(self, factories: ExecutionFactories) -> None:
|
|
self.factories = factories
|
|
|
|
def execute(
|
|
self,
|
|
assignment: AssignmentModel,
|
|
*,
|
|
should_stop: Callable[[], bool] | None = None,
|
|
) -> AssignmentExecutionResult:
|
|
with bind_planner_execution_context(assignment):
|
|
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, should_stop=should_stop)
|
|
if assignment.goal is not None:
|
|
return self._execute_goal(assignment, should_stop=should_stop)
|
|
return AssignmentExecutionResult(
|
|
status="failed",
|
|
failure_reason="assignment has neither goal nor workflow definition",
|
|
)
|
|
|
|
def _execute_goal(
|
|
self,
|
|
assignment: AssignmentModel,
|
|
*,
|
|
should_stop: Callable[[], bool] | None,
|
|
) -> AssignmentExecutionResult:
|
|
task = Task(goal=assignment.goal or "", device_id=assignment.device_id)
|
|
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,
|
|
metadata={
|
|
"runtime_task_id": completed.id,
|
|
"runtime_status": completed.status,
|
|
},
|
|
)
|
|
|
|
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)
|
|
if definition is None:
|
|
return AssignmentExecutionResult(
|
|
status="failed",
|
|
failure_reason=f"unknown workflow definition {definition_id!r}",
|
|
)
|
|
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=(
|
|
None if run.status == "completed" else f"workflow ended as {run.status}"
|
|
),
|
|
metadata={
|
|
"workflow_run_id": run.id,
|
|
"workflow_status": run.status,
|
|
},
|
|
)
|