from __future__ import annotations import logging from dataclasses import dataclass, replace from agents.config import CollaborationConfig, load_config from agents.models import Observation, VerificationVerdict from agents.observer import Observer from agents.reflector import Reflector from agents.verifier import Verifier 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 runtime.task import TaskRunner, TaskRunnerConfig from semantic.llm_client import AnthropicSemanticClient logger = logging.getLogger(__name__) @dataclass class CollaborativeTaskRunnerConfig: max_steps: int = 20 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.task_runner = task_runner 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() 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: runner = self.task_runner or TaskRunner( planner=self.planner, executor=self.executor, config=TaskRunnerConfig(max_steps=self.config.max_steps), ) return runner.run(task) def _run_collaborative(self, task: Task) -> Task: context = TaskContext(task_id=task.id, goal=task.goal) task.status = "running" # type: ignore[assignment] task.updated_at = utc_now() 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, ): task.status = "completed" # type: ignore[assignment] task.updated_at = utc_now() task.completed_at = utc_now() return 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) 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, ) if verdict.result == "achieved": continue if recovery_attempts >= self.collaboration_config.max_recovery_attempts: task.status = "failed" # type: ignore[assignment] task.updated_at = utc_now() task.completed_at = utc_now() task.failure_reason = ( f"Reflection recovery ceiling reached " f"({self.collaboration_config.max_recovery_attempts} attempts)" ) return task 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, ) recovery_result = self.executor.execute( self._step_for_device(recovery_step, task.device_id), context=context, ) context.add_step_result(recovery_result) if not recovery_result.success: task.status = "failed" # type: ignore[assignment] task.updated_at = utc_now() task.completed_at = utc_now() task.failure_reason = ( f"Recovery action failed: {recovery_result.error}" ) return task task.status = "failed" # type: ignore[assignment] task.updated_at = utc_now() task.completed_at = utc_now() task.failure_reason = f"max steps exceeded: {self.config.max_steps}" 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})