from __future__ import annotations import logging from collections.abc import Callable from dataclasses import dataclass, replace from inspect import Parameter, signature from core.models import Scene, Task, utc_now from runtime.ai_planner import AIPlanner from runtime.context import TaskContext from runtime.executor import Executor from runtime.planner import PlannedStep, Planner from runtime.planner_config import PlannerConfig, load_config as load_planner_config from semantic.models import SemanticScene from skills_learning.config import ( SkillAuthoringConfig, load_config as load_skill_authoring_config, ) from skills_learning.embeddings import EmbeddingClient, embed_skill_text from skills_learning.models import skill_embedding_text from skills_learning.store import SkillStore, get_default_store from skills_learning.synthesis import synthesize_flow_skill from skills_learning.versioning import store_synthesized_skill from storage.task_metadata import TaskMetadataStore from storage.timeline import Timeline from tools.describe_screen import describe_screen from tools.screenshot import take_screenshot from world.config import WorldConfig, load_config as load_world_config from world.model import TaskWorldView, WorldModel logger = logging.getLogger(__name__) @dataclass class TaskRunnerConfig: max_steps: int = 999999 Observer = Callable[[str], Scene] ScreenshotProvider = Callable[[str], bytes] TaskSucceededHook = Callable[[str, str, Timeline], None] StopRequested = Callable[[], bool] StepProgressCallback = Callable[[int, str, str], None] 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, world_model: WorldModel | None = None, world_config: WorldConfig | None = None, on_task_succeeded: TaskSucceededHook | None = None, skill_authoring_config: SkillAuthoringConfig | None = None, skill_store: SkillStore | None = None, skill_embedding_client: EmbeddingClient | None = None, planner_config: PlannerConfig | None = None, on_step_progress: StepProgressCallback | None = None, ) -> None: self.planner_config = planner_config or load_planner_config() self.planner = planner or self._default_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) ) self.world_config = world_config or load_world_config() if world_model is not None: self.world_model = world_model elif self.world_config.enabled: self.world_model = WorldModel(config=self.world_config) else: self.world_model = None self.skill_authoring_config = ( skill_authoring_config or load_skill_authoring_config() ) self.skill_store = skill_store self.skill_embedding_client = skill_embedding_client if on_task_succeeded is not None: self.on_task_succeeded = on_task_succeeded elif self.skill_authoring_config.enabled: self.on_task_succeeded = self._default_task_succeeded_hook else: self.on_task_succeeded = None self.on_step_progress = on_step_progress def run( self, task: Task, *, should_stop: StopRequested | None = None, ) -> Task: if self.metadata_store: self.metadata_store.create_task(task) context = TaskContext(task_id=task.id, goal=task.goal) world_handle = self._start_world_view(task.id) if world_handle is not None: context.world = world_handle.state self._update_task(task, status="running") for _ in range(self.config.max_steps): if should_stop is not None and should_stop(): return self._interrupt_task(task) try: scene = self.observer(task.device_id) context.add_scene(scene) screenshot = self._planning_screenshot(task.device_id) steps = self._plan(task.goal, scene, context, screenshot=screenshot) except Exception as exc: reason = ( f"{type(exc).__name__}: {exc}" if str(exc) else type(exc).__name__ ) self._emit_step_progress(len(context.step_results), "failed", reason) self._update_task( task, status="failed", completed=True, failure_reason=reason, ) return task if not steps or self.planner.goal_reached( goal=task.goal, scene=scene, context=context, ): return self._complete_task(task) for step_index, step in enumerate(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) if step_index == 0 and screenshot is not None: # Reuse the screenshot already captured for planning instead of # taking a new one, so the "before action" image shown alongside # `scene`'s OCR/UI-tree overlay matches what was actually planned # against (avoids drift from the LLM planning round-trip). before_screenshot = screenshot else: before_screenshot = self._planning_screenshot(task.device_id) result = self.executor.execute( executable_step, context=context, ) after_screenshot = self._planning_screenshot(task.device_id) context.add_step_result(result) self._record_step_result( world_handle, context, task, scene, step, result, before_screenshot=before_screenshot, after_screenshot=after_screenshot, ) self._emit_step_progress( len(context.step_results), "running", f"{step.action}: {step.description}", ) if not result.success: self._emit_step_progress( len(context.step_results), "failed", result.error or "step failed", ) self._update_task( task, status="failed", completed=True, failure_reason=result.error or "step failed", ) return task self._emit_step_progress( len(context.step_results), "failed", f"max steps exceeded: {self.config.max_steps}", ) self._update_task( task, status="failed", completed=True, failure_reason=f"max steps exceeded: {self.config.max_steps}", ) return task def _interrupt_task(self, task: Task) -> Task: self._emit_step_progress(-1, "failed", "execution interrupted") self._update_task( task, status="failed", completed=True, failure_reason="execution interrupted", ) return task def _start_world_view(self, task_id: str) -> TaskWorldView | None: """Start (or resume) this task's isolated `WorldModel` handle, if enabled. Shared by `run()` and by any other driver (e.g. `CollaborativeTaskRunner`) that composes this `TaskRunner` for its world-model bookkeeping. """ if self.world_model is None: return None return self.world_model.start_task(task_id) def _record_step_result( self, world_handle: TaskWorldView | None, context: TaskContext, task: Task, scene: Scene, step: PlannedStep, result: object, *, before_screenshot: bytes | None = None, after_screenshot: bytes | None = None, ) -> None: """Record one executed step's bookkeeping: world-model observe + timeline append. Shared by `run()` and by any other driver (e.g. `CollaborativeTaskRunner`) that executes steps outside of this class's own loop, so the two paths cannot drift apart on what gets recorded for a step. """ self._update_world(world_handle, context, scene, step, result) self._append_timeline( task, scene, step, result, before_screenshot=before_screenshot, after_screenshot=after_screenshot, ) def _emit_step_progress( self, step_index: int, step_status: str, summary: str ) -> None: if self.on_step_progress is None: return try: self.on_step_progress(max(step_index, 0), step_status, summary[:200]) except Exception as exc: logger.debug("step progress callback failed: %s", exc) def _complete_task(self, task: Task) -> Task: self._emit_step_progress(-1, "completed", "task completed") self._update_task(task, status="completed", completed=True) self._notify_task_succeeded(task) return task def _notify_task_succeeded(self, task: Task) -> None: if self.on_task_succeeded is None or self.timeline is None: return try: self.on_task_succeeded(task.id, task.goal, self.timeline) except Exception as exc: logger.info("task succeeded hook failed: %s", exc) def _default_task_succeeded_hook( self, task_id: str, goal: str, timeline: Timeline, ) -> None: store = self.skill_store or get_default_store() candidate = synthesize_flow_skill( goal, timeline, task_id=task_id, store=store, ) stored = store_synthesized_skill(store, candidate).skill vector = embed_skill_text( skill_embedding_text(stored), client=self.skill_embedding_client, config=self.skill_authoring_config, ) if vector is not None: store.store_embedding( stored, vector, model_name=self.skill_authoring_config.embedding_model, ) def _default_planner(self) -> Planner: if self.planner_config.enabled: return AIPlanner(config=self.planner_config) return Planner() def _plan( self, goal: str, scene: Scene, context: TaskContext, screenshot: bytes | None = None, ) -> list[PlannedStep]: kwargs: dict[str, object] = {} if context.world is not None and self._planner_accepts("world"): kwargs["world"] = context.world if screenshot is not None and self._planner_accepts("screenshot"): kwargs["screenshot"] = screenshot return self.planner.plan(goal=goal, scene=scene, context=context, **kwargs) def _planner_accepts(self, name: str) -> bool: try: parameters = signature(self.planner.plan).parameters except (TypeError, ValueError): return True return name in parameters or any( parameter.kind is Parameter.VAR_KEYWORD for parameter in parameters.values() ) def _planning_screenshot(self, device_id: str) -> bytes | None: try: return self.screenshot_provider(device_id) except Exception: return None def _update_world( self, world_handle: TaskWorldView | None, context: TaskContext, scene: Scene, step: PlannedStep, result: object, ) -> None: if world_handle is None: return world_handle.observe( scene, self._semantic_scene_from_result(result), step, result, # type: ignore[arg-type] ) context.world = world_handle.state def _semantic_scene_from_result(self, result: object) -> SemanticScene | None: value = getattr(result, "result", None) if isinstance(value, dict): semantic_scene = value.get("semantic_scene") if isinstance(semantic_scene, SemanticScene): return semantic_scene return None def _append_timeline( self, task: Task, scene: Scene, step: PlannedStep, result: object, *, before_screenshot: bytes | None = None, after_screenshot: bytes | None = None, ) -> None: if not self.timeline: return ocr_results = scene.ocr_results_to_dict() if not ocr_results: ocr_results = [ element.to_dict() for element in scene.elements if element.source == "ocr" ] ui_tree_results = [ element.to_dict() for element in scene.elements if element.source == "ui" ] self.timeline.append( task_id=task.id, scene=scene.to_dict(), prompt=step.prompt or 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}, before_screenshot=before_screenshot, after_screenshot=after_screenshot, ocr_results=ocr_results, ui_tree_results=ui_tree_results, ) 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", "describe_screen_semantic", "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, )