from __future__ import annotations from collections.abc import Callable from dataclasses import dataclass, field from typing import TYPE_CHECKING, 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 from runtime.task import is_cancellation_reason if TYPE_CHECKING: from host_agent.mcp_lock import McpBusyTracker @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, *, mcp_busy_tracker: McpBusyTracker | None = None, ) -> None: self.factories = factories self._progress = TaskProgressHolder() self._mcp_busy_tracker = mcp_busy_tracker 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, stop_reason: Callable[[], str | None] | None = None, ) -> AssignmentExecutionResult: self._progress.clear() if self._mcp_busy_tracker is not None and ( assignment.device_id in self._mcp_busy_tracker.busy_device_ids() ): return AssignmentExecutionResult( status="failed", failure_reason=( f"device {assignment.device_id} is held by an active " "MCP session" ), ) with bind_planner_execution_context(assignment): if should_stop is not None and should_stop(): reason = stop_reason() if stop_reason is not None else None return AssignmentExecutionResult( status="cancelled" if is_cancellation_reason(reason) else "failed", failure_reason=reason or "execution interrupted", ) if assignment.workflow_definition_id is not None: return self._execute_workflow( assignment, should_stop=should_stop, stop_reason=stop_reason ) if assignment.goal is not None: return self._execute_goal( assignment, should_stop=should_stop, stop_reason=stop_reason ) return AssignmentExecutionResult( status="failed", failure_reason="assignment has neither goal nor workflow definition", ) def _execute_goal( self, assignment: AssignmentModel, *, should_stop: Callable[[], bool] | None, stop_reason: Callable[[], str | None] | None, ) -> AssignmentExecutionResult: task = Task(goal=assignment.goal or "", device_id=assignment.device_id) if self.factories.metadata_store is not None: self.factories.metadata_store.create_task( task, source_task_id=assignment.task_id, source_attempt=assignment.attempt, ) 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, stop_reason=stop_reason ) return AssignmentExecutionResult( status=_terminal_status(completed.status), 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, stop_reason: Callable[[], str | None] | 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, stop_reason=stop_reason, ) return AssignmentExecutionResult( status=_terminal_status(run.status), failure_reason=( None if run.status == "completed" else f"workflow ended as {run.status}" ), metadata={ "workflow_run_id": run.id, "workflow_status": run.status, }, ) def _terminal_status(runtime_status: str) -> str: if runtime_status == "completed": return "done" if runtime_status == "cancelled": return "cancelled" return "failed"