from __future__ import annotations import logging from dataclasses import dataclass, replace from agents.config import CollaborationConfig, load_config from agents.observer import Observer from agents.reflector import Reflector from agents.verifier import Verifier from core.models import Scene, Task from runtime.context import TaskContext from runtime.executor import Executor from runtime.planner import PlannedStep, Planner from runtime.task import TaskRunner, TaskRunnerConfig from semantic.llm_client import AnthropicSemanticClient from world.model import TaskWorldView logger = logging.getLogger(__name__) @dataclass class CollaborativeTaskRunnerConfig: max_steps: int = 999999 max_recovery_attempts: int = 3 class CollaborativeTaskRunner: def __init__( self, *, planner: Planner | None = None, executor: Executor | None = None, task_runner: TaskRunner | None = None, observer: Observer | None = None, verifier: Verifier | None = None, reflector: Reflector | None = None, config: CollaborativeTaskRunnerConfig | None = None, collaboration_config: CollaborationConfig | None = None, llm_client: AnthropicSemanticClient | None = None, ) -> None: self.planner = planner or Planner() self.executor = executor or Executor() self.observer = observer or Observer() self.verifier = verifier or Verifier(client=llm_client) self.reflector = reflector or Reflector(client=llm_client) self.config = config or CollaborativeTaskRunnerConfig() self.collaboration_config = collaboration_config or load_config() # Shared with `_run_plain`: also holds the Timeline/TaskMetadataStore/ # WorldModel/on_task_succeeded wiring that `_run_collaborative` reuses # for its own per-step bookkeeping, so the two paths cannot drift apart. self.task_runner = task_runner or TaskRunner( planner=self.planner, executor=self.executor, config=TaskRunnerConfig(max_steps=self.config.max_steps), ) def run(self, task: Task) -> Task: if not self.collaboration_config.enabled: return self._run_plain(task) return self._run_collaborative(task) def _run_plain(self, task: Task) -> Task: return self.task_runner.run(task) def _run_collaborative(self, task: Task) -> Task: context = TaskContext(task_id=task.id, goal=task.goal) world_handle: TaskWorldView | None = self.task_runner._start_world_view(task.id) if world_handle is not None: context.world = world_handle.state self.task_runner._update_task(task, status="running") recovery_attempts = 0 for _ in range(self.config.max_steps): scene = self._observe_scene(task.device_id) context.add_scene(scene) pre_observation = self.observer.observe( scene=scene, world=context.world, ) 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, ): return self.task_runner._complete_task(task) for step in steps: executable_step = self._step_for_device(step, task.device_id) before_screenshot = self.task_runner._planning_screenshot( task.device_id ) result = self.executor.execute(executable_step, context=context) after_screenshot = self.task_runner._planning_screenshot(task.device_id) context.add_step_result(result) self.task_runner._record_step_result( world_handle, context, task, scene, step, result, before_screenshot=before_screenshot, after_screenshot=after_screenshot, ) post_scene = self._observe_scene(task.device_id) post_observation = self.observer.observe( scene=post_scene, world=context.world, ) verdict = self.verifier.verify( pre_observation=pre_observation, post_observation=post_observation, planned_step=step, step_result=result, ) # Refresh pre_observation for the *next* step in this plan (or # the next outer-loop iteration) so it always reflects the # most recent known device state, never the stale # pre-whole-plan observation. pre_observation = post_observation if verdict.result == "achieved": continue if recovery_attempts >= self.collaboration_config.max_recovery_attempts: return self._fail( task, f"Reflection recovery ceiling reached " f"({self.collaboration_config.max_recovery_attempts} attempts)", ) outcome = self.reflector.reflect( observation=post_observation, planned_step=step, step_result=result, verdict=verdict, ) recovery_attempts += 1 if outcome.replan: break if outcome.action is not None: recovery_step = PlannedStep( action=outcome.action.action, description=outcome.action.description, args=outcome.action.args, ) before_screenshot = self.task_runner._planning_screenshot( task.device_id ) recovery_result = self.executor.execute( self._step_for_device(recovery_step, task.device_id), context=context, ) after_screenshot = self.task_runner._planning_screenshot( task.device_id ) context.add_step_result(recovery_result) self.task_runner._record_step_result( world_handle, context, task, post_scene, recovery_step, recovery_result, before_screenshot=before_screenshot, after_screenshot=after_screenshot, ) if not recovery_result.success: return self._fail( task, f"Recovery action failed: {recovery_result.error}" ) recovery_scene = self._observe_scene(task.device_id) pre_observation = self.observer.observe( scene=recovery_scene, world=context.world, ) return self._fail(task, f"max steps exceeded: {self.config.max_steps}") def _fail(self, task: Task, reason: str) -> Task: self.task_runner._update_task( task, status="failed", completed=True, failure_reason=reason, ) return task def _observe_scene(self, device_id: str) -> Scene: from tools.describe_screen import describe_screen return describe_screen(device_id) 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})