feat: checkpoint device agent runtime milestones
This commit is contained in:
+2
-1
@@ -7,6 +7,7 @@ from core.models import Scene
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from runtime.executor import StepResult
|
||||
from world.models import WorldState
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -15,6 +16,7 @@ class TaskContext:
|
||||
goal: str
|
||||
scenes: list[Scene] = field(default_factory=list)
|
||||
step_results: list["StepResult"] = field(default_factory=list)
|
||||
world: "WorldState | None" = None
|
||||
|
||||
def add_scene(self, scene: Scene) -> None:
|
||||
self.scenes.append(scene)
|
||||
@@ -27,4 +29,3 @@ class TaskContext:
|
||||
if not self.scenes:
|
||||
return None
|
||||
return self.scenes[-1]
|
||||
|
||||
|
||||
+3
-1
@@ -5,8 +5,8 @@ from dataclasses import dataclass, field
|
||||
from time import sleep
|
||||
from typing import Any
|
||||
|
||||
from core.device_manager import DeviceManager
|
||||
from core.errors import ElementNotFoundError
|
||||
from device.manager import DeviceManager
|
||||
from runtime.context import TaskContext
|
||||
from runtime.planner import PlannedStep
|
||||
|
||||
@@ -110,6 +110,7 @@ def default_tool_registry(
|
||||
manager: DeviceManager | None = None,
|
||||
) -> dict[str, ToolCallable]:
|
||||
from tools.describe_screen import describe_screen
|
||||
from tools.describe_screen_semantic import describe_screen_semantic
|
||||
from tools.find_icon import find_icon, find_icon_on_screen
|
||||
from tools.find_text import find_text, find_text_on_screen
|
||||
from tools.input_text import input_text
|
||||
@@ -130,6 +131,7 @@ def default_tool_registry(
|
||||
"get_ui_tree": _bind_manager(get_ui_tree, manager),
|
||||
"ui_tree": _bind_manager(get_ui_tree, manager),
|
||||
"describe_screen": _bind_manager(describe_screen, manager),
|
||||
"describe_screen_semantic": _bind_manager(describe_screen_semantic, manager),
|
||||
"find_text": find_text,
|
||||
"find_text_on_screen": _bind_manager(find_text_on_screen, manager),
|
||||
"find_icon": find_icon,
|
||||
|
||||
+5
-2
@@ -1,11 +1,14 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from core.models import Scene
|
||||
from runtime.context import TaskContext
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from world.models import WorldState
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class PlannedStep:
|
||||
@@ -22,6 +25,7 @@ class Planner:
|
||||
goal: str,
|
||||
scene: Scene,
|
||||
context: TaskContext,
|
||||
world: "WorldState | None" = None,
|
||||
) -> list[PlannedStep]:
|
||||
if context.step_results:
|
||||
return []
|
||||
@@ -37,4 +41,3 @@ class Planner:
|
||||
return bool(context.step_results) and all(
|
||||
result.success for result in context.step_results
|
||||
)
|
||||
|
||||
|
||||
+70
-2
@@ -2,15 +2,19 @@ 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
|
||||
@@ -33,6 +37,8 @@ class TaskRunner:
|
||||
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()
|
||||
@@ -43,15 +49,24 @@ class TaskRunner:
|
||||
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.planner.plan(goal=task.goal, scene=scene, context=context)
|
||||
steps = self._plan(task.goal, scene, context)
|
||||
if not steps or self.planner.goal_reached(
|
||||
goal=task.goal,
|
||||
scene=scene,
|
||||
@@ -61,11 +76,13 @@ class TaskRunner:
|
||||
return task
|
||||
|
||||
for step in steps:
|
||||
executable_step = self._step_for_device(step, task.device_id)
|
||||
result = self.executor.execute(
|
||||
self._step_for_device(step, task.device_id),
|
||||
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(
|
||||
@@ -84,6 +101,56 @@ class TaskRunner:
|
||||
)
|
||||
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,
|
||||
@@ -122,6 +189,7 @@ class TaskRunner:
|
||||
"get_ui_tree",
|
||||
"ui_tree",
|
||||
"describe_screen",
|
||||
"describe_screen_semantic",
|
||||
"find_text_on_screen",
|
||||
"find_icon_on_screen",
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user