121 lines
4.3 KiB
Python
121 lines
4.3 KiB
Python
from __future__ import annotations
|
|
|
|
import os
|
|
from collections.abc import Callable
|
|
from dataclasses import dataclass, replace
|
|
|
|
from device.manager import DeviceManager
|
|
from host_agent.cloud_planner_client import CloudProxyToolCallingClient
|
|
from host_agent.config import HostAgentConfig, load_host_agent_config
|
|
from runtime.ai_planner import AIPlanner
|
|
from runtime.executor import Executor, default_tool_registry
|
|
from runtime.planner import Planner
|
|
from runtime.planner_config import PlannerConfig, load_config as load_planner_config
|
|
from runtime.task import TaskRunner
|
|
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 workflow.runner import WorkflowRunner
|
|
from workflow.store import WorkflowStore
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class ExecutionFactories:
|
|
task_runner_factory: Callable[[], TaskRunner]
|
|
workflow_runner_factory: Callable[[], WorkflowRunner]
|
|
workflow_store: WorkflowStore
|
|
metadata_store: TaskMetadataStore | None = None
|
|
|
|
|
|
def create_execution_factories(
|
|
manager: DeviceManager,
|
|
*,
|
|
workflow_store: WorkflowStore | None = None,
|
|
metadata_store: TaskMetadataStore | None = None,
|
|
timeline: Timeline | None = None,
|
|
host_agent_config: HostAgentConfig | None = None,
|
|
) -> ExecutionFactories:
|
|
shared_workflow_store = workflow_store or WorkflowStore()
|
|
resolved_host_agent_config = host_agent_config
|
|
|
|
def create_task_runner() -> TaskRunner:
|
|
return TaskRunner(
|
|
executor=Executor(tools=default_tool_registry(manager=manager)),
|
|
observer=lambda device_id: describe_screen(device_id, manager=manager),
|
|
screenshot_provider=lambda device_id: take_screenshot(
|
|
device_id, manager=manager
|
|
),
|
|
metadata_store=metadata_store,
|
|
timeline=timeline,
|
|
planner=_host_agent_planner(resolved_host_agent_config),
|
|
planner_config=_host_agent_planner_config(),
|
|
device_platform_provider=lambda device_id: _device_platform(
|
|
manager, device_id
|
|
),
|
|
)
|
|
|
|
def create_workflow_runner() -> WorkflowRunner:
|
|
return WorkflowRunner(
|
|
shared_workflow_store,
|
|
task_runner_factory=create_task_runner,
|
|
)
|
|
|
|
return ExecutionFactories(
|
|
task_runner_factory=create_task_runner,
|
|
workflow_runner_factory=create_workflow_runner,
|
|
workflow_store=shared_workflow_store,
|
|
metadata_store=metadata_store,
|
|
)
|
|
|
|
|
|
def _host_agent_planner_config() -> PlannerConfig:
|
|
"""Host Agent defaults to the AI planner unless an operator opts out.
|
|
|
|
`runtime.planner_config` defaults `enabled=False` for the shared Runtime
|
|
library (local dev/tests/cloud dispatcher keep the deterministic stub
|
|
planner unless asked). The Host Agent is the actual device-control path,
|
|
so it flips that default on here -- an explicit `AI_PLANNER_ENABLED=false`
|
|
still disables it.
|
|
"""
|
|
config = load_planner_config()
|
|
if os.environ.get("AI_PLANNER_ENABLED") is None:
|
|
config = replace(config, enabled=True)
|
|
return config
|
|
|
|
|
|
def _host_agent_planner(
|
|
host_agent_config: HostAgentConfig | None,
|
|
) -> Planner | None:
|
|
"""Build the `AIPlanner` explicitly when the cloud-proxy transport is
|
|
selected, so its `ToolCallingClient` is a `CloudProxyToolCallingClient`
|
|
instead of a local Anthropic/OpenAI SDK client.
|
|
|
|
Returns `None` (letting `TaskRunner` fall back to its own
|
|
`_default_planner()`) only for the explicit `direct` transport.
|
|
"""
|
|
planner_config = _host_agent_planner_config()
|
|
if not planner_config.enabled:
|
|
return None
|
|
|
|
resolved_config = host_agent_config or load_host_agent_config()
|
|
if resolved_config.ai_planner_transport != "cloud":
|
|
return None
|
|
|
|
return AIPlanner(
|
|
client=CloudProxyToolCallingClient(resolved_config),
|
|
config=planner_config,
|
|
)
|
|
|
|
|
|
def _device_platform(manager: DeviceManager, device_id: str) -> str | None:
|
|
for device in manager.list_devices():
|
|
if device.id != device_id:
|
|
continue
|
|
if device.driver_type == "wda":
|
|
return "ios"
|
|
if device.driver_type == "uiautomator2":
|
|
return "android"
|
|
return None
|
|
return None
|