Files
agentic-mobile-control/agents/collab_runner.py
T
2026-07-06 23:23:44 +08:00

185 lines
6.7 KiB
Python

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})