feat: checkpoint device agent runtime milestones
This commit is contained in:
+148
@@ -0,0 +1,148 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from uuid import uuid4
|
||||
|
||||
from device.manager import DeviceManager
|
||||
from driver.registry import build_driver_factory
|
||||
from pydantic import BaseModel, Field
|
||||
from runtime.task import TaskRunner
|
||||
from storage.device_config import DeviceConfigStore
|
||||
from storage.task_metadata import TaskMetadataStore
|
||||
from storage.timeline import Timeline
|
||||
|
||||
|
||||
class RegisterDeviceRequest(BaseModel):
|
||||
driver_type: str
|
||||
connection_info: dict[str, Any] = Field(default_factory=dict)
|
||||
name: str | None = None
|
||||
|
||||
|
||||
class RuntimeConfigRequest(BaseModel):
|
||||
max_steps: int
|
||||
|
||||
|
||||
def create_console_router(
|
||||
*,
|
||||
device_manager: DeviceManager,
|
||||
metadata_store: TaskMetadataStore,
|
||||
timeline: Timeline,
|
||||
config_store: DeviceConfigStore,
|
||||
task_runner: TaskRunner,
|
||||
) -> Any:
|
||||
from fastapi import APIRouter, HTTPException, Response, status
|
||||
|
||||
router = APIRouter(prefix="/console", tags=["console"])
|
||||
|
||||
@router.get("/devices")
|
||||
def devices() -> list[dict[str, Any]]:
|
||||
return [device.to_dict() for device in device_manager.list_devices()]
|
||||
|
||||
@router.post("/devices", status_code=status.HTTP_201_CREATED)
|
||||
def register_device(request: RegisterDeviceRequest) -> dict[str, Any]:
|
||||
try:
|
||||
driver_factory = build_driver_factory(
|
||||
request.driver_type,
|
||||
request.connection_info,
|
||||
)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
|
||||
device_id = uuid4().hex
|
||||
device = device_manager.register_device(
|
||||
device_id,
|
||||
driver_factory,
|
||||
name=request.name,
|
||||
driver_type=request.driver_type,
|
||||
connection_info=request.connection_info,
|
||||
)
|
||||
config_store.add(
|
||||
device_id=device_id,
|
||||
name=request.name,
|
||||
driver_type=request.driver_type,
|
||||
connection_info=request.connection_info,
|
||||
)
|
||||
return device.to_dict()
|
||||
|
||||
@router.delete(
|
||||
"/devices/{device_id}",
|
||||
status_code=status.HTTP_204_NO_CONTENT,
|
||||
response_model=None,
|
||||
)
|
||||
def unregister_device(device_id: str) -> Response:
|
||||
if not _has_device(device_manager, device_id):
|
||||
raise HTTPException(status_code=404, detail="device not found")
|
||||
device_manager.unregister_device(device_id)
|
||||
config_store.remove(device_id)
|
||||
return Response(status_code=status.HTTP_204_NO_CONTENT)
|
||||
|
||||
@router.get("/tasks")
|
||||
def tasks(
|
||||
device_id: str | None = None,
|
||||
status: str | None = None,
|
||||
) -> list[dict[str, Any]]:
|
||||
rows = metadata_store.list_tasks()
|
||||
if device_id is not None:
|
||||
rows = [row for row in rows if row["device_id"] == device_id]
|
||||
if status is not None:
|
||||
rows = [row for row in rows if row["status"] == status]
|
||||
return rows
|
||||
|
||||
@router.get("/tasks/{task_id}")
|
||||
def task_detail(task_id: str) -> dict[str, Any]:
|
||||
task = metadata_store.get_task(task_id)
|
||||
if task is None:
|
||||
raise HTTPException(status_code=404, detail="task not found")
|
||||
return task
|
||||
|
||||
@router.get("/tasks/{task_id}/timeline")
|
||||
def task_timeline(task_id: str) -> list[dict[str, Any]]:
|
||||
if metadata_store.get_task(task_id) is None:
|
||||
raise HTTPException(status_code=404, detail="task not found")
|
||||
return [_inline_screenshot(record) for record in timeline.read(task_id)]
|
||||
|
||||
@router.get("/config")
|
||||
def runtime_config() -> dict[str, int]:
|
||||
return {"max_steps": _runner_max_steps(task_runner)}
|
||||
|
||||
@router.put("/config")
|
||||
def update_runtime_config(request: RuntimeConfigRequest) -> dict[str, int]:
|
||||
if request.max_steps <= 0:
|
||||
raise HTTPException(status_code=400, detail="max_steps must be positive")
|
||||
_set_runner_max_steps(task_runner, request.max_steps)
|
||||
config_store.set_setting("max_steps", request.max_steps)
|
||||
return {"max_steps": request.max_steps}
|
||||
|
||||
return router
|
||||
|
||||
|
||||
def _inline_screenshot(record: dict[str, Any]) -> dict[str, Any]:
|
||||
payload = dict(record)
|
||||
screenshot_path = payload.get("screenshot_path")
|
||||
if screenshot_path:
|
||||
path = Path(str(screenshot_path))
|
||||
if path.exists():
|
||||
payload["image_base64"] = base64.b64encode(path.read_bytes()).decode(
|
||||
"ascii"
|
||||
)
|
||||
return payload
|
||||
|
||||
|
||||
def _has_device(device_manager: DeviceManager, device_id: str) -> bool:
|
||||
return any(device.id == device_id for device in device_manager.list_devices())
|
||||
|
||||
|
||||
def _runner_max_steps(task_runner: TaskRunner) -> int:
|
||||
config = getattr(task_runner, "config", None)
|
||||
if config is None or not hasattr(config, "max_steps"):
|
||||
raise RuntimeError("task runner config unavailable")
|
||||
return int(config.max_steps)
|
||||
|
||||
|
||||
def _set_runner_max_steps(task_runner: TaskRunner, max_steps: int) -> None:
|
||||
config = getattr(task_runner, "config", None)
|
||||
if config is None or not hasattr(config, "max_steps"):
|
||||
raise RuntimeError("task runner config unavailable")
|
||||
config.max_steps = max_steps
|
||||
+1
-1
@@ -3,7 +3,7 @@ from collections.abc import Callable
|
||||
from typing import Any
|
||||
|
||||
from api.errors import call_with_semantic_errors
|
||||
from core.device_manager import DEFAULT_MANAGER, DeviceManager
|
||||
from device.manager import DEFAULT_MANAGER, DeviceManager
|
||||
from tools.describe_screen import describe_screen
|
||||
from tools.find_icon import find_icon_on_screen
|
||||
from tools.find_text import find_text_on_screen
|
||||
|
||||
+63
-2
@@ -1,11 +1,15 @@
|
||||
from typing import Any
|
||||
|
||||
from api.console import create_console_router
|
||||
from api.errors import semantic_error
|
||||
from core.device_manager import DEFAULT_MANAGER, DeviceManager
|
||||
from core.models import Task
|
||||
from device.manager import DEFAULT_MANAGER, DeviceManager
|
||||
from driver.registry import build_driver_factory
|
||||
from runtime.executor import Executor, default_tool_registry
|
||||
from runtime.task import TaskRunner
|
||||
from runtime.task import TaskRunner, TaskRunnerConfig
|
||||
from storage.device_config import DEFAULT_MAX_STEPS, DeviceConfigStore
|
||||
from storage.task_metadata import TaskMetadataStore
|
||||
from storage.timeline import Timeline
|
||||
from tools.launch_app import launch_app
|
||||
from tools.screenshot import take_screenshot
|
||||
from tools.tap import tap
|
||||
@@ -16,17 +20,33 @@ def create_app(
|
||||
manager: DeviceManager | None = None,
|
||||
task_runner: TaskRunner | None = None,
|
||||
metadata_store: TaskMetadataStore | None = None,
|
||||
device_config_store: DeviceConfigStore | None = None,
|
||||
timeline: Timeline | None = None,
|
||||
) -> Any:
|
||||
from fastapi import BackgroundTasks, FastAPI, HTTPException
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
from pydantic import BaseModel
|
||||
|
||||
device_manager = manager or DEFAULT_MANAGER
|
||||
store = metadata_store or TaskMetadataStore()
|
||||
config_store = device_config_store or DeviceConfigStore()
|
||||
timeline_store = timeline or Timeline()
|
||||
max_steps = _load_max_steps(config_store)
|
||||
_reload_device_configs(device_manager, config_store)
|
||||
runner = task_runner or TaskRunner(
|
||||
metadata_store=store,
|
||||
executor=Executor(tools=default_tool_registry(manager=device_manager)),
|
||||
timeline=timeline_store,
|
||||
config=TaskRunnerConfig(max_steps=max_steps),
|
||||
)
|
||||
_apply_max_steps(runner, max_steps)
|
||||
app = FastAPI(title="Apex Agent API")
|
||||
app.add_middleware(
|
||||
CORSMiddleware,
|
||||
allow_origins=["*"],
|
||||
allow_methods=["*"],
|
||||
allow_headers=["*"],
|
||||
)
|
||||
|
||||
class TapRequest(BaseModel):
|
||||
x: float
|
||||
@@ -97,6 +117,16 @@ def create_app(
|
||||
raise HTTPException(status_code=404, detail="task not found")
|
||||
return task
|
||||
|
||||
app.include_router(
|
||||
create_console_router(
|
||||
device_manager=device_manager,
|
||||
metadata_store=store,
|
||||
timeline=timeline_store,
|
||||
config_store=config_store,
|
||||
task_runner=runner,
|
||||
)
|
||||
)
|
||||
|
||||
return app
|
||||
|
||||
|
||||
@@ -107,3 +137,34 @@ def _raise_semantic(func: Any) -> Any:
|
||||
return func()
|
||||
except Exception as exc:
|
||||
raise HTTPException(status_code=400, detail=semantic_error(exc)) from exc
|
||||
|
||||
|
||||
def _load_max_steps(config_store: DeviceConfigStore) -> int:
|
||||
raw_value = config_store.get_setting("max_steps")
|
||||
try:
|
||||
max_steps = int(raw_value) if raw_value is not None else DEFAULT_MAX_STEPS
|
||||
except ValueError:
|
||||
return DEFAULT_MAX_STEPS
|
||||
if max_steps <= 0:
|
||||
return DEFAULT_MAX_STEPS
|
||||
return max_steps
|
||||
|
||||
|
||||
def _apply_max_steps(task_runner: Any, max_steps: int) -> None:
|
||||
config = getattr(task_runner, "config", None)
|
||||
if config is not None and hasattr(config, "max_steps"):
|
||||
config.max_steps = max_steps
|
||||
|
||||
|
||||
def _reload_device_configs(
|
||||
device_manager: DeviceManager,
|
||||
config_store: DeviceConfigStore,
|
||||
) -> None:
|
||||
for config in config_store.list():
|
||||
device_manager.register_device(
|
||||
config["device_id"],
|
||||
build_driver_factory(config["driver_type"], config["connection_info"]),
|
||||
name=config["name"],
|
||||
driver_type=config["driver_type"],
|
||||
connection_info=config["connection_info"],
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user