185 lines
6.7 KiB
Python
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})
|