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.context import TaskContext from runtime.executor import Executor from runtime.planner import PlannedStep, Planner 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 = 20 Observer = Callable[[str], Scene] ScreenshotProvider = Callable[[str], bytes] TaskSucceededHook = Callable[[str, str, Timeline], 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, ) -> 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) ) 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 def run(self, task: 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): scene = self.observer(task.device_id) context.add_scene(scene) steps = self._plan(task.goal, scene, context) if not steps or self.planner.goal_reached( goal=task.goal, scene=scene, context=context, ): return self._complete_task(task) for step in steps: executable_step = self._step_for_device(step, task.device_id) result = self.executor.execute( executable_step, context=context, ) context.add_step_result(result) self._record_step_result(world_handle, context, 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 _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, ) -> 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) def _complete_task(self, task: Task) -> Task: 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 _plan( self, goal: str, scene: Scene, context: TaskContext, ) -> list[PlannedStep]: if context.world is not None and self._planner_accepts_world(): return self.planner.plan( goal=goal, scene=scene, context=context, world=context.world, ) return self.planner.plan(goal=goal, scene=scene, context=context) def _planner_accepts_world(self) -> bool: try: parameters = signature(self.planner.plan).parameters except (TypeError, ValueError): return True return "world" in parameters or any( parameter.kind is Parameter.VAR_KEYWORD for parameter in parameters.values() ) 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, ) -> 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", "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, )