from __future__ import annotations from collections.abc import Callable from dataclasses import dataclass, replace from core.models import Scene, Task, utc_now from runtime.context import TaskContext from runtime.executor import Executor from runtime.planner import PlannedStep, Planner from storage.task_metadata import TaskMetadataStore from storage.timeline import Timeline from tools.describe_screen import describe_screen from tools.screenshot import take_screenshot @dataclass class TaskRunnerConfig: max_steps: int = 20 Observer = Callable[[str], Scene] ScreenshotProvider = Callable[[str], bytes] class TaskRunner: def __init__( self, *, planner: Planner | None = None, executor: Executor | None = None, timeline: Timeline | None = None, metadata_store: TaskMetadataStore | None = None, config: TaskRunnerConfig | None = None, observer: Observer | None = None, screenshot_provider: ScreenshotProvider | None = None, ) -> None: self.planner = planner or Planner() self.executor = executor or Executor() self.timeline = timeline self.metadata_store = metadata_store self.config = config or TaskRunnerConfig() self.observer = observer or (lambda device_id: describe_screen(device_id)) self.screenshot_provider = screenshot_provider or ( lambda device_id: take_screenshot(device_id) ) def run(self, task: Task) -> Task: context = TaskContext(task_id=task.id, goal=task.goal) self._update_task(task, status="running") for _ in range(self.config.max_steps): scene = self.observer(task.device_id) context.add_scene(scene) steps = self.planner.plan(goal=task.goal, scene=scene, context=context) if not steps or self.planner.goal_reached( goal=task.goal, scene=scene, context=context, ): self._update_task(task, status="completed", completed=True) return task for step in steps: result = self.executor.execute( self._step_for_device(step, task.device_id), context=context, ) context.add_step_result(result) self._append_timeline(task, scene, step, result) if not result.success: self._update_task( task, status="failed", completed=True, failure_reason=result.error or "step failed", ) return task self._update_task( task, status="failed", completed=True, failure_reason=f"max steps exceeded: {self.config.max_steps}", ) return task def _append_timeline( self, task: Task, scene: Scene, step: PlannedStep, result: object, ) -> None: if not self.timeline: return try: screenshot = self.screenshot_provider(task.device_id) except Exception: screenshot = None self.timeline.append( task_id=task.id, scene=scene.to_dict(), prompt=task.goal, tool_call={ "action": step.action, "description": step.description, "args": step.args, }, result=result.to_dict() if hasattr(result, "to_dict") else {"result": result}, screenshot=screenshot, ) def _step_for_device(self, step: PlannedStep, device_id: str) -> PlannedStep: device_scoped_actions = { "take_screenshot", "screenshot", "tap", "swipe", "input_text", "launch_app", "terminate_app", "get_ui_tree", "ui_tree", "describe_screen", "find_text_on_screen", "find_icon_on_screen", } if step.action not in device_scoped_actions or "device_id" in step.args: return step return replace(step, args={**step.args, "device_id": device_id}) def _update_task( self, task: Task, *, status: str, completed: bool = False, failure_reason: str | None = None, ) -> None: task.status = status # type: ignore[assignment] task.updated_at = utc_now() if completed: task.completed_at = utc_now() task.failure_reason = failure_reason if self.metadata_store: self.metadata_store.update_task( task.id, status=task.status, completed=completed, failure_reason=failure_reason, )