220 lines
7.2 KiB
Python
220 lines
7.2 KiB
Python
from __future__ import annotations
|
|
|
|
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 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 WorldModel
|
|
|
|
|
|
@dataclass
|
|
class TaskRunnerConfig:
|
|
max_steps: int = 20
|
|
|
|
|
|
Observer = Callable[[str], Scene]
|
|
ScreenshotProvider = Callable[[str], bytes]
|
|
|
|
|
|
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,
|
|
) -> 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
|
|
|
|
def run(self, task: Task) -> Task:
|
|
context = TaskContext(task_id=task.id, goal=task.goal)
|
|
if self.world_model is not None:
|
|
context.world = self.world_model.start_task(task.id)
|
|
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,
|
|
):
|
|
self._update_task(task, status="completed", completed=True)
|
|
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)
|
|
self._update_world(context, scene, step, result)
|
|
self._append_timeline(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 _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,
|
|
context: TaskContext,
|
|
scene: Scene,
|
|
step: PlannedStep,
|
|
result: object,
|
|
) -> None:
|
|
if self.world_model is None:
|
|
return
|
|
self.world_model.observe(
|
|
scene,
|
|
self._semantic_scene_from_result(result),
|
|
step,
|
|
result, # type: ignore[arg-type]
|
|
)
|
|
context.world = self.world_model.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,
|
|
)
|