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 from host_agent.progress import TaskProgressHolder, TaskProgressSnapshot @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 self._progress = TaskProgressHolder() def latest_progress(self) -> TaskProgressSnapshot | None: """Latest step progress reported by the currently-running assignment.""" return self._progress.snapshot() def execute( self, assignment: AssignmentModel, *, should_stop: Callable[[], bool] | None = None, ) -> AssignmentExecutionResult: self._progress.clear() 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() runner.on_step_progress = self._progress.update 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, }, )