Compare commits
36
Commits
358f4623ba
...
master
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
9076f8ddb0 | ||
|
|
8c99dc015a | ||
|
|
60ee157e97 | ||
|
|
dd8df33910 | ||
|
|
5458f3b8a4 | ||
|
|
44e1a6651a | ||
|
|
5b8daab457 | ||
|
|
697e54427b | ||
|
|
fdaca7539b | ||
|
|
3c9e65c78e | ||
|
|
71fd182f50 | ||
|
|
a315c62f3a | ||
|
|
050d1329c4 | ||
|
|
433ab41f95 | ||
|
|
70e0624a47 | ||
|
|
e69cea0245 | ||
|
|
6d9237a592 | ||
|
|
6241fb9d6d | ||
|
|
2d0c740c88 | ||
|
|
ce64c4eb47 | ||
|
|
ce2469616e | ||
|
|
dcb4798408 | ||
|
|
ab15218b27 | ||
|
|
b0932dd398 | ||
|
|
9e3007e7f6 | ||
|
|
98089b6748 | ||
|
|
b73db01626 | ||
|
|
cf8affe4d7 | ||
|
|
61c923b92b | ||
|
|
c7faee8da3 | ||
|
|
29b9a8c39a | ||
|
|
d1b0fffabb | ||
|
|
47eac0f2a7 | ||
|
|
2f8f4a36c6 | ||
|
|
28dccb908c | ||
|
|
3b62195a36 |
@@ -60,6 +60,83 @@ operator view for the tasks that actually execute on that Host, including
|
||||
per-step screenshots, OCR observations, and UI-tree results. The Cloud Console
|
||||
remains the fleet-level view for dispatch status and Cloud-proxy planner history.
|
||||
|
||||
## Local-only Host Agent
|
||||
|
||||
Run without a Cloud Control Plane by setting `HOST_AGENT_MODE=local`. Tasks
|
||||
submitted in the Host console are queued and executed in the same process:
|
||||
|
||||
```bash
|
||||
export HOST_AGENT_MODE=local
|
||||
export AI_PLANNER_ENABLED=true
|
||||
export AI_PLANNER_PROVIDER=openai-compatible
|
||||
export AI_PLANNER_MODEL=qwen2.5
|
||||
export AI_PLANNER_API_KEY=local-key
|
||||
export AI_PLANNER_BASE_URL=http://127.0.0.1:11434/v1
|
||||
export AI_PLANNER_MULTIMODAL=true
|
||||
uv run --package device-host-agent device-host-agent setup
|
||||
uv run --package device-host-agent device-host-agent
|
||||
```
|
||||
|
||||
`openai-compatible` works with Ollama, LM Studio, vLLM, or another server that
|
||||
implements OpenAI `/chat/completions`. Hosted `openai` and `anthropic` providers
|
||||
also accept `AI_PLANNER_API_KEY` and their conventional API key variables.
|
||||
|
||||
The local Host console also exposes an authenticated conversational Agent API:
|
||||
|
||||
```text
|
||||
POST http://127.0.0.1:8765/api/chat
|
||||
```
|
||||
|
||||
Send a JSON body containing `device_id` and `messages` (`user`/`assistant` roles). The Agent can
|
||||
return ordinary assistant text or call the same device-operation tool contracts
|
||||
used by the Runtime; each tool result is fed back to the model before the final
|
||||
reply is returned. The endpoint uses the Host console session cookie, so it is
|
||||
not an unauthenticated device-control endpoint.
|
||||
|
||||
Every chat request is bound to exactly one registered `device_id`. Keep a
|
||||
separate message history and client session for each phone; the Agent injects
|
||||
the bound device into device tools and rejects cross-device tool arguments.
|
||||
|
||||
For vision-capable OpenAI models, a user message may contain standard OpenAI
|
||||
multimodal blocks:
|
||||
|
||||
```json
|
||||
{
|
||||
"messages": [{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "点击图片中的登录按钮"},
|
||||
{"type": "image_url", "image_url": {"url": "data:image/png;base64,..."}}
|
||||
]
|
||||
}]
|
||||
}
|
||||
```
|
||||
|
||||
The compact form `{ "role": "user", "text": "...", "image_base64": "..." }`
|
||||
is also accepted. The model can inspect the supplied image and then call a
|
||||
phone tool such as `tap` in the same conversation.
|
||||
|
||||
Set `AI_PLANNER_MULTIMODAL=true` for a vision model. The Planner sends the
|
||||
screenshot and omits OCR-only elements and OCR metadata from the structured
|
||||
scene payload, avoiding duplicate OCR text.
|
||||
|
||||
Local mode records LLM responses, reasoning fields, tool calls, tool results,
|
||||
and final replies in a local SQLite database. View them at
|
||||
`http://127.0.0.1:8765/conversations`; image bytes are excluded. Set
|
||||
`HOST_AGENT_CONVERSATION_LOG_PATH` to change the database path.
|
||||
|
||||
The authenticated `Devices` page has an on-demand `Get screenshot` button for
|
||||
each connected device. Screenshots are captured only after the operator clicks
|
||||
the button; the page does not auto-refresh or capture screenshots as part of
|
||||
heartbeat synchronization.
|
||||
|
||||
In local mode, Appium supervision is enabled by default. Host Agent probes
|
||||
`/status`, adopts a healthy existing Appium instance, starts Appium when no
|
||||
listener exists, restarts only processes it started if they crash, and stops
|
||||
those child processes on shutdown. Override `HOST_AGENT_APPIUM_*` or set
|
||||
`HOST_AGENT_DEPENDENCY_SUPERVISOR_ENABLED=false` when an external process
|
||||
manager owns Appium.
|
||||
|
||||
## Project Direction
|
||||
|
||||
The durable roadmap is in [docs/ROADMAP.md](docs/ROADMAP.md). The architecture
|
||||
|
||||
+27
-25
@@ -3,7 +3,7 @@ from collections.abc import Callable
|
||||
from typing import Any
|
||||
|
||||
from api.errors import call_with_semantic_errors
|
||||
from device.manager import DEFAULT_MANAGER, DeviceManager
|
||||
from device.manager import 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
|
||||
@@ -17,12 +17,11 @@ from tools.ui_tree import get_ui_tree
|
||||
|
||||
def tool_handlers(
|
||||
*,
|
||||
manager: DeviceManager | None = None,
|
||||
manager: DeviceManager,
|
||||
) -> dict[str, Callable[..., Any]]:
|
||||
device_manager = manager or DEFAULT_MANAGER
|
||||
|
||||
def _screenshot(device_id: str | None = None) -> dict[str, Any]:
|
||||
image = take_screenshot(device_id, manager=device_manager)
|
||||
image = take_screenshot(device_id, manager=manager)
|
||||
return {
|
||||
"ok": True,
|
||||
"image_base64": base64.b64encode(image).decode("ascii"),
|
||||
@@ -39,65 +38,65 @@ def tool_handlers(
|
||||
x,
|
||||
y,
|
||||
device_id=device_id,
|
||||
manager=device_manager,
|
||||
manager=manager,
|
||||
),
|
||||
"swipe": lambda start_x, start_y, end_x, end_y, duration_ms=500, device_id=None: call_with_semantic_errors(
|
||||
swipe,
|
||||
start_x,
|
||||
start_y,
|
||||
end_x,
|
||||
end_y,
|
||||
duration_ms=duration_ms,
|
||||
device_id=device_id,
|
||||
manager=device_manager,
|
||||
"swipe": lambda start_x, start_y, end_x, end_y, duration_ms=500, device_id=None: (
|
||||
call_with_semantic_errors(
|
||||
swipe,
|
||||
start_x,
|
||||
start_y,
|
||||
end_x,
|
||||
end_y,
|
||||
duration_ms=duration_ms,
|
||||
device_id=device_id,
|
||||
manager=manager,
|
||||
)
|
||||
),
|
||||
"input_text": lambda text, device_id=None: call_with_semantic_errors(
|
||||
input_text,
|
||||
text,
|
||||
device_id=device_id,
|
||||
manager=device_manager,
|
||||
manager=manager,
|
||||
),
|
||||
"launch_app": lambda app_id, device_id=None: call_with_semantic_errors(
|
||||
launch_app,
|
||||
app_id,
|
||||
device_id=device_id,
|
||||
manager=device_manager,
|
||||
manager=manager,
|
||||
),
|
||||
"find_text": lambda query, device_id=None: call_with_semantic_errors(
|
||||
find_text_on_screen,
|
||||
query,
|
||||
device_id=device_id,
|
||||
manager=device_manager,
|
||||
manager=manager,
|
||||
),
|
||||
"find_icon": lambda name, device_id=None: call_with_semantic_errors(
|
||||
find_icon_on_screen,
|
||||
name,
|
||||
device_id=device_id,
|
||||
manager=device_manager,
|
||||
manager=manager,
|
||||
),
|
||||
"get_ui_tree": lambda device_id=None, include_app_info=False: (
|
||||
call_with_semantic_errors(
|
||||
get_ui_tree,
|
||||
device_id,
|
||||
manager=device_manager,
|
||||
manager=manager,
|
||||
include_app_info=include_app_info,
|
||||
)
|
||||
),
|
||||
"describe_screen": lambda device_id=None: call_with_semantic_errors(
|
||||
lambda: describe_screen(device_id, manager=device_manager).to_dict()
|
||||
lambda: describe_screen(device_id, manager=manager).to_dict()
|
||||
),
|
||||
"list_devices": lambda: [
|
||||
device.to_dict() for device in device_manager.list_devices()
|
||||
],
|
||||
"list_devices": lambda: [device.to_dict() for device in manager.list_devices()],
|
||||
"device_status": lambda device_id: call_with_semantic_errors(
|
||||
lambda: {"device_id": device_id, "status": device_manager.status(device_id)}
|
||||
lambda: {"device_id": device_id, "status": manager.status(device_id)}
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
def create_mcp_server(
|
||||
*,
|
||||
manager: DeviceManager | None = None,
|
||||
manager: DeviceManager,
|
||||
skill_catalog_store: Any | None = None,
|
||||
skill_active_subscriptions: set[str] | None = None,
|
||||
skill_local_store: Any | None = None,
|
||||
@@ -107,6 +106,9 @@ def create_mcp_server(
|
||||
except ImportError as exc:
|
||||
raise RuntimeError("mcp SDK is not installed") from exc
|
||||
|
||||
if manager is None:
|
||||
raise ValueError("create_mcp_server requires a non-None manager")
|
||||
|
||||
handlers = tool_handlers(manager=manager)
|
||||
server = FastMCP("apex-agent")
|
||||
|
||||
|
||||
@@ -1,8 +1,9 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
from contextlib import suppress
|
||||
from dataclasses import dataclass
|
||||
from dataclasses import dataclass, replace
|
||||
|
||||
import uvicorn
|
||||
|
||||
@@ -11,6 +12,8 @@ from device.manager import DeviceManager
|
||||
from host_agent.assignment import AssignmentExecutor
|
||||
from host_agent.client import HostAgentClient, HostAgentEnrollmentClient
|
||||
from host_agent.config import HostAgentConfig, load_host_agent_config
|
||||
from host_agent.conversation import ConversationAgent
|
||||
from host_agent.conversation_log import ConversationLogStore
|
||||
from host_agent.dependency_supervisor import DependencySupervisor
|
||||
from host_agent.devices import register_local_device
|
||||
from host_agent.enrollment import resolve_host_identity
|
||||
@@ -21,6 +24,9 @@ from host_agent.identity import HostIdentityStore
|
||||
from host_agent.instance_lock import InstanceLock
|
||||
from host_agent.lease import ActiveAssignmentRunner
|
||||
from host_agent.local_account import LocalAccountStore
|
||||
from host_agent.local_client import LocalHostAgentClient
|
||||
from host_agent.mcp_lock import McpBusyTracker
|
||||
from host_agent.mcp_token import McpTokenStore
|
||||
from host_agent.policy_cache import HostPolicyCacheStore
|
||||
from host_agent.processor import AssignmentProcessingResult, AssignmentProcessor
|
||||
from host_agent.retention import prune_task_history
|
||||
@@ -28,6 +34,9 @@ from host_agent.skill_sync import HostAgentSkillSync
|
||||
from host_agent.status import AgentStatusTracker
|
||||
from host_agent.web.app import create_console_app
|
||||
from host_agent.web.auth import SessionManager
|
||||
from host_agent.web.mcp import build_mcp_server
|
||||
from runtime.executor import default_tool_registry
|
||||
from runtime.planner_config import load_config as load_planner_config
|
||||
from storage.artifact_store import ArtifactStore
|
||||
from storage.device_config import DeviceConfigStore
|
||||
from storage.task_metadata import TaskMetadataStore
|
||||
@@ -169,25 +178,37 @@ def create_application(
|
||||
startup_config.identity_path
|
||||
)
|
||||
owned_enrollment_client = enrollment_client is None
|
||||
bootstrap_client = enrollment_client or HostAgentEnrollmentClient(
|
||||
startup_config
|
||||
)
|
||||
bootstrap_client = enrollment_client or HostAgentEnrollmentClient(startup_config)
|
||||
try:
|
||||
resolved_config = resolve_host_identity(
|
||||
startup_config,
|
||||
identity_store=resolved_identity_store,
|
||||
client=bootstrap_client,
|
||||
)
|
||||
bootstrap_client.config = resolved_config
|
||||
resolved_manager = manager or _configured_device_manager(
|
||||
config_store,
|
||||
config=resolved_config,
|
||||
enrollment_client=bootstrap_client,
|
||||
)
|
||||
if startup_config.mode == "local":
|
||||
resolved_config = replace(
|
||||
startup_config, control_plane_url="", host_id="local-host"
|
||||
)
|
||||
resolved_manager = manager or _configured_device_manager(
|
||||
config_store, config=resolved_config, enrollment_client=None
|
||||
)
|
||||
else:
|
||||
resolved_config = resolve_host_identity(
|
||||
startup_config,
|
||||
identity_store=resolved_identity_store,
|
||||
client=bootstrap_client,
|
||||
)
|
||||
bootstrap_client.config = resolved_config
|
||||
resolved_manager = manager or _configured_device_manager(
|
||||
config_store,
|
||||
config=resolved_config,
|
||||
enrollment_client=bootstrap_client,
|
||||
)
|
||||
finally:
|
||||
if owned_enrollment_client:
|
||||
bootstrap_client.close()
|
||||
client = HostAgentClient(resolved_config)
|
||||
if resolved_config.mode == "local":
|
||||
client = LocalHostAgentClient(
|
||||
host_id="local-host",
|
||||
device_ids=lambda: [device.id for device in resolved_manager.list_devices()],
|
||||
)
|
||||
else:
|
||||
client = HostAgentClient(resolved_config)
|
||||
|
||||
history_store = ConsoleHistoryStore(
|
||||
resolved_config.identity_path.parent / "host_console_history.sqlite3",
|
||||
@@ -201,13 +222,40 @@ def create_application(
|
||||
db_path=resolved_config.task_progress_db_path
|
||||
)
|
||||
timeline = Timeline(ArtifactStore(root=resolved_config.task_artifact_dir))
|
||||
mcp_token_path = resolved_config.identity_path.parent / "host_mcp_token.json"
|
||||
mcp_token_existed = mcp_token_path.exists()
|
||||
mcp_token_store = McpTokenStore(mcp_token_path)
|
||||
mcp_token_store.load_or_create()
|
||||
if not mcp_token_existed:
|
||||
logging.getLogger(__name__).info(
|
||||
"MCP token generated at %s", mcp_token_path
|
||||
)
|
||||
mcp_busy_tracker = McpBusyTracker(ttl_seconds=20.0)
|
||||
conversation_log = ConversationLogStore(resolved_config.conversation_log_path) if resolved_config.mode == "local" else None
|
||||
executor = AssignmentExecutor(
|
||||
create_execution_factories(
|
||||
resolved_manager,
|
||||
metadata_store=metadata_store,
|
||||
timeline=timeline,
|
||||
host_agent_config=resolved_config,
|
||||
conversation_log=conversation_log,
|
||||
),
|
||||
mcp_busy_tracker=mcp_busy_tracker,
|
||||
)
|
||||
mcp_server = build_mcp_server(
|
||||
manager=resolved_manager,
|
||||
mcp_busy_tracker=mcp_busy_tracker,
|
||||
status_tracker=status_tracker,
|
||||
)
|
||||
planner_config = load_planner_config()
|
||||
conversation_agent = (
|
||||
ConversationAgent(
|
||||
config=planner_config,
|
||||
tools=default_tool_registry(manager=resolved_manager),
|
||||
event_logger=conversation_log.append if conversation_log else None,
|
||||
)
|
||||
if resolved_config.ai_planner_transport == "direct"
|
||||
else None
|
||||
)
|
||||
console_app = create_console_app(
|
||||
config=resolved_config,
|
||||
@@ -225,6 +273,11 @@ def create_application(
|
||||
metadata_store=metadata_store,
|
||||
timeline=timeline,
|
||||
executor=executor,
|
||||
mcp_server=mcp_server,
|
||||
mcp_token_store=mcp_token_store,
|
||||
mcp_busy_tracker=mcp_busy_tracker,
|
||||
conversation_agent=conversation_agent,
|
||||
conversation_log=conversation_log,
|
||||
)
|
||||
console_server = _EmbeddedConsoleServer(
|
||||
uvicorn.Config(
|
||||
@@ -240,6 +293,7 @@ def create_application(
|
||||
client,
|
||||
resolved_config,
|
||||
status_tracker=status_tracker,
|
||||
mcp_busy_tracker=mcp_busy_tracker,
|
||||
on_sync=lambda device_count: history_store.record_heartbeat(
|
||||
device_count=device_count
|
||||
),
|
||||
@@ -329,7 +383,7 @@ def _configured_device_manager(
|
||||
config_store: DeviceConfigStore,
|
||||
*,
|
||||
config: HostAgentConfig,
|
||||
enrollment_client: HostAgentEnrollmentClient,
|
||||
enrollment_client: HostAgentEnrollmentClient | None,
|
||||
) -> DeviceManager:
|
||||
manager = DeviceManager()
|
||||
for device in config_store.list():
|
||||
|
||||
@@ -2,7 +2,7 @@ from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from cloud.internal_api.models import AssignmentModel
|
||||
from core.models import Task
|
||||
@@ -11,6 +11,9 @@ from host_agent.planner_context import bind_planner_execution_context
|
||||
from host_agent.progress import TaskProgressHolder, TaskProgressSnapshot
|
||||
from runtime.task import is_cancellation_reason
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from host_agent.mcp_lock import McpBusyTracker
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class AssignmentExecutionResult:
|
||||
@@ -20,9 +23,15 @@ class AssignmentExecutionResult:
|
||||
|
||||
|
||||
class AssignmentExecutor:
|
||||
def __init__(self, factories: ExecutionFactories) -> None:
|
||||
def __init__(
|
||||
self,
|
||||
factories: ExecutionFactories,
|
||||
*,
|
||||
mcp_busy_tracker: McpBusyTracker | None = None,
|
||||
) -> None:
|
||||
self.factories = factories
|
||||
self._progress = TaskProgressHolder()
|
||||
self._mcp_busy_tracker = mcp_busy_tracker
|
||||
|
||||
def latest_progress(self) -> TaskProgressSnapshot | None:
|
||||
"""Latest step progress reported by the currently-running assignment."""
|
||||
@@ -36,6 +45,15 @@ class AssignmentExecutor:
|
||||
stop_reason: Callable[[], str | None] | None = None,
|
||||
) -> AssignmentExecutionResult:
|
||||
self._progress.clear()
|
||||
if self._mcp_busy_tracker is not None and (
|
||||
assignment.device_id in self._mcp_busy_tracker.busy_device_ids()
|
||||
):
|
||||
return AssignmentExecutionResult(
|
||||
status="failed",
|
||||
failure_reason=(
|
||||
f"device {assignment.device_id} is held by an active MCP session"
|
||||
),
|
||||
)
|
||||
with bind_planner_execution_context(assignment):
|
||||
if should_stop is not None and should_stop():
|
||||
reason = stop_reason() if stop_reason is not None else None
|
||||
|
||||
@@ -2,7 +2,9 @@ from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import getpass
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from collections.abc import Sequence
|
||||
from dataclasses import replace
|
||||
|
||||
@@ -10,6 +12,7 @@ from host_agent.app import create_application
|
||||
from host_agent.config import load_host_agent_config
|
||||
from host_agent.instance_lock import InstanceAlreadyRunningError
|
||||
from host_agent.local_account import LocalAccountStore
|
||||
from host_agent.mcp_token import McpTokenStore
|
||||
|
||||
|
||||
class LocalAccountSetupError(RuntimeError):
|
||||
@@ -17,11 +20,20 @@ class LocalAccountSetupError(RuntimeError):
|
||||
|
||||
|
||||
def main(argv: Sequence[str] | None = None) -> None:
|
||||
_load_dotenv()
|
||||
parser = argparse.ArgumentParser(description="Run the Device Host Agent")
|
||||
subparsers = parser.add_subparsers(dest="command")
|
||||
subparsers.add_parser("setup", help="Create the local operator account")
|
||||
subparsers.add_parser(
|
||||
"mcp-token",
|
||||
help="Print the MCP server bearer token (generating if missing)",
|
||||
)
|
||||
args = parser.parse_args(argv)
|
||||
|
||||
if args.command == "mcp-token":
|
||||
_print_mcp_token()
|
||||
return
|
||||
|
||||
try:
|
||||
if args.command == "setup":
|
||||
_run_setup()
|
||||
@@ -54,6 +66,12 @@ def _run_setup() -> None:
|
||||
print(f"Local account '{account.username}' created.")
|
||||
|
||||
|
||||
def _print_mcp_token() -> None:
|
||||
config = load_host_agent_config()
|
||||
store = McpTokenStore(config.identity_path.parent / "host_mcp_token.json")
|
||||
print(store.load_or_create().token)
|
||||
|
||||
|
||||
def _resolve_config_with_local_account():
|
||||
config = load_host_agent_config()
|
||||
store = LocalAccountStore(config.local_account_path)
|
||||
@@ -81,3 +99,19 @@ def _prompt_and_create(store: LocalAccountStore):
|
||||
if password != confirm:
|
||||
raise LocalAccountSetupError("passwords do not match")
|
||||
return store.create(username, password)
|
||||
|
||||
|
||||
def _load_dotenv() -> None:
|
||||
"""Load a simple repository-root .env without overriding shell values."""
|
||||
path = Path.cwd() / ".env"
|
||||
if not path.is_file():
|
||||
return
|
||||
for line in path.read_text(encoding="utf-8").splitlines():
|
||||
line = line.strip()
|
||||
if not line or line.startswith("#") or "=" not in line:
|
||||
continue
|
||||
name, value = line.split("=", 1)
|
||||
name = name.strip()
|
||||
value = value.strip()
|
||||
if name and name not in os.environ:
|
||||
os.environ[name] = value
|
||||
|
||||
@@ -168,17 +168,21 @@ class HostAgentClient:
|
||||
*,
|
||||
address: str | None = None,
|
||||
policy_revision: int = 0,
|
||||
mcp_busy_device_ids: list[str] | None = None,
|
||||
) -> HeartbeatResponse:
|
||||
payload: dict[str, Any] = {
|
||||
"host_id": self.config.host_id,
|
||||
"address": address,
|
||||
"devices": [device.model_dump(mode="json") for device in devices],
|
||||
"policy_revision": policy_revision,
|
||||
"planner_transport": self.config.ai_planner_transport,
|
||||
}
|
||||
if mcp_busy_device_ids:
|
||||
payload["mcp_busy_device_ids"] = list(mcp_busy_device_ids)
|
||||
response = await self._request(
|
||||
"PUT",
|
||||
f"/internal/v1/hosts/{self.config.host_id}/heartbeat",
|
||||
json={
|
||||
"host_id": self.config.host_id,
|
||||
"address": address,
|
||||
"devices": [device.model_dump(mode="json") for device in devices],
|
||||
"policy_revision": policy_revision,
|
||||
"planner_transport": self.config.ai_planner_transport,
|
||||
},
|
||||
json=payload,
|
||||
)
|
||||
return HeartbeatResponse.model_validate(response.json())
|
||||
|
||||
|
||||
@@ -62,11 +62,13 @@ class CloudProxyToolCallingClient:
|
||||
screenshot: bytes | None,
|
||||
tools: list[ToolSpec],
|
||||
timeout: float,
|
||||
history: list[dict[str, Any]] | None = None,
|
||||
) -> ToolCallDecision:
|
||||
payload: dict[str, Any] = {
|
||||
"host_id": self.config.host_id,
|
||||
"system_prompt": system_prompt,
|
||||
"user_prompt": user_prompt,
|
||||
"history": history or [],
|
||||
"screenshot_base64": (
|
||||
base64.b64encode(screenshot).decode("ascii")
|
||||
if screenshot is not None
|
||||
|
||||
@@ -13,6 +13,7 @@ class HostAgentConfigurationError(ValueError):
|
||||
|
||||
_LOOPBACK_BIND_HOSTS = frozenset({"127.0.0.1", "localhost", "::1"})
|
||||
_AI_PLANNER_TRANSPORTS = frozenset({"direct", "cloud"})
|
||||
_HOST_AGENT_MODES = frozenset({"cloud", "local"})
|
||||
_REMOVED_RUNTIME_SUPERVISION_SETTINGS = (
|
||||
"HOST_AGENT_RUNTIME_SUPERVISED",
|
||||
"HOST_AGENT_RUNTIME_HOST",
|
||||
@@ -22,7 +23,8 @@ _REMOVED_RUNTIME_SUPERVISION_SETTINGS = (
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class HostAgentConfig:
|
||||
control_plane_url: str
|
||||
control_plane_url: str = ""
|
||||
mode: str = "cloud"
|
||||
host_id: str = ""
|
||||
token: str = field(default="", repr=False)
|
||||
identity_path: Path = Path("tasks/host_identity.json")
|
||||
@@ -47,6 +49,7 @@ class HostAgentConfig:
|
||||
dependency_restart_max_attempts: int = 5
|
||||
task_progress_db_path: Path = Path("host_agent_data/task_progress.sqlite3")
|
||||
task_artifact_dir: Path = Path("host_agent_data/history")
|
||||
conversation_log_path: Path = Path("host_agent_data/conversations.sqlite3")
|
||||
task_retention_max_count: int = 50
|
||||
task_retention_max_age_days: int = 7
|
||||
skill_sync_interval_seconds: float = 300.0
|
||||
@@ -57,16 +60,19 @@ def load_host_agent_config(
|
||||
) -> HostAgentConfig:
|
||||
values = os.environ if env is None else env
|
||||
_reject_removed_runtime_supervision_settings(values)
|
||||
mode = values.get("HOST_AGENT_MODE", "cloud").strip().lower()
|
||||
if mode not in _HOST_AGENT_MODES:
|
||||
raise HostAgentConfigurationError("HOST_AGENT_MODE must be cloud or local")
|
||||
control_plane_url = (
|
||||
values.get(
|
||||
"HOST_AGENT_CONTROL_PLANE_URL",
|
||||
"https://amcp.home.jerryyan.top",
|
||||
"https://amcp.home.jerryyan.top" if mode == "cloud" else "",
|
||||
)
|
||||
.strip()
|
||||
.rstrip("/")
|
||||
)
|
||||
parsed_url = urlparse(control_plane_url)
|
||||
if parsed_url.scheme not in {"http", "https"} or not parsed_url.netloc:
|
||||
if mode == "cloud" and (parsed_url.scheme not in {"http", "https"} or not parsed_url.netloc):
|
||||
raise HostAgentConfigurationError(
|
||||
"HOST_AGENT_CONTROL_PLANE_URL must be an HTTP(S) URL"
|
||||
)
|
||||
@@ -82,9 +88,10 @@ def load_host_agent_config(
|
||||
|
||||
config = HostAgentConfig(
|
||||
control_plane_url=control_plane_url,
|
||||
mode=mode,
|
||||
identity_path=identity_path,
|
||||
local_account_path=local_account_path,
|
||||
enrollment_managed=True,
|
||||
enrollment_managed=mode == "cloud",
|
||||
display_name=values.get("HOST_AGENT_DISPLAY_NAME") or None,
|
||||
heartbeat_interval_seconds=_positive_float(
|
||||
values,
|
||||
@@ -128,13 +135,15 @@ def load_host_agent_config(
|
||||
"HOST_AGENT_CONSOLE_HISTORY_LIMIT",
|
||||
200,
|
||||
),
|
||||
ai_planner_transport=_parse_ai_planner_transport(
|
||||
values.get("AI_PLANNER_TRANSPORT")
|
||||
),
|
||||
ai_planner_transport=("direct" if mode == "local" else _parse_ai_planner_transport(values.get("AI_PLANNER_TRANSPORT"))),
|
||||
dependency_supervisor_enabled=_truthy(
|
||||
values, "HOST_AGENT_DEPENDENCY_SUPERVISOR_ENABLED", False
|
||||
values,
|
||||
"HOST_AGENT_DEPENDENCY_SUPERVISOR_ENABLED",
|
||||
mode == "local",
|
||||
),
|
||||
appium_supervised=_truthy(
|
||||
values, "HOST_AGENT_APPIUM_SUPERVISED", mode == "local"
|
||||
),
|
||||
appium_supervised=_truthy(values, "HOST_AGENT_APPIUM_SUPERVISED", False),
|
||||
appium_host=values.get("HOST_AGENT_APPIUM_HOST", "127.0.0.1").strip(),
|
||||
appium_port=_positive_int(values, "HOST_AGENT_APPIUM_PORT", 4723),
|
||||
dependency_restart_max_attempts=_positive_int(
|
||||
@@ -151,6 +160,7 @@ def load_host_agent_config(
|
||||
"HOST_AGENT_TASK_ARTIFACT_DIR", "host_agent_data/history"
|
||||
).strip()
|
||||
),
|
||||
conversation_log_path=Path(values.get("HOST_AGENT_CONVERSATION_LOG_PATH", "host_agent_data/conversations.sqlite3").strip()),
|
||||
task_retention_max_count=_positive_int(
|
||||
values, "HOST_AGENT_TASK_RETENTION_MAX_COUNT", 50
|
||||
),
|
||||
|
||||
@@ -0,0 +1,291 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Callable
|
||||
|
||||
from runtime.planner_config import PlannerConfig
|
||||
from runtime.tool_specs import ACTION_TOOL_SPECS, ToolSpec
|
||||
|
||||
READ_TOOL_SPECS = [
|
||||
ToolSpec("take_screenshot", "Capture the current device screen.", {"type": "object", "properties": {"device_id": {"type": "string"}}}),
|
||||
ToolSpec("describe_screen", "Inspect the current screen and UI.", {"type": "object", "properties": {"device_id": {"type": "string"}}}),
|
||||
ToolSpec("find_text", "Find visible text on the current screen.", {"type": "object", "required": ["query"], "properties": {"query": {"type": "string"}, "device_id": {"type": "string"}}}),
|
||||
ToolSpec("get_ui_tree", "Read the current accessibility/UI tree.", {"type": "object", "properties": {"device_id": {"type": "string"}, "include_app_info": {"type": "boolean"}}}),
|
||||
ToolSpec("list_devices", "List configured devices.", {"type": "object", "properties": {}}),
|
||||
ToolSpec("device_status", "Read one device status.", {"type": "object", "required": ["device_id"], "properties": {"device_id": {"type": "string"}}}),
|
||||
]
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ChatResult:
|
||||
content: str
|
||||
tool_calls: int
|
||||
|
||||
|
||||
class ConversationAgent:
|
||||
"""Small OpenAI-compatible agent loop backed by Host Agent tools."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
config: PlannerConfig,
|
||||
tools: dict[str, Callable[..., Any]],
|
||||
max_rounds: int = 8,
|
||||
event_logger: Callable[[dict[str, Any]], None] | None = None,
|
||||
) -> None:
|
||||
self.config = config
|
||||
self.tools = tools
|
||||
self.max_rounds = max_rounds
|
||||
self.event_logger = event_logger
|
||||
self._client: Any | None = None
|
||||
|
||||
def chat(self, messages: list[dict[str, Any]], *, tools: dict[str, Callable[..., Any]] | None = None) -> ChatResult:
|
||||
if self.config.provider == "anthropic":
|
||||
return self._chat_anthropic(messages, tools=tools)
|
||||
if self.config.provider not in {"openai", "openai_compatible"}:
|
||||
raise RuntimeError("unsupported conversation provider")
|
||||
client = self._get_client()
|
||||
active_tools = tools or self.tools
|
||||
history: list[dict[str, Any]] = [
|
||||
{
|
||||
"role": "system",
|
||||
"content": (
|
||||
"You are a mobile device assistant. Reply naturally when no "
|
||||
"device action is needed. Call a tool when the user asks you "
|
||||
"to inspect or operate a device. Never claim an action was "
|
||||
"completed unless the tool result confirms it."
|
||||
),
|
||||
},
|
||||
*[_normalize_message(message) for message in messages],
|
||||
]
|
||||
calls = 0
|
||||
for _ in range(self.max_rounds):
|
||||
response = client.chat.completions.create(
|
||||
model=self.config.resolved_model(),
|
||||
messages=history,
|
||||
tools=[_openai_tool(spec) for spec in [*ACTION_TOOL_SPECS, *READ_TOOL_SPECS]],
|
||||
tool_choice="auto",
|
||||
parallel_tool_calls=False,
|
||||
timeout=self.config.timeout,
|
||||
)
|
||||
message = response.choices[0].message
|
||||
text = getattr(message, "content", None)
|
||||
thinking = getattr(message, "reasoning_content", None) or getattr(message, "reasoning", None)
|
||||
tool_calls = getattr(message, "tool_calls", None) or []
|
||||
self._log({"type": "llm_response", "content": text, "thinking": thinking,
|
||||
"tool_calls": [{"id": c.id, "name": c.function.name, "arguments": c.function.arguments} for c in tool_calls]})
|
||||
if not tool_calls:
|
||||
self._log({"type": "final_reply", "content": str(text or "")})
|
||||
return ChatResult(content=str(text or ""), tool_calls=calls)
|
||||
history.append(_message_dict(message))
|
||||
for tool_call in tool_calls:
|
||||
name = tool_call.function.name
|
||||
arguments: dict[str, Any] = {}
|
||||
try:
|
||||
decoded_arguments = json.loads(tool_call.function.arguments or "{}")
|
||||
if not isinstance(decoded_arguments, dict):
|
||||
raise ValueError("tool arguments must be an object")
|
||||
arguments = decoded_arguments
|
||||
arguments.pop("purpose", None)
|
||||
arguments.pop("expected_outcome", None)
|
||||
result = active_tools[name](**arguments)
|
||||
except Exception as exc:
|
||||
result = {"ok": False, "error": str(exc)}
|
||||
calls += 1
|
||||
self._log({"type": "tool_result", "tool_call_id": tool_call.id,
|
||||
"tool_name": name, "arguments": arguments, "result": result})
|
||||
history.append(
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": tool_call.id,
|
||||
"content": json.dumps(result, ensure_ascii=True, default=str),
|
||||
}
|
||||
)
|
||||
raise RuntimeError("conversation exceeded the maximum tool-call rounds")
|
||||
|
||||
def _chat_anthropic(self, messages: list[dict[str, Any]], *, tools: dict[str, Callable[..., Any]] | None = None) -> ChatResult:
|
||||
import anthropic
|
||||
|
||||
kwargs: dict[str, Any] = {}
|
||||
if self.config.api_key:
|
||||
kwargs["api_key"] = self.config.api_key
|
||||
if self.config.base_url:
|
||||
kwargs["base_url"] = self.config.base_url
|
||||
client = anthropic.Anthropic(**kwargs)
|
||||
history = [_normalize_anthropic_message(message) for message in messages]
|
||||
active_tools = tools or self.tools
|
||||
calls = 0
|
||||
for _ in range(self.max_rounds):
|
||||
response = client.messages.create(
|
||||
model=self.config.resolved_model(),
|
||||
max_tokens=2048,
|
||||
system="You are a mobile device assistant. Reply naturally, or call one device tool when needed. Never claim an action succeeded unless its tool result confirms it.",
|
||||
messages=history,
|
||||
tools=[_anthropic_tool(spec) for spec in [*ACTION_TOOL_SPECS, *READ_TOOL_SPECS]],
|
||||
timeout=self.config.timeout,
|
||||
)
|
||||
blocks = getattr(response, "content", []) or []
|
||||
text_parts = [getattr(block, "text", "") for block in blocks if getattr(block, "type", None) == "text"]
|
||||
thinking_parts = [getattr(block, "thinking", "") for block in blocks if getattr(block, "type", None) == "thinking"]
|
||||
uses = [block for block in blocks if getattr(block, "type", None) == "tool_use"]
|
||||
self._log({"type": "llm_response", "content": "\n".join(x for x in text_parts if x), "thinking": "\n".join(x for x in thinking_parts if x), "tool_calls": [{"id": u.id, "name": u.name, "arguments": u.input} for u in uses]})
|
||||
if not uses:
|
||||
content = "\n".join(x for x in text_parts if x)
|
||||
self._log({"type": "final_reply", "content": content})
|
||||
return ChatResult(content=content, tool_calls=calls)
|
||||
history.append({"role": "assistant", "content": [_anthropic_block_dict(block) for block in blocks if getattr(block, "type", None) != "thinking"]})
|
||||
results = []
|
||||
for use in uses:
|
||||
arguments = dict(use.input) if isinstance(use.input, dict) else {}
|
||||
try:
|
||||
result = active_tools[use.name](**arguments)
|
||||
except Exception as exc:
|
||||
result = {"ok": False, "error": str(exc)}
|
||||
calls += 1
|
||||
self._log({"type": "tool_result", "tool_call_id": use.id, "tool_name": use.name, "arguments": arguments, "result": result})
|
||||
results.append({"type": "tool_result", "tool_use_id": use.id, "content": json.dumps(result, ensure_ascii=True, default=str)})
|
||||
history.append({"role": "user", "content": results})
|
||||
raise RuntimeError("conversation exceeded the maximum tool-call rounds")
|
||||
|
||||
def chat_for_device(self, device_id: str, messages: list[dict[str, Any]]) -> ChatResult:
|
||||
"""Run a conversation bound to exactly one device."""
|
||||
device_tools = {
|
||||
name: _bind_device(tool, device_id, name=name)
|
||||
for name, tool in self.tools.items()
|
||||
}
|
||||
return self.chat(messages, tools=device_tools)
|
||||
|
||||
def _log(self, event: dict[str, Any]) -> None:
|
||||
if self.event_logger is not None:
|
||||
try:
|
||||
self.event_logger(event)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def _get_client(self) -> Any:
|
||||
if self._client is None:
|
||||
from openai import OpenAI
|
||||
|
||||
kwargs: dict[str, Any] = {}
|
||||
if self.config.api_key:
|
||||
kwargs["api_key"] = self.config.api_key
|
||||
if self.config.base_url:
|
||||
kwargs["base_url"] = self.config.base_url
|
||||
self._client = OpenAI(**kwargs)
|
||||
return self._client
|
||||
|
||||
|
||||
def _openai_tool(spec: ToolSpec) -> dict[str, Any]:
|
||||
parameters = dict(spec.parameters)
|
||||
properties = dict(parameters.get("properties") or {})
|
||||
properties.setdefault(
|
||||
"device_id",
|
||||
{"type": "string", "description": "Target device ID when needed."},
|
||||
)
|
||||
parameters["properties"] = properties
|
||||
return {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": spec.name,
|
||||
"description": spec.description,
|
||||
"parameters": parameters,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _anthropic_tool(spec: ToolSpec) -> dict[str, Any]:
|
||||
return {"name": spec.name, "description": spec.description, "input_schema": spec.parameters}
|
||||
|
||||
|
||||
def _anthropic_block_dict(block: Any) -> dict[str, Any]:
|
||||
block_type = getattr(block, "type", "")
|
||||
if block_type == "text":
|
||||
return {"type": "text", "text": getattr(block, "text", "")}
|
||||
if block_type == "thinking":
|
||||
return {"type": "thinking", "thinking": getattr(block, "thinking", "")}
|
||||
return {"type": "tool_use", "id": block.id, "name": block.name, "input": block.input}
|
||||
|
||||
|
||||
def _normalize_anthropic_message(message: dict[str, Any]) -> dict[str, Any]:
|
||||
normalized = _normalize_message(message)
|
||||
content = normalized["content"]
|
||||
if isinstance(content, str):
|
||||
return normalized
|
||||
blocks = []
|
||||
for block in content:
|
||||
if block.get("type") == "text":
|
||||
blocks.append(block)
|
||||
elif block.get("type") == "image_url":
|
||||
url = block.get("image_url", {}).get("url", "")
|
||||
if isinstance(url, str) and url.startswith("data:"):
|
||||
header, data = url.split(",", 1)
|
||||
media_type = header[5:].split(";", 1)[0]
|
||||
blocks.append({"type": "image", "source": {"type": "base64", "media_type": media_type, "data": data}})
|
||||
return {"role": normalized["role"], "content": blocks}
|
||||
|
||||
|
||||
def _message_dict(message: Any) -> dict[str, Any]:
|
||||
result: dict[str, Any] = {"role": "assistant"}
|
||||
content = getattr(message, "content", None)
|
||||
if content is not None:
|
||||
result["content"] = content
|
||||
calls = getattr(message, "tool_calls", None) or []
|
||||
result["tool_calls"] = [
|
||||
{
|
||||
"id": call.id,
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": call.function.name,
|
||||
"arguments": call.function.arguments,
|
||||
},
|
||||
}
|
||||
for call in calls
|
||||
]
|
||||
return result
|
||||
|
||||
|
||||
def _normalize_message(message: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Accept OpenAI text/image content blocks and a compact image_base64 form."""
|
||||
role = message.get("role")
|
||||
content = message.get("content")
|
||||
if isinstance(content, str):
|
||||
return {"role": role, "content": content}
|
||||
if isinstance(content, list):
|
||||
blocks: list[dict[str, Any]] = []
|
||||
for block in content:
|
||||
if not isinstance(block, dict):
|
||||
continue
|
||||
if block.get("type") == "text" and isinstance(block.get("text"), str):
|
||||
blocks.append({"type": "text", "text": block["text"]})
|
||||
elif block.get("type") == "image_url":
|
||||
image_url = block.get("image_url")
|
||||
if isinstance(image_url, str):
|
||||
image_url = {"url": image_url}
|
||||
if isinstance(image_url, dict) and isinstance(image_url.get("url"), str):
|
||||
blocks.append({"type": "image_url", "image_url": {"url": image_url["url"]}})
|
||||
if blocks:
|
||||
return {"role": role, "content": blocks}
|
||||
image = message.get("image_base64")
|
||||
if isinstance(image, str) and image:
|
||||
mime = str(message.get("mime_type") or "image/png")
|
||||
text = message.get("text")
|
||||
blocks = []
|
||||
if isinstance(text, str) and text:
|
||||
blocks.append({"type": "text", "text": text})
|
||||
blocks.append({"type": "image_url", "image_url": {"url": f"data:{mime};base64,{image}"}})
|
||||
return {"role": role, "content": blocks}
|
||||
raise ValueError("message content must be text, image blocks, or image_base64")
|
||||
|
||||
|
||||
def _bind_device(tool: Callable[..., Any], device_id: str, *, name: str) -> Callable[..., Any]:
|
||||
def bound(**arguments: Any) -> Any:
|
||||
if name == "list_devices":
|
||||
return tool(**arguments)
|
||||
requested = arguments.get("device_id")
|
||||
if requested is not None and requested != device_id:
|
||||
raise ValueError(f"conversation is bound to device {device_id}")
|
||||
arguments["device_id"] = device_id
|
||||
return tool(**arguments)
|
||||
|
||||
return bound
|
||||
@@ -0,0 +1,48 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import sqlite3
|
||||
from datetime import UTC, datetime
|
||||
from pathlib import Path
|
||||
from threading import Lock
|
||||
from typing import Any
|
||||
|
||||
|
||||
class ConversationLogStore:
|
||||
"""SQLite-backed local audit log for chat and tool activity."""
|
||||
|
||||
def __init__(self, path: str | Path) -> None:
|
||||
self.path = Path(path)
|
||||
self._lock = Lock()
|
||||
self.path.parent.mkdir(parents=True, exist_ok=True)
|
||||
with self._connect() as connection:
|
||||
connection.execute("CREATE TABLE IF NOT EXISTS conversation_events (id INTEGER PRIMARY KEY AUTOINCREMENT, occurred_at TEXT NOT NULL, event_type TEXT NOT NULL, payload_json TEXT NOT NULL)")
|
||||
|
||||
def append(self, event: dict[str, Any]) -> None:
|
||||
with self._lock, self._connect() as connection:
|
||||
connection.execute("INSERT INTO conversation_events (occurred_at, event_type, payload_json) VALUES (?, ?, ?)", (datetime.now(UTC).isoformat(), str(event.get("type") or "event"), json.dumps(_safe(event), ensure_ascii=True, default=str)))
|
||||
|
||||
def list_recent(self, *, limit: int = 200) -> list[dict[str, Any]]:
|
||||
with self._connect() as connection:
|
||||
rows = connection.execute("SELECT id, occurred_at, event_type, payload_json FROM conversation_events ORDER BY id DESC LIMIT ?", (max(1, min(limit, 1000)),)).fetchall()
|
||||
result = []
|
||||
for row in rows:
|
||||
try:
|
||||
payload = json.loads(row["payload_json"])
|
||||
except (TypeError, ValueError):
|
||||
payload = {}
|
||||
result.append({"id": row["id"], "occurred_at": row["occurred_at"], "event_type": row["event_type"], **payload})
|
||||
return result
|
||||
|
||||
def _connect(self) -> sqlite3.Connection:
|
||||
connection = sqlite3.connect(self.path)
|
||||
connection.row_factory = sqlite3.Row
|
||||
return connection
|
||||
|
||||
|
||||
def _safe(value: Any) -> Any:
|
||||
if isinstance(value, dict):
|
||||
return {str(k): _safe(v) for k, v in value.items()}
|
||||
if isinstance(value, list):
|
||||
return [_safe(v) for v in value]
|
||||
return value if isinstance(value, (str, int, float, bool)) or value is None else str(value)
|
||||
@@ -3,6 +3,7 @@ from __future__ import annotations
|
||||
import os
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, replace
|
||||
from typing import Any
|
||||
|
||||
from device.manager import DeviceManager
|
||||
from host_agent.cloud_planner_client import CloudProxyToolCallingClient
|
||||
@@ -35,6 +36,7 @@ def create_execution_factories(
|
||||
metadata_store: TaskMetadataStore | None = None,
|
||||
timeline: Timeline | None = None,
|
||||
host_agent_config: HostAgentConfig | None = None,
|
||||
conversation_log: Any | None = None,
|
||||
) -> ExecutionFactories:
|
||||
shared_workflow_store = workflow_store or WorkflowStore()
|
||||
resolved_host_agent_config = host_agent_config
|
||||
@@ -48,7 +50,10 @@ def create_execution_factories(
|
||||
),
|
||||
metadata_store=metadata_store,
|
||||
timeline=timeline,
|
||||
planner=_host_agent_planner(resolved_host_agent_config),
|
||||
planner=_host_agent_planner(
|
||||
resolved_host_agent_config,
|
||||
event_logger=conversation_log.append if conversation_log is not None else None,
|
||||
),
|
||||
planner_config=_host_agent_planner_config(),
|
||||
device_platform_provider=lambda device_id: _device_platform(
|
||||
manager, device_id
|
||||
@@ -86,6 +91,8 @@ def _host_agent_planner_config() -> PlannerConfig:
|
||||
|
||||
def _host_agent_planner(
|
||||
host_agent_config: HostAgentConfig | None,
|
||||
*,
|
||||
event_logger: Callable[[dict[str, Any]], None] | None = None,
|
||||
) -> Planner | None:
|
||||
"""Build the `AIPlanner` explicitly when the cloud-proxy transport is
|
||||
selected, so its `ToolCallingClient` is a `CloudProxyToolCallingClient`
|
||||
@@ -100,11 +107,12 @@ def _host_agent_planner(
|
||||
|
||||
resolved_config = host_agent_config or load_host_agent_config()
|
||||
if resolved_config.ai_planner_transport != "cloud":
|
||||
return None
|
||||
return AIPlanner(config=planner_config, event_logger=event_logger)
|
||||
|
||||
return AIPlanner(
|
||||
client=CloudProxyToolCallingClient(resolved_config),
|
||||
config=planner_config,
|
||||
event_logger=event_logger,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -14,6 +14,8 @@ from host_agent.status import AgentStatusTracker
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Awaitable, Callable
|
||||
|
||||
from host_agent.mcp_lock import McpBusyTracker
|
||||
|
||||
|
||||
def build_device_snapshot(manager: DeviceManager) -> list[DeviceSnapshotModel]:
|
||||
return [
|
||||
@@ -40,6 +42,7 @@ class HeartbeatSynchronizer:
|
||||
on_sync: Callable[[int], None] | None = None,
|
||||
policy_cache: HostPolicyCacheStore | None = None,
|
||||
on_policy_sync: Callable[[int], None] | None = None,
|
||||
mcp_busy_tracker: McpBusyTracker | None = None,
|
||||
) -> None:
|
||||
self.manager = manager
|
||||
self.client = client
|
||||
@@ -50,17 +53,26 @@ class HeartbeatSynchronizer:
|
||||
self.on_sync = on_sync
|
||||
self.policy_cache = policy_cache
|
||||
self.on_policy_sync = on_policy_sync
|
||||
self.mcp_busy_tracker = mcp_busy_tracker
|
||||
self.policy = policy_cache.load() if policy_cache is not None else None
|
||||
self.policy_revision = self.policy.revision if self.policy is not None else 0
|
||||
if self.status_tracker is not None:
|
||||
self.status_tracker.mark_host_policy(self.policy)
|
||||
|
||||
async def sync_once(self) -> HeartbeatResponse:
|
||||
self.probe_connected_devices()
|
||||
self.connect_devices()
|
||||
snapshot = build_device_snapshot(self.manager)
|
||||
mcp_busy_ids = (
|
||||
self.mcp_busy_tracker.busy_device_ids()
|
||||
if self.mcp_busy_tracker is not None
|
||||
else []
|
||||
)
|
||||
response = await self.client.heartbeat(
|
||||
snapshot,
|
||||
address=self.address,
|
||||
policy_revision=self.policy_revision,
|
||||
mcp_busy_device_ids=mcp_busy_ids,
|
||||
)
|
||||
self.policy_revision = response.policy_revision
|
||||
if response.policy is not None:
|
||||
@@ -98,9 +110,19 @@ class HeartbeatSynchronizer:
|
||||
|
||||
def connect_devices(self) -> None:
|
||||
for device in self.manager.list_devices():
|
||||
if device.status != "idle":
|
||||
if device.status not in {"idle", "offline", "error"}:
|
||||
continue
|
||||
try:
|
||||
self.manager.connect(device.id)
|
||||
except DeviceRuntimeError:
|
||||
continue
|
||||
|
||||
def probe_connected_devices(self) -> None:
|
||||
active_id: str | None = None
|
||||
if self.status_tracker is not None:
|
||||
current = self.status_tracker.snapshot().get("current_assignment")
|
||||
if isinstance(current, dict) and isinstance(current.get("device_id"), str):
|
||||
active_id = current["device_id"]
|
||||
for device in self.manager.list_devices():
|
||||
if device.id != active_id and device.status == "busy":
|
||||
self.manager.probe(device.id)
|
||||
|
||||
@@ -0,0 +1,79 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import subprocess
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
|
||||
class IOSDiscoveryError(RuntimeError):
|
||||
pass
|
||||
|
||||
|
||||
def discover_connected_ios_devices(*, timeout_seconds: int = 10) -> list[dict[str, str]]:
|
||||
"""Return paired, currently connected physical iOS devices from CoreDevice."""
|
||||
with tempfile.TemporaryDirectory(prefix="ios-device-discovery-") as temp_dir:
|
||||
output_path = Path(temp_dir) / "devices.json"
|
||||
try:
|
||||
completed = subprocess.run(
|
||||
[
|
||||
"xcrun",
|
||||
"devicectl",
|
||||
"list",
|
||||
"devices",
|
||||
"--json-output",
|
||||
str(output_path),
|
||||
"--timeout",
|
||||
str(timeout_seconds),
|
||||
"--quiet",
|
||||
],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=timeout_seconds + 2,
|
||||
check=False,
|
||||
)
|
||||
except FileNotFoundError as exc:
|
||||
raise IOSDiscoveryError("xcrun is unavailable; install Xcode command line tools") from exc
|
||||
except subprocess.TimeoutExpired as exc:
|
||||
raise IOSDiscoveryError("iOS device discovery timed out") from exc
|
||||
|
||||
if completed.returncode != 0:
|
||||
detail = completed.stderr.strip() or completed.stdout.strip()
|
||||
raise IOSDiscoveryError(detail or "devicectl failed to discover iOS devices")
|
||||
try:
|
||||
payload = json.loads(output_path.read_text(encoding="utf-8"))
|
||||
except (OSError, ValueError) as exc:
|
||||
raise IOSDiscoveryError("devicectl returned invalid device data") from exc
|
||||
|
||||
raw_devices = payload.get("result", {}).get("devices", [])
|
||||
devices: list[dict[str, str]] = []
|
||||
for item in raw_devices if isinstance(raw_devices, list) else []:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
hardware = item.get("hardwareProperties", {})
|
||||
properties = item.get("deviceProperties", {})
|
||||
connection = item.get("connectionProperties", {})
|
||||
if not all(isinstance(value, dict) for value in (hardware, properties, connection)):
|
||||
continue
|
||||
udid = hardware.get("udid")
|
||||
if (
|
||||
hardware.get("platform") != "iOS"
|
||||
or hardware.get("reality") != "physical"
|
||||
or connection.get("pairingState") != "paired"
|
||||
or connection.get("tunnelState") != "connected"
|
||||
or not isinstance(udid, str)
|
||||
or not udid
|
||||
):
|
||||
continue
|
||||
devices.append(
|
||||
{
|
||||
"udid": udid,
|
||||
"name": str(properties.get("name") or hardware.get("marketingName") or "iPhone"),
|
||||
"model": str(hardware.get("marketingName") or hardware.get("productType") or "iPhone"),
|
||||
"os_version": str(properties.get("osVersionNumber") or ""),
|
||||
"transport": str(connection.get("transportType") or "unknown"),
|
||||
}
|
||||
)
|
||||
return sorted(devices, key=lambda device: (device["name"], device["udid"]))
|
||||
|
||||
@@ -0,0 +1,72 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import uuid
|
||||
from collections.abc import Callable
|
||||
from datetime import UTC, datetime, timedelta
|
||||
|
||||
from cloud.internal_api.models import (
|
||||
ClaimResponse,
|
||||
HeartbeatResponse,
|
||||
HostTaskCancellationResponse,
|
||||
HostTaskSubmissionResponse,
|
||||
LeaseRenewalResponse,
|
||||
TerminalResultResponse,
|
||||
)
|
||||
from host_agent.progress import TaskProgressSnapshot
|
||||
|
||||
|
||||
class LocalHostAgentClient:
|
||||
"""In-process task broker used when the Host Agent runs without Cloud."""
|
||||
|
||||
def __init__(self, *, host_id: str, device_ids: Callable[[], list[str]]) -> None:
|
||||
self.host_id = host_id
|
||||
self._device_ids = device_ids
|
||||
self._queue: asyncio.Queue[tuple[str, str, str | None]] = asyncio.Queue()
|
||||
self._cancelled: set[str] = set()
|
||||
|
||||
async def submit_self_task(self, *, goal: str, device_id: str | None = None) -> HostTaskSubmissionResponse:
|
||||
task_id = f"local-{uuid.uuid4().hex}"
|
||||
await self._queue.put((task_id, goal, device_id))
|
||||
return HostTaskSubmissionResponse(task_id=task_id)
|
||||
|
||||
async def claim(self):
|
||||
task_id, goal, requested_device = await self._queue.get()
|
||||
devices = self._device_ids()
|
||||
device_id = requested_device or (devices[0] if devices else "")
|
||||
if not device_id:
|
||||
return None
|
||||
from cloud.internal_api.models import AssignmentModel
|
||||
|
||||
return AssignmentModel(
|
||||
task_id=task_id,
|
||||
attempt=1,
|
||||
lease_id=f"local-lease-{uuid.uuid4().hex}",
|
||||
lease_expires_at=datetime.now(UTC) + timedelta(days=3650),
|
||||
host_id=self.host_id,
|
||||
device_id=device_id,
|
||||
goal=goal,
|
||||
)
|
||||
|
||||
async def heartbeat(self, *args, **kwargs) -> HeartbeatResponse:
|
||||
return HeartbeatResponse(
|
||||
host_id=self.host_id,
|
||||
accepted_devices=len(self._device_ids()),
|
||||
received_at=datetime.now(UTC),
|
||||
)
|
||||
|
||||
async def renew(self, assignment, *, progress: TaskProgressSnapshot | None = None) -> LeaseRenewalResponse:
|
||||
return LeaseRenewalResponse(
|
||||
status="renewed",
|
||||
lease_expires_at=datetime.now(UTC) + timedelta(days=3650),
|
||||
)
|
||||
|
||||
async def report_result(self, assignment, *, status: str, failure_reason: str | None = None, result: dict | None = None) -> TerminalResultResponse:
|
||||
return TerminalResultResponse(status="recorded")
|
||||
|
||||
async def cancel_task(self, task_id: str) -> HostTaskCancellationResponse:
|
||||
self._cancelled.add(task_id)
|
||||
return HostTaskCancellationResponse(task_id=task_id, status="cancel_requested")
|
||||
|
||||
async def aclose(self) -> None:
|
||||
return None
|
||||
@@ -0,0 +1,155 @@
|
||||
"""Per-device MCP session-level busy tracker.
|
||||
|
||||
The cloud-side assignment path and the MCP-driven path both drive devices
|
||||
through the same in-process ``DeviceManager``. This tracker records which
|
||||
devices are currently held by an MCP session so that:
|
||||
|
||||
- MCP tool calls against a device held by another session (or by a cloud
|
||||
assignment — checked separately by the caller via ``AgentStatusTracker``)
|
||||
can fail fast with a busy error.
|
||||
- The heartbeat payload can advertise ``mcp_busy_device_ids`` so the cloud
|
||||
scheduler won't dispatch conflicting assignments to the same device.
|
||||
|
||||
Leases expire ``ttl_seconds`` after the last ``renew()`` call (set on every
|
||||
tool call from the holding session). Expired leases are lazy-swept on read.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from datetime import UTC, datetime
|
||||
from threading import Lock
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Callable
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class McpDeviceLease:
|
||||
device_id: str
|
||||
session_id: str
|
||||
acquired_at: datetime
|
||||
last_seen_at: datetime
|
||||
|
||||
|
||||
class McpBusyTracker:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
ttl_seconds: float = 20.0,
|
||||
now: Callable[[], datetime] | None = None,
|
||||
) -> None:
|
||||
self._ttl = float(ttl_seconds)
|
||||
self._now = now or (lambda: datetime.now(UTC))
|
||||
self._lock = Lock()
|
||||
# device_id -> McpDeviceLease
|
||||
self._leases: dict[str, McpDeviceLease] = {}
|
||||
|
||||
def acquire(self, device_id: str, session_id: str) -> bool:
|
||||
with self._lock:
|
||||
self._sweep_locked()
|
||||
existing = self._leases.get(device_id)
|
||||
if existing is not None and existing.session_id != session_id:
|
||||
return False
|
||||
now = self._now()
|
||||
lease = McpDeviceLease(
|
||||
device_id=device_id,
|
||||
session_id=session_id,
|
||||
acquired_at=(existing.acquired_at if existing is not None else now),
|
||||
last_seen_at=now,
|
||||
)
|
||||
self._leases[device_id] = lease
|
||||
return True
|
||||
|
||||
def renew(self, device_id: str, session_id: str) -> bool:
|
||||
with self._lock:
|
||||
self._sweep_locked()
|
||||
existing = self._leases.get(device_id)
|
||||
# Tolerate boundary: lease may have been swept, but if the caller
|
||||
# is the legitimate previous holder, re-acquire on their behalf.
|
||||
if existing is None:
|
||||
now = self._now()
|
||||
self._leases[device_id] = McpDeviceLease(
|
||||
device_id=device_id,
|
||||
session_id=session_id,
|
||||
acquired_at=now,
|
||||
last_seen_at=now,
|
||||
)
|
||||
return True
|
||||
if existing.session_id != session_id:
|
||||
return False
|
||||
self._leases[device_id] = McpDeviceLease(
|
||||
device_id=device_id,
|
||||
session_id=session_id,
|
||||
acquired_at=existing.acquired_at,
|
||||
last_seen_at=self._now(),
|
||||
)
|
||||
return True
|
||||
|
||||
def release(self, session_id: str) -> list[str]:
|
||||
with self._lock:
|
||||
freed = [
|
||||
device_id
|
||||
for device_id, lease in self._leases.items()
|
||||
if lease.session_id == session_id
|
||||
]
|
||||
for device_id in freed:
|
||||
del self._leases[device_id]
|
||||
return freed
|
||||
|
||||
def release_device(self, device_id: str, session_id: str) -> bool:
|
||||
with self._lock:
|
||||
existing = self._leases.get(device_id)
|
||||
if existing is None or existing.session_id != session_id:
|
||||
return False
|
||||
del self._leases[device_id]
|
||||
return True
|
||||
|
||||
def busy_device_ids(self) -> list[str]:
|
||||
with self._lock:
|
||||
self._sweep_locked()
|
||||
return sorted(self._leases)
|
||||
|
||||
def snapshot(self) -> list[McpDeviceLease]:
|
||||
with self._lock:
|
||||
self._sweep_locked()
|
||||
return sorted(self._leases.values(), key=lambda lease: lease.device_id)
|
||||
|
||||
def wait_until_usable(
|
||||
self,
|
||||
device_id: str,
|
||||
session_id: str,
|
||||
*,
|
||||
timeout: float,
|
||||
poll_interval: float = 1.0,
|
||||
cloud_busy_check: Callable[[], bool] | None = None,
|
||||
) -> bool:
|
||||
"""Block until ``device_id`` is acquirable by ``session_id`` or timeout.
|
||||
|
||||
Reserved capability. MVP callers use try-acquire (``acquire`` -> False
|
||||
means busy). This method exists for future wiring where the cloud
|
||||
assignment path or an explicit MCP tool may opt to wait.
|
||||
"""
|
||||
deadline = time.monotonic() + timeout
|
||||
while True:
|
||||
cloud_busy = cloud_busy_check() if cloud_busy_check else False
|
||||
if not cloud_busy:
|
||||
if self.acquire(device_id, session_id):
|
||||
return True
|
||||
if time.monotonic() >= deadline:
|
||||
return False
|
||||
remaining = deadline - time.monotonic()
|
||||
time.sleep(max(0.0, min(poll_interval, remaining)))
|
||||
|
||||
def _sweep_locked(self) -> None:
|
||||
"""Caller holds ``self._lock``. Drops leases past their TTL."""
|
||||
cutoff = self._now()
|
||||
expired = [
|
||||
device_id
|
||||
for device_id, lease in self._leases.items()
|
||||
if (cutoff - lease.last_seen_at).total_seconds() > self._ttl
|
||||
]
|
||||
for device_id in expired:
|
||||
del self._leases[device_id]
|
||||
@@ -0,0 +1,116 @@
|
||||
"""Bearer-token persistence for the host-agent MCP server.
|
||||
|
||||
The token is generated on first start and persisted to a JSON file with
|
||||
0o600 permissions (POSIX) alongside the host identity. Rotation = delete
|
||||
the file and restart host-agent.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
import secrets
|
||||
import tempfile
|
||||
from dataclasses import dataclass
|
||||
from datetime import UTC, datetime
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Callable
|
||||
|
||||
|
||||
_TOKEN_BYTES = 32
|
||||
|
||||
|
||||
class McpTokenStoreError(RuntimeError):
|
||||
"""Raised when the MCP token file cannot be read or written."""
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class McpToken:
|
||||
version: int
|
||||
token: str
|
||||
created_at: datetime
|
||||
|
||||
|
||||
class McpTokenStore:
|
||||
def __init__(
|
||||
self,
|
||||
path: Path,
|
||||
*,
|
||||
now: Callable[[], datetime] | None = None,
|
||||
) -> None:
|
||||
self._path = Path(path)
|
||||
self._now = now or (lambda: datetime.now(UTC))
|
||||
|
||||
def load_or_create(self) -> McpToken:
|
||||
if self._path.exists():
|
||||
return self._read_existing()
|
||||
return self._generate_and_write()
|
||||
|
||||
def verify(self, presented: str) -> bool:
|
||||
try:
|
||||
token = self.load_or_create()
|
||||
except McpTokenStoreError:
|
||||
return False
|
||||
import hmac
|
||||
|
||||
return hmac.compare_digest(token.token, presented)
|
||||
|
||||
def _read_existing(self) -> McpToken:
|
||||
try:
|
||||
data = json.loads(self._path.read_text())
|
||||
except (OSError, json.JSONDecodeError) as exc:
|
||||
raise McpTokenStoreError(
|
||||
f"cannot read MCP token file {self._path}: {exc}"
|
||||
) from exc
|
||||
if not isinstance(data, dict):
|
||||
raise McpTokenStoreError("MCP token file is not a JSON object")
|
||||
try:
|
||||
return McpToken(
|
||||
version=int(data["version"]),
|
||||
token=str(data["token"]),
|
||||
created_at=datetime.fromisoformat(str(data["created_at"])),
|
||||
)
|
||||
except (KeyError, TypeError, ValueError) as exc:
|
||||
raise McpTokenStoreError(f"MCP token file schema invalid: {exc}") from exc
|
||||
|
||||
def _generate_and_write(self) -> McpToken:
|
||||
token = McpToken(
|
||||
version=1,
|
||||
token=secrets.token_urlsafe(_TOKEN_BYTES),
|
||||
created_at=self._now(),
|
||||
)
|
||||
payload = {
|
||||
"version": token.version,
|
||||
"token": token.token,
|
||||
"created_at": token.created_at.isoformat(),
|
||||
}
|
||||
try:
|
||||
self._atomic_write(json.dumps(payload, indent=2))
|
||||
except OSError as exc:
|
||||
raise McpTokenStoreError(
|
||||
f"cannot write MCP token file {self._path}: {exc}"
|
||||
) from exc
|
||||
return token
|
||||
|
||||
def _atomic_write(self, content: str) -> None:
|
||||
self._path.parent.mkdir(parents=True, exist_ok=True)
|
||||
# Atomic on POSIX; on Windows os.replace is also atomic per docs.
|
||||
fd, tmp_name = tempfile.mkstemp(
|
||||
prefix=".host_mcp_token.",
|
||||
suffix=".tmp",
|
||||
dir=str(self._path.parent),
|
||||
)
|
||||
try:
|
||||
with os.fdopen(fd, "w", encoding="utf-8") as fh:
|
||||
fh.write(content)
|
||||
os.chmod(tmp_name, 0o600)
|
||||
os.replace(tmp_name, self._path)
|
||||
except BaseException:
|
||||
try:
|
||||
os.unlink(tmp_name)
|
||||
except OSError:
|
||||
pass
|
||||
raise
|
||||
@@ -5,13 +5,15 @@ import base64
|
||||
import json
|
||||
from collections.abc import Awaitable, Callable
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import jinja2
|
||||
from fastapi import Depends, FastAPI, HTTPException, Request
|
||||
from fastapi.responses import HTMLResponse, JSONResponse, RedirectResponse, Response
|
||||
|
||||
from core.errors import DeviceNotFoundError, DeviceOfflineError, DeviceRuntimeError
|
||||
from device.manager import DeviceManager
|
||||
from driver.registry import build_driver_factory
|
||||
from host_agent.assignment import AssignmentExecutor
|
||||
from host_agent.client import (
|
||||
HostAgentClient,
|
||||
@@ -20,10 +22,15 @@ from host_agent.client import (
|
||||
HostTaskSubmissionUnknownError,
|
||||
)
|
||||
from host_agent.config import HostAgentConfig
|
||||
from host_agent.conversation import ConversationAgent
|
||||
from host_agent.conversation_log import ConversationLogStore
|
||||
from host_agent.devices import register_local_device, unregister_local_device
|
||||
from host_agent.history import ConsoleHistoryStore
|
||||
from host_agent.identity import HostIdentityStore
|
||||
from host_agent.ios_discovery import IOSDiscoveryError, discover_connected_ios_devices
|
||||
from host_agent.local_account import LocalAccountStore
|
||||
from host_agent.mcp_lock import McpBusyTracker
|
||||
from host_agent.mcp_token import McpTokenStore
|
||||
from host_agent.status import AgentStatusTracker
|
||||
from host_agent.web.auth import (
|
||||
SessionManager,
|
||||
@@ -31,7 +38,11 @@ from host_agent.web.auth import (
|
||||
attempt_login,
|
||||
change_password,
|
||||
)
|
||||
from host_agent.web.mcp_auth import BearerAuthMiddleware
|
||||
from storage.device_config import DeviceConfigStore
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from mcp.server.fastmcp import FastMCP
|
||||
from storage.task_metadata import TaskMetadataStore
|
||||
from storage.timeline import Timeline
|
||||
|
||||
@@ -223,6 +234,11 @@ def create_console_app(
|
||||
metadata_store: TaskMetadataStore | None = None,
|
||||
timeline: Timeline | None = None,
|
||||
executor: AssignmentExecutor | None = None,
|
||||
mcp_server: FastMCP | None = None,
|
||||
mcp_token_store: McpTokenStore | None = None,
|
||||
mcp_busy_tracker: McpBusyTracker | None = None,
|
||||
conversation_agent: ConversationAgent | None = None,
|
||||
conversation_log: ConversationLogStore | None = None,
|
||||
) -> FastAPI:
|
||||
app = FastAPI(title="Host Agent Console")
|
||||
cookie_secure = config.console_bind_host not in _LOOPBACK_BIND_HOSTS
|
||||
@@ -231,6 +247,18 @@ def create_console_app(
|
||||
if cancel_task is None and host_client is not None:
|
||||
cancel_task = host_client.cancel_task
|
||||
submission_available = submit_self_task is not None
|
||||
mcp_mounted = mcp_server is not None and mcp_token_store is not None
|
||||
if mcp_mounted:
|
||||
from starlette.applications import Starlette
|
||||
from starlette.middleware import Middleware
|
||||
|
||||
mcp_asgi = mcp_server.streamable_http_app()
|
||||
authed = Starlette(
|
||||
routes=[],
|
||||
middleware=[Middleware(BearerAuthMiddleware, token_store=mcp_token_store)],
|
||||
)
|
||||
authed.router.mount("/", mcp_asgi)
|
||||
app.mount("/mcp", authed)
|
||||
|
||||
def _running_devices() -> list[dict[str, str]]:
|
||||
return [
|
||||
@@ -340,6 +368,10 @@ def create_console_app(
|
||||
for d in manager.list_devices()
|
||||
]
|
||||
texts = _dashboard_texts(snapshot=snapshot)
|
||||
mcp_endpoint = "/mcp" if mcp_mounted else None
|
||||
mcp_busy_devices = (
|
||||
mcp_busy_tracker.busy_device_ids() if mcp_busy_tracker is not None else []
|
||||
)
|
||||
return _render(
|
||||
"dashboard.html",
|
||||
title="Status",
|
||||
@@ -347,6 +379,8 @@ def create_console_app(
|
||||
identity=identity,
|
||||
devices=devices,
|
||||
config=config,
|
||||
mcp_endpoint=mcp_endpoint,
|
||||
mcp_busy_devices=mcp_busy_devices,
|
||||
**texts,
|
||||
)
|
||||
|
||||
@@ -375,7 +409,63 @@ def create_console_app(
|
||||
}
|
||||
for device in manager.list_devices()
|
||||
]
|
||||
return JSONResponse({"status": snapshot, "devices": devices})
|
||||
return JSONResponse(
|
||||
{
|
||||
"status": snapshot,
|
||||
"devices": devices,
|
||||
"mcp_endpoint": "/mcp" if mcp_mounted else None,
|
||||
"mcp_busy_devices": (
|
||||
mcp_busy_tracker.busy_device_ids()
|
||||
if mcp_busy_tracker is not None
|
||||
else []
|
||||
),
|
||||
}
|
||||
)
|
||||
|
||||
@app.post("/api/chat")
|
||||
async def api_chat(
|
||||
request: Request,
|
||||
session: SessionState = Depends(require_csrf),
|
||||
) -> JSONResponse:
|
||||
if conversation_agent is None:
|
||||
raise HTTPException(status_code=503, detail="chat agent is not configured")
|
||||
payload = await request.json()
|
||||
device_id = payload.get("device_id") if isinstance(payload, dict) else None
|
||||
if not isinstance(device_id, str) or not device_id.strip():
|
||||
raise HTTPException(status_code=400, detail="device_id is required")
|
||||
if device_id not in {device.id for device in manager.list_devices()}:
|
||||
raise HTTPException(status_code=404, detail="unknown device")
|
||||
raw_messages = payload.get("messages") if isinstance(payload, dict) else None
|
||||
if not isinstance(raw_messages, list) or not raw_messages:
|
||||
raise HTTPException(status_code=400, detail="messages must be a non-empty list")
|
||||
messages = [
|
||||
{key: value for key, value in item.items() if key in {"role", "content", "image_base64", "mime_type", "text"}}
|
||||
for item in raw_messages
|
||||
if isinstance(item, dict)
|
||||
and item.get("role") in {"user", "assistant"}
|
||||
and (isinstance(item.get("content"), (str, list)) or isinstance(item.get("image_base64"), str))
|
||||
]
|
||||
if not messages:
|
||||
raise HTTPException(status_code=400, detail="messages are invalid")
|
||||
if conversation_log is not None:
|
||||
await asyncio.to_thread(
|
||||
conversation_log.append,
|
||||
{"type": "user_request", "device_id": device_id.strip(), "messages": messages},
|
||||
)
|
||||
try:
|
||||
result = await asyncio.to_thread(
|
||||
conversation_agent.chat_for_device, device_id.strip(), messages
|
||||
)
|
||||
except Exception as exc:
|
||||
if conversation_log is not None:
|
||||
await asyncio.to_thread(
|
||||
conversation_log.append,
|
||||
{"type": "agent_error", "device_id": device_id.strip(), "error": str(exc)},
|
||||
)
|
||||
raise HTTPException(status_code=502, detail=str(exc)) from exc
|
||||
return JSONResponse(
|
||||
{"content": result.content, "tool_calls": result.tool_calls}
|
||||
)
|
||||
|
||||
@app.get("/devices", response_class=HTMLResponse)
|
||||
async def devices_page(
|
||||
@@ -401,6 +491,124 @@ def create_console_app(
|
||||
error=None,
|
||||
)
|
||||
|
||||
@app.post("/api/devices/{device_id}/screenshot")
|
||||
async def api_device_screenshot(
|
||||
device_id: str,
|
||||
session: SessionState = Depends(require_csrf),
|
||||
) -> Response:
|
||||
"""Capture one on-demand screenshot for a connected local device."""
|
||||
try:
|
||||
screenshot = await asyncio.to_thread(
|
||||
lambda: manager.active_driver(device_id).screenshot()
|
||||
)
|
||||
except DeviceNotFoundError as exc:
|
||||
raise HTTPException(status_code=404, detail=str(exc)) from exc
|
||||
except DeviceOfflineError as exc:
|
||||
raise HTTPException(status_code=503, detail=str(exc)) from exc
|
||||
except DeviceRuntimeError as exc:
|
||||
raise HTTPException(status_code=502, detail=str(exc)) from exc
|
||||
except Exception as exc:
|
||||
raise HTTPException(
|
||||
status_code=502,
|
||||
detail=str(exc) or "failed to capture device screenshot",
|
||||
) from exc
|
||||
|
||||
if not isinstance(screenshot, bytes) or not screenshot:
|
||||
raise HTTPException(
|
||||
status_code=502,
|
||||
detail="device returned an empty screenshot",
|
||||
)
|
||||
return Response(
|
||||
content=screenshot,
|
||||
media_type="image/png",
|
||||
headers={
|
||||
"Cache-Control": "no-store",
|
||||
"Pragma": "no-cache",
|
||||
"X-Content-Type-Options": "nosniff",
|
||||
},
|
||||
)
|
||||
|
||||
@app.get("/api/devices/discover-ios")
|
||||
async def api_discover_ios_devices(
|
||||
session: SessionState = Depends(require_session),
|
||||
) -> JSONResponse:
|
||||
try:
|
||||
discovered = await asyncio.to_thread(discover_connected_ios_devices)
|
||||
except IOSDiscoveryError as exc:
|
||||
raise HTTPException(status_code=502, detail=str(exc)) from exc
|
||||
|
||||
configured = await asyncio.to_thread(config_store.list)
|
||||
configured_udids = {
|
||||
str(record["connection_info"].get("udid"))
|
||||
for record in configured
|
||||
if record["driver_type"] == "wda"
|
||||
}
|
||||
used_wda_ports = {
|
||||
record["connection_info"].get("wda_local_port") for record in configured
|
||||
}
|
||||
used_mjpeg_ports = {
|
||||
record["connection_info"].get("mjpegServerPort") for record in configured
|
||||
}
|
||||
next_wda_port = 8100
|
||||
next_mjpeg_port = 9100
|
||||
result = []
|
||||
for device in discovered:
|
||||
while next_wda_port in used_wda_ports:
|
||||
next_wda_port += 1
|
||||
while next_mjpeg_port in used_mjpeg_ports:
|
||||
next_mjpeg_port += 1
|
||||
result.append(
|
||||
{
|
||||
**device,
|
||||
"configured": device["udid"] in configured_udids,
|
||||
"suggested_wda_port": next_wda_port,
|
||||
"suggested_mjpeg_port": next_mjpeg_port,
|
||||
}
|
||||
)
|
||||
used_wda_ports.add(next_wda_port)
|
||||
used_mjpeg_ports.add(next_mjpeg_port)
|
||||
next_wda_port += 1
|
||||
next_mjpeg_port += 1
|
||||
return JSONResponse({"devices": result})
|
||||
|
||||
@app.post("/api/devices/test-connection")
|
||||
async def api_device_test_connection(
|
||||
request: Request,
|
||||
session: SessionState = Depends(require_csrf),
|
||||
) -> JSONResponse:
|
||||
payload = await request.json()
|
||||
if not isinstance(payload, dict):
|
||||
raise HTTPException(status_code=400, detail="Request must be an object.")
|
||||
driver_type = str(payload.get("driver_type", "")).strip()
|
||||
connection_info = payload.get("connection_info")
|
||||
if not driver_type or not isinstance(connection_info, dict):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="Driver type and connection info are required.",
|
||||
)
|
||||
|
||||
driver = None
|
||||
connected = False
|
||||
try:
|
||||
driver = build_driver_factory(driver_type, connection_info)()
|
||||
await asyncio.to_thread(driver.connect)
|
||||
connected = True
|
||||
await asyncio.to_thread(driver.health_check)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
except Exception as exc:
|
||||
raise HTTPException(
|
||||
status_code=502,
|
||||
detail=str(exc) or "device connection test failed",
|
||||
) from exc
|
||||
finally:
|
||||
if connected and driver is not None:
|
||||
try:
|
||||
await asyncio.to_thread(driver.disconnect)
|
||||
except Exception:
|
||||
pass
|
||||
return JSONResponse({"ok": True, "driver_type": driver_type})
|
||||
|
||||
@app.post("/devices/save")
|
||||
async def devices_save(
|
||||
request: Request,
|
||||
@@ -427,6 +635,26 @@ def create_console_app(
|
||||
else:
|
||||
connection_info = parsed
|
||||
|
||||
if error is None:
|
||||
# The console exposes the common Appium settings as regular form
|
||||
# fields. Advanced JSON remains available for uncommon capabilities.
|
||||
field_map = {
|
||||
"server_url": "server_url",
|
||||
"udid": "udid",
|
||||
"device_name": "device_name",
|
||||
}
|
||||
for form_key, config_key in field_map.items():
|
||||
value = str(form.get(form_key, "")).strip()
|
||||
if value:
|
||||
connection_info[config_key] = value
|
||||
port_key = "wda_local_port" if driver_type == "wda" else "system_port"
|
||||
port_value = str(form.get(port_key, "")).strip()
|
||||
if port_value:
|
||||
try:
|
||||
connection_info[port_key] = int(port_value)
|
||||
except ValueError:
|
||||
error = f"{port_key} must be an integer."
|
||||
|
||||
if error is None:
|
||||
try:
|
||||
await asyncio.to_thread(
|
||||
@@ -540,6 +768,15 @@ def create_console_app(
|
||||
entries=entries,
|
||||
)
|
||||
|
||||
@app.get("/conversations", response_class=HTMLResponse)
|
||||
async def conversations_page(
|
||||
session: SessionState = Depends(require_session),
|
||||
) -> HTMLResponse:
|
||||
events = await asyncio.to_thread(
|
||||
conversation_log.list_recent if conversation_log is not None else (lambda: [])
|
||||
)
|
||||
return _render("conversations.html", title="Conversations", session=session, events=events)
|
||||
|
||||
def _tasks_list_context(
|
||||
session: SessionState,
|
||||
*,
|
||||
|
||||
@@ -0,0 +1,255 @@
|
||||
"""FastMCP server builder for the host-agent MCP endpoint.
|
||||
|
||||
Wraps ``api.mcp.tool_handlers(manager=...)`` with:
|
||||
|
||||
- Cloud-busy and MCP-busy checks (per-device, fail-fast on conflict).
|
||||
- Lazy session-level device lock acquire / renew.
|
||||
- Display-status mapping for ``list_devices`` / ``device_status`` so
|
||||
connected-but-idle devices don't appear "busy" (which they do at the
|
||||
``DeviceManager`` layer because an Appium/WDA session is open).
|
||||
|
||||
The builder returns a ``FastMCP`` instance. The caller
|
||||
(``create_console_app``) is responsible for wrapping it in
|
||||
``BearerAuthMiddleware`` and mounting at ``/mcp``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextvars
|
||||
from collections.abc import Callable
|
||||
from typing import Any
|
||||
|
||||
from device.manager import DeviceManager
|
||||
from host_agent.mcp_lock import McpBusyTracker
|
||||
from host_agent.status import AgentStatusTracker
|
||||
from mcp.server.fastmcp import Context, FastMCP
|
||||
|
||||
# Tool names that don't target a specific device — skip busy check.
|
||||
_NON_DEVICE_TOOLS = frozenset({"list_devices", "device_status"})
|
||||
# Tools that report status and should use the display-status mapping.
|
||||
_STATUS_TOOLS = frozenset({"list_devices", "device_status"})
|
||||
# Name of the wrapper kwarg FastMCP injects the live ``Context`` into.
|
||||
# We set ``tool.context_kwarg = _CONTEXT_KWARG`` after swapping the tool's
|
||||
# ``fn`` (see ``build_mcp_server``) so FastMCP passes ``ctx`` into our
|
||||
# wrapper alongside the validated arguments.
|
||||
_CONTEXT_KWARG = "ctx"
|
||||
|
||||
|
||||
class McpDeviceBusyError(Exception):
|
||||
"""Raised by the wrapper when the target device is held by the cloud
|
||||
assignment path or another MCP session."""
|
||||
|
||||
def __init__(self, device_id: str, busy_owner: str) -> None:
|
||||
super().__init__(f"device {device_id} is busy (held by {busy_owner})")
|
||||
self.device_id = device_id
|
||||
self.busy_owner = busy_owner
|
||||
|
||||
|
||||
class FastMcpSdkIncompatibilityError(RuntimeError):
|
||||
"""Raised when the FastMCP SDK layout diverges from what this module
|
||||
expects (e.g. ``Tool.fn`` rename or ``Tool.context_kwarg`` removal)."""
|
||||
|
||||
|
||||
# contextvars fallback used by tests and any call that originates outside a
|
||||
# live FastMCP request lifecycle. Production handlers run inside an MCP
|
||||
# request whose context exposes ``request_id`` and the underlying
|
||||
# ``session``; ``_current_session_id`` reads from that context first and
|
||||
# falls back to this ContextVar.
|
||||
_TEST_SESSION_ID: contextvars.ContextVar[str] = contextvars.ContextVar(
|
||||
"_TEST_SESSION_ID", default=""
|
||||
)
|
||||
|
||||
|
||||
def _current_session_id(ctx: Context | None = None) -> str:
|
||||
"""Extract a stable per-MCP-session identifier from the live context.
|
||||
|
||||
The mcp SDK 1.28.1 ``Context`` exposes ``session`` (a long-lived
|
||||
``ServerSession`` instance per Streamable HTTP session). Its Python
|
||||
object identity (``id(ctx.session)``) is stable across every tool call
|
||||
the same client makes within that session, which is exactly the
|
||||
identity the busy tracker needs to renew leases.
|
||||
|
||||
Falls back to ``_TEST_SESSION_ID`` when no Context is supplied (i.e.
|
||||
when invoked outside a FastMCP request lifecycle, as ``_call_tool_sync``
|
||||
does in tests).
|
||||
"""
|
||||
if ctx is not None:
|
||||
session_obj = getattr(ctx, "session", None)
|
||||
if session_obj is not None:
|
||||
return f"mcp_session:{id(session_obj)}"
|
||||
return _TEST_SESSION_ID.get("")
|
||||
|
||||
|
||||
def build_mcp_server(
|
||||
*,
|
||||
manager: DeviceManager,
|
||||
mcp_busy_tracker: McpBusyTracker,
|
||||
status_tracker: AgentStatusTracker,
|
||||
) -> FastMCP:
|
||||
"""Construct the FastMCP server wrapping ``tool_handlers``."""
|
||||
# Imported lazily to keep the package import graph flat.
|
||||
from api.mcp import tool_handlers
|
||||
|
||||
handlers = tool_handlers(manager=manager)
|
||||
server = FastMCP("apex-host-agent")
|
||||
|
||||
for tool_name, raw_handler in handlers.items():
|
||||
wrapped = _wrap_tool(
|
||||
tool_name,
|
||||
raw_handler,
|
||||
mcp_busy_tracker=mcp_busy_tracker,
|
||||
status_tracker=status_tracker,
|
||||
)
|
||||
# Register the raw handler so FastMCP captures its signature (the
|
||||
# MCP wire schema is derived from the function signature). Then
|
||||
# swap ``tool.fn`` for our busy-check / status-mapping wrapper.
|
||||
# Using ``*args, **kwargs`` directly breaks the schema, so we have
|
||||
# to keep the signature and only replace the underlying callable.
|
||||
server._tool_manager.add_tool( # type: ignore[attr-defined]
|
||||
raw_handler, name=tool_name
|
||||
)
|
||||
try:
|
||||
tool = server._tool_manager._tools[tool_name] # type: ignore[attr-defined]
|
||||
tool.fn = wrapped
|
||||
# FastMCP injects the live Context into the kwarg named by
|
||||
# ``tool.context_kwarg``. The raw handler doesn't declare one,
|
||||
# so the cached value is None; we override it so the wrapper
|
||||
# receives the Context via its ``ctx`` kwarg.
|
||||
tool.context_kwarg = _CONTEXT_KWARG
|
||||
except AttributeError as exc:
|
||||
raise FastMcpSdkIncompatibilityError(
|
||||
"FastMCP SDK layout changed: cannot swap Tool.fn or set "
|
||||
f"context_kwarg (tool={tool_name!r}). Underlying error: {exc}"
|
||||
) from exc
|
||||
|
||||
return server
|
||||
|
||||
|
||||
def _wrap_tool(
|
||||
tool_name: str,
|
||||
handler: Callable[..., Any],
|
||||
*,
|
||||
mcp_busy_tracker: McpBusyTracker,
|
||||
status_tracker: AgentStatusTracker,
|
||||
) -> Callable[..., Any]:
|
||||
def wrapped(*args: Any, **kwargs: Any) -> Any:
|
||||
ctx = kwargs.pop(_CONTEXT_KWARG, None)
|
||||
session_id = _current_session_id(ctx)
|
||||
device_id = kwargs.get("device_id")
|
||||
|
||||
if tool_name in _STATUS_TOOLS:
|
||||
return _with_display_status(handler, status_tracker, *args, **kwargs)
|
||||
|
||||
if device_id is not None and tool_name not in _NON_DEVICE_TOOLS:
|
||||
_check_and_acquire(device_id, session_id, mcp_busy_tracker, status_tracker)
|
||||
|
||||
return handler(*args, **kwargs)
|
||||
|
||||
return wrapped
|
||||
|
||||
|
||||
def _check_and_acquire(
|
||||
device_id: str,
|
||||
session_id: str,
|
||||
mcp_busy_tracker: McpBusyTracker,
|
||||
status_tracker: AgentStatusTracker,
|
||||
) -> None:
|
||||
cloud_busy = _cloud_busy_device_id(status_tracker)
|
||||
if cloud_busy == device_id:
|
||||
raise McpDeviceBusyError(device_id, "cloud_assignment")
|
||||
if device_id in mcp_busy_tracker.busy_device_ids():
|
||||
existing = next(
|
||||
(
|
||||
lease
|
||||
for lease in mcp_busy_tracker.snapshot()
|
||||
if lease.device_id == device_id
|
||||
),
|
||||
None,
|
||||
)
|
||||
if existing is not None and existing.session_id != session_id:
|
||||
prefix = existing.session_id[:8]
|
||||
raise McpDeviceBusyError(device_id, f"mcp_session:{prefix}")
|
||||
if not mcp_busy_tracker.acquire(device_id, session_id):
|
||||
# Race: someone else got it between check and acquire.
|
||||
raise McpDeviceBusyError(device_id, "another_session")
|
||||
mcp_busy_tracker.renew(device_id, session_id)
|
||||
|
||||
|
||||
def _cloud_busy_device_id(status_tracker: AgentStatusTracker) -> str | None:
|
||||
"""Return the device_id currently bound to the cloud assignment, if any."""
|
||||
snap = status_tracker.snapshot()
|
||||
current = snap.get("current_assignment")
|
||||
if not isinstance(current, dict):
|
||||
return None
|
||||
device_id = current.get("device_id")
|
||||
return device_id if isinstance(device_id, str) else None
|
||||
|
||||
|
||||
def _with_display_status(
|
||||
handler: Callable[..., Any],
|
||||
status_tracker: AgentStatusTracker,
|
||||
*args: Any,
|
||||
**kwargs: Any,
|
||||
) -> Any:
|
||||
busy_device_id = _cloud_busy_device_id(status_tracker)
|
||||
result = handler(*args, **kwargs)
|
||||
if isinstance(result, list):
|
||||
for item in result:
|
||||
if isinstance(item, dict) and "status" in item:
|
||||
item["status"] = _display_status(
|
||||
item["status"], item.get("id"), busy_device_id
|
||||
)
|
||||
return result
|
||||
if isinstance(result, dict) and "status" in result:
|
||||
result["status"] = _display_status(
|
||||
result["status"], result.get("device_id"), busy_device_id
|
||||
)
|
||||
return result
|
||||
|
||||
|
||||
def _display_status(raw: str, device_id: Any, busy_device_id: str | None) -> str:
|
||||
"""Mirror ``host_agent.web.app._device_display_status`` semantics.
|
||||
|
||||
A device that's locally "busy" because it's connected-but-idle reports
|
||||
"connected" instead, unless it's the device currently running a cloud
|
||||
assignment (in which case "busy" is the truthful status).
|
||||
"""
|
||||
if raw == "busy" and device_id != busy_device_id:
|
||||
return "connected"
|
||||
return raw
|
||||
|
||||
|
||||
def _call_tool_sync(
|
||||
server: FastMCP,
|
||||
tool_name: str,
|
||||
arguments: dict[str, Any],
|
||||
*,
|
||||
session_id: str,
|
||||
) -> Any:
|
||||
"""Test helper: invoke a registered tool synchronously with a forced
|
||||
``session_id``. Bypasses the HTTP/MCP transport layer (and the live
|
||||
FastMCP Context) so tests don't need an MCP client.
|
||||
|
||||
Walks FastMCP's tool registry (``_tool_manager._tools[tool_name].fn``) —
|
||||
the exact attribute path follows mcp SDK 1.28.1's
|
||||
``ToolManager._tools`` layout.
|
||||
"""
|
||||
token = _TEST_SESSION_ID.set(session_id)
|
||||
try:
|
||||
manager = getattr(server, "_tool_manager", None)
|
||||
if manager is None:
|
||||
raise KeyError(f"tool {tool_name!r} not registered (no tool manager)")
|
||||
registry = getattr(manager, "_tools", None) or getattr(manager, "tools", None)
|
||||
if isinstance(registry, dict):
|
||||
tool = registry.get(tool_name)
|
||||
else:
|
||||
tool = manager.get_tool(tool_name) # type: ignore[union-attr]
|
||||
if tool is None:
|
||||
raise KeyError(f"tool {tool_name!r} not registered")
|
||||
# FastMCP Tool wraps a callable; our wrappers are sync, so unwrap.
|
||||
fn = getattr(tool, "fn", None) or getattr(tool, "func", None)
|
||||
if fn is None:
|
||||
raise KeyError(f"tool {tool_name!r} has no callable")
|
||||
return fn(**arguments)
|
||||
finally:
|
||||
_TEST_SESSION_ID.reset(token)
|
||||
@@ -0,0 +1,36 @@
|
||||
"""Bearer-token auth middleware for the MCP sub-app.
|
||||
|
||||
Mounted on the FastMCP ``streamable_http_app()`` (NOT the console FastAPI),
|
||||
so cookie-session auth on console routes is unaffected.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from starlette.middleware.base import BaseHTTPMiddleware
|
||||
from starlette.requests import Request
|
||||
from starlette.responses import JSONResponse, Response
|
||||
|
||||
from host_agent.mcp_token import McpTokenStore
|
||||
|
||||
|
||||
class BearerAuthMiddleware(BaseHTTPMiddleware):
|
||||
def __init__(self, app, token_store: McpTokenStore) -> None:
|
||||
super().__init__(app)
|
||||
self._store = token_store
|
||||
|
||||
async def dispatch(self, request: Request, call_next) -> Response: # type: ignore[no-untyped-def]
|
||||
header = request.headers.get("Authorization")
|
||||
if not header or not header.lower().startswith("bearer "):
|
||||
return _unauthorized()
|
||||
presented = header.split(" ", 1)[1].strip()
|
||||
if not self._store.verify(presented):
|
||||
return _unauthorized()
|
||||
return await call_next(request)
|
||||
|
||||
|
||||
def _unauthorized() -> JSONResponse:
|
||||
return JSONResponse(
|
||||
status_code=401,
|
||||
content={"error": "invalid token"},
|
||||
headers={"WWW-Authenticate": "Bearer"},
|
||||
)
|
||||
@@ -26,6 +26,7 @@ form.inline { display: inline; margin: 0; }
|
||||
<a href="/tasks">Tasks</a>
|
||||
<a href="/account">Account</a>
|
||||
<a href="/history">History</a>
|
||||
<a href="/conversations">Conversations</a>
|
||||
<form class="inline" method="post" action="/logout">
|
||||
<input type="hidden" name="csrf_token" value="{{ session.csrf_token }}">
|
||||
<button type="submit">Logout</button>
|
||||
|
||||
@@ -0,0 +1,14 @@
|
||||
{% extends "base.html" %}
|
||||
{% block body %}
|
||||
<h1>Conversations</h1>
|
||||
{% if not events %}<p>No conversation activity recorded yet.</p>{% endif %}
|
||||
{% for event in events %}
|
||||
<article>
|
||||
<h2>{{ event["event_type"] }} <small>{{ event["occurred_at"] }}</small></h2>
|
||||
{% if event.get("content") %}<pre>{{ event["content"] }}</pre>{% endif %}
|
||||
{% if event.get("thinking") %}<details><summary>LLM reasoning</summary><pre>{{ event["thinking"] }}</pre></details>{% endif %}
|
||||
{% if event.get("tool_calls") %}<details open><summary>Tool calls</summary><pre>{{ event["tool_calls"] | tojson(indent=2) }}</pre></details>{% endif %}
|
||||
{% if event.get("tool_name") %}<p><strong>{{ event["tool_name"] }}</strong></p><pre>{{ event.get("arguments") | tojson(indent=2) }}</pre><pre>{{ event.get("result") | tojson(indent=2) }}</pre>{% endif %}
|
||||
</article>
|
||||
{% endfor %}
|
||||
{% endblock %}
|
||||
@@ -20,6 +20,24 @@
|
||||
<p id="current-assignment">{{ assignment_text }}</p>
|
||||
<p id="current-progress">{{ progress_text }}</p>
|
||||
</section>
|
||||
<section>
|
||||
<h2>MCP</h2>
|
||||
<table>
|
||||
<tbody>
|
||||
<tr>
|
||||
<td>MCP</td>
|
||||
<td>
|
||||
{% if mcp_endpoint %}
|
||||
endpoint <code>{{ mcp_endpoint }}</code>;
|
||||
{% if mcp_busy_devices %}busy: {{ mcp_busy_devices|join(", ") }}{% else %}idle{% endif %}
|
||||
{% else %}
|
||||
not configured
|
||||
{% endif %}
|
||||
</td>
|
||||
</tr>
|
||||
</tbody>
|
||||
</table>
|
||||
</section>
|
||||
<section>
|
||||
<h2>Devices</h2>
|
||||
<table>
|
||||
|
||||
@@ -5,13 +5,20 @@
|
||||
<p class="error">{{ error }}</p>
|
||||
{% endif %}
|
||||
<table>
|
||||
<thead><tr><th>ID</th><th>Name</th><th>Driver</th><th>Cloud ID</th><th></th></tr></thead>
|
||||
<thead><tr><th>ID</th><th>Name</th><th>Driver</th><th>Cloud ID</th><th>Screenshot</th><th></th></tr></thead>
|
||||
<tbody>{% for device in devices %}
|
||||
<tr>
|
||||
<td>{{ device["device_id"] }}</td>
|
||||
<td>{{ device["name"] or "" }}</td>
|
||||
<td>{{ device["driver_type"] }}</td>
|
||||
<td>{{ device["cloud_device_id"] or "" }}</td>
|
||||
<td>
|
||||
<button type="button" class="screenshot-button" data-device-id="{{ device["device_id"] }}">Get screenshot</button>
|
||||
<div class="screenshot-preview" data-screenshot-preview hidden>
|
||||
<p class="screenshot-status" data-screenshot-status></p>
|
||||
<img alt="Current screen for {{ device["device_id"] }}" data-screenshot-image hidden>
|
||||
</div>
|
||||
</td>
|
||||
<td>
|
||||
<a href="/devices?edit={{ device["device_id"] }}">Edit</a>
|
||||
<form class="inline" method="post" action="/devices/remove">
|
||||
@@ -24,14 +31,173 @@
|
||||
{% endfor %}</tbody>
|
||||
</table>
|
||||
<h2>{{ "Edit device" if edit_record else "Add device" }}</h2>
|
||||
<button type="button" id="discover-ios">Scan connected iPhones</button>
|
||||
<p id="discovery-status" class="screenshot-status"></p>
|
||||
<div id="discovered-devices"></div>
|
||||
<form method="post" action="/devices/save">
|
||||
<input type="hidden" name="csrf_token" value="{{ csrf_token }}">
|
||||
<label>Device ID <input type="text" name="device_id" value="{{ edit_record["device_id"] if edit_record else "" }}" required></label><br>
|
||||
<label>Name <input type="text" name="name" value="{{ edit_record["name"] if edit_record else "" }}"></label><br>
|
||||
<label>Driver type <input type="text" name="driver_type" value="{{ edit_record["driver_type"] if edit_record else "wda" }}" required></label><br>
|
||||
<label>Connection info (JSON)<br>
|
||||
<textarea name="connection_info" rows="3" cols="50">{{ connection_info_json }}</textarea>
|
||||
<label>Platform and protocol
|
||||
<select name="driver_type" id="driver-type" required>
|
||||
<option value="wda"{% if not edit_record or edit_record["driver_type"] == "wda" %} selected{% endif %}>iOS - Appium / XCUITest (WDA)</option>
|
||||
<option value="uiautomator2"{% if edit_record and edit_record["driver_type"] == "uiautomator2" %} selected{% endif %}>Android - Appium / UiAutomator2</option>
|
||||
</select>
|
||||
</label><br>
|
||||
<label>Appium server URL <input type="url" name="server_url" value="{{ edit_record["connection_info"].get("server_url", "http://127.0.0.1:4723") if edit_record else "http://127.0.0.1:4723" }}" required></label><br>
|
||||
<label>Device UDID <input type="text" name="udid" value="{{ edit_record["connection_info"].get("udid", "") if edit_record else "" }}" required></label><br>
|
||||
<label>Device name <input type="text" name="device_name" value="{{ edit_record["connection_info"].get("device_name", "") if edit_record else "" }}" placeholder="iPhone"></label><br>
|
||||
<label class="ios-setting">WDA local port <input type="number" min="1" max="65535" name="wda_local_port" value="{{ edit_record["connection_info"].get("wda_local_port", "") if edit_record else "" }}" placeholder="8100"></label>
|
||||
<label class="android-setting">UiAutomator2 system port <input type="number" min="1" max="65535" name="system_port" value="{{ edit_record["connection_info"].get("system_port", "") if edit_record else "" }}" placeholder="8200"></label><br>
|
||||
<details>
|
||||
<summary>Advanced connection capabilities (JSON)</summary>
|
||||
<textarea name="connection_info" rows="4" cols="60">{{ connection_info_json }}</textarea>
|
||||
</details>
|
||||
<p id="connection-test-status" class="screenshot-status"></p>
|
||||
<button type="button" id="test-connection">Test connection</button>
|
||||
<button type="submit">Save</button>
|
||||
</form>
|
||||
<style>
|
||||
.screenshot-preview { margin-top: 0.5rem; max-width: 260px; }
|
||||
.screenshot-preview img { display: block; width: 100%; height: auto; border: 1px solid #c8d0d6; }
|
||||
.screenshot-status { margin: 0 0 0.35rem; color: #5e6b73; }
|
||||
.screenshot-status.error { color: #b00020; }
|
||||
</style>
|
||||
<script>
|
||||
(() => {
|
||||
const driverType = document.getElementById("driver-type");
|
||||
const syncPlatformFields = () => {
|
||||
const ios = driverType.value === "wda";
|
||||
document.querySelectorAll(".ios-setting").forEach((el) => { el.hidden = !ios; });
|
||||
document.querySelectorAll(".android-setting").forEach((el) => { el.hidden = ios; });
|
||||
};
|
||||
driverType.addEventListener("change", syncPlatformFields);
|
||||
syncPlatformFields();
|
||||
const csrfInput = document.querySelector('input[name="csrf_token"]');
|
||||
const csrfToken = csrfInput ? csrfInput.value : "";
|
||||
const form = document.querySelector('form[action="/devices/save"]');
|
||||
const discoverButton = document.getElementById("discover-ios");
|
||||
const discoveryStatus = document.getElementById("discovery-status");
|
||||
const discoveredDevices = document.getElementById("discovered-devices");
|
||||
discoverButton.addEventListener("click", async () => {
|
||||
discoverButton.disabled = true;
|
||||
discoveryStatus.classList.remove("error");
|
||||
discoveryStatus.textContent = "Scanning...";
|
||||
discoveredDevices.replaceChildren();
|
||||
try {
|
||||
const response = await fetch("/api/devices/discover-ios");
|
||||
const payload = await response.json();
|
||||
if (!response.ok) throw new Error(payload.detail || "Discovery failed.");
|
||||
discoveryStatus.textContent = payload.devices.length
|
||||
? `Found ${payload.devices.length} connected iPhone(s).`
|
||||
: "No connected, paired iPhones found.";
|
||||
payload.devices.forEach((device, index) => {
|
||||
const row = document.createElement("p");
|
||||
const description = document.createElement("span");
|
||||
description.textContent = `${device.name} - ${device.model} - iOS ${device.os_version} (${device.transport}) `;
|
||||
const select = document.createElement("button");
|
||||
select.type = "button";
|
||||
select.textContent = device.configured ? "Already added" : "Use this iPhone";
|
||||
select.disabled = device.configured;
|
||||
select.addEventListener("click", () => {
|
||||
form.elements.driver_type.value = "wda";
|
||||
form.elements.udid.value = device.udid;
|
||||
form.elements.device_name.value = device.name;
|
||||
form.elements.wda_local_port.value = device.suggested_wda_port;
|
||||
form.elements.device_id.value ||= `ios-phone-${index + 1}`;
|
||||
form.elements.name.value ||= device.name;
|
||||
const advanced = JSON.parse(form.elements.connection_info.value || "{}");
|
||||
advanced.mjpegServerPort = device.suggested_mjpeg_port;
|
||||
advanced.derivedDataPath = `/tmp/wda-${device.udid}`;
|
||||
form.elements.connection_info.value = JSON.stringify(advanced, null, 2);
|
||||
syncPlatformFields();
|
||||
form.scrollIntoView({ behavior: "smooth", block: "start" });
|
||||
});
|
||||
row.append(description, select);
|
||||
discoveredDevices.append(row);
|
||||
});
|
||||
} catch (error) {
|
||||
discoveryStatus.classList.add("error");
|
||||
discoveryStatus.textContent = error.message || "Discovery failed.";
|
||||
} finally {
|
||||
discoverButton.disabled = false;
|
||||
}
|
||||
});
|
||||
const testButton = document.getElementById("test-connection");
|
||||
const testStatus = document.getElementById("connection-test-status");
|
||||
const connectionInfo = () => {
|
||||
const data = new FormData(form);
|
||||
let info = {};
|
||||
const advanced = String(data.get("connection_info") || "{}");
|
||||
info = JSON.parse(advanced);
|
||||
["server_url", "udid", "device_name"].forEach((key) => {
|
||||
const value = String(data.get(key) || "").trim();
|
||||
if (value) info[key] = value;
|
||||
});
|
||||
const portKey = data.get("driver_type") === "wda" ? "wda_local_port" : "system_port";
|
||||
const port = String(data.get(portKey) || "").trim();
|
||||
if (port) info[portKey] = Number(port);
|
||||
return { driver_type: data.get("driver_type"), connection_info: info };
|
||||
};
|
||||
testButton.addEventListener("click", async () => {
|
||||
testButton.disabled = true;
|
||||
testStatus.classList.remove("error");
|
||||
testStatus.textContent = "Testing connection...";
|
||||
try {
|
||||
const response = await fetch("/api/devices/test-connection", {
|
||||
method: "POST",
|
||||
headers: { "Content-Type": "application/json", "X-CSRF-Token": csrfToken },
|
||||
body: JSON.stringify(connectionInfo()),
|
||||
});
|
||||
const payload = await response.json();
|
||||
if (!response.ok) throw new Error(payload.detail || "Connection test failed.");
|
||||
testStatus.textContent = "Connection successful.";
|
||||
} catch (error) {
|
||||
testStatus.classList.add("error");
|
||||
testStatus.textContent = error.message || "Connection test failed.";
|
||||
} finally {
|
||||
testButton.disabled = false;
|
||||
}
|
||||
});
|
||||
document.querySelectorAll(".screenshot-button").forEach((button) => {
|
||||
button.addEventListener("click", async () => {
|
||||
const deviceId = button.dataset.deviceId;
|
||||
const preview = button.parentElement.querySelector("[data-screenshot-preview]");
|
||||
const status = preview.querySelector("[data-screenshot-status]");
|
||||
const image = preview.querySelector("[data-screenshot-image]");
|
||||
const previousUrl = image.dataset.objectUrl;
|
||||
if (previousUrl) URL.revokeObjectURL(previousUrl);
|
||||
button.disabled = true;
|
||||
preview.hidden = false;
|
||||
image.hidden = true;
|
||||
status.classList.remove("error");
|
||||
status.textContent = "Capturing...";
|
||||
try {
|
||||
const response = await fetch(
|
||||
"/api/devices/" + encodeURIComponent(deviceId) + "/screenshot",
|
||||
{ method: "POST", headers: { "X-CSRF-Token": csrfToken } }
|
||||
);
|
||||
if (!response.ok) {
|
||||
let detail = "Screenshot failed.";
|
||||
try {
|
||||
const payload = await response.json();
|
||||
if (payload.detail) detail = payload.detail;
|
||||
} catch (_) {}
|
||||
throw new Error(detail);
|
||||
}
|
||||
const objectUrl = URL.createObjectURL(await response.blob());
|
||||
image.src = objectUrl;
|
||||
image.dataset.objectUrl = objectUrl;
|
||||
image.hidden = false;
|
||||
status.textContent = "Captured.";
|
||||
} catch (error) {
|
||||
status.classList.add("error");
|
||||
status.textContent = error.message || "Screenshot failed.";
|
||||
} finally {
|
||||
button.disabled = false;
|
||||
}
|
||||
});
|
||||
});
|
||||
})();
|
||||
</script>
|
||||
{% endblock %}
|
||||
|
||||
@@ -10,6 +10,7 @@ dependencies = [
|
||||
"filelock>=3.0",
|
||||
"httpx>=0.27.0",
|
||||
"jinja2>=3.1",
|
||||
"mcp>=1.28,<2",
|
||||
"uvicorn[standard]>=0.30.0",
|
||||
]
|
||||
|
||||
|
||||
@@ -47,6 +47,7 @@ def test_devices_renders(env, sample_session) -> None:
|
||||
**make_devices_context(sample_session)
|
||||
)
|
||||
assert '<form method="post" action="/devices/save">' in html
|
||||
assert 'class="screenshot-button"' in html
|
||||
|
||||
|
||||
def test_account_renders(env, sample_session) -> None:
|
||||
|
||||
@@ -7,6 +7,7 @@ from datetime import UTC, datetime, timedelta
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from starlette.testclient import TestClient
|
||||
|
||||
from cloud.internal_api.models import (
|
||||
AssignmentModel,
|
||||
@@ -15,10 +16,23 @@ from cloud.internal_api.models import (
|
||||
)
|
||||
from device.manager import DeviceManager
|
||||
from host_agent.app import HostAgentApplication, create_application
|
||||
from host_agent.assignment import AssignmentExecutor
|
||||
from host_agent.config import HostAgentConfig
|
||||
from host_agent.execution import create_execution_factories
|
||||
from host_agent.history import ConsoleHistoryStore
|
||||
from host_agent.identity import HostIdentityStore
|
||||
from host_agent.instance_lock import InstanceAlreadyRunningError
|
||||
from host_agent.local_account import LocalAccountStore
|
||||
from host_agent.mcp_lock import McpBusyTracker
|
||||
from host_agent.mcp_token import McpTokenStore
|
||||
from host_agent.status import AgentStatusTracker
|
||||
from host_agent.web.app import create_console_app
|
||||
from host_agent.web.auth import SessionManager
|
||||
from host_agent.web.mcp import build_mcp_server
|
||||
from storage.artifact_store import ArtifactStore
|
||||
from storage.device_config import DeviceConfigStore
|
||||
from storage.task_metadata import TaskMetadataStore
|
||||
from storage.timeline import Timeline
|
||||
|
||||
|
||||
def _free_loopback_port() -> int:
|
||||
@@ -653,6 +667,77 @@ def test_create_application_with_independent_identity_paths_coexist(
|
||||
asyncio.run(app_a.client.aclose())
|
||||
|
||||
|
||||
def test_create_application_wires_mcp_components(tmp_path, monkeypatch) -> None:
|
||||
"""create_application produces a console app with /mcp mounted (auth-protected)
|
||||
and persists the host_mcp_token.json file alongside the identity."""
|
||||
monkeypatch.chdir(tmp_path)
|
||||
config = _config()
|
||||
config_store = DeviceConfigStore(tmp_path / "devices.sqlite3")
|
||||
identity_store = HostIdentityStore(config.identity_path)
|
||||
history_store = ConsoleHistoryStore(
|
||||
tmp_path / "host_console_history.sqlite3",
|
||||
limit=config.console_history_limit,
|
||||
)
|
||||
metadata_store = TaskMetadataStore(db_path=config.task_progress_db_path)
|
||||
timeline = Timeline(ArtifactStore(root=config.task_artifact_dir))
|
||||
status_tracker = AgentStatusTracker()
|
||||
|
||||
application = create_application(
|
||||
config=config,
|
||||
device_config_store=config_store,
|
||||
identity_store=identity_store,
|
||||
manager=DeviceManager(),
|
||||
)
|
||||
|
||||
# Token file must exist after create_application.
|
||||
assert (config.identity_path.parent / "host_mcp_token.json").exists()
|
||||
|
||||
# Heartbeat must hold the in-process McpBusyTracker.
|
||||
assert application.heartbeat.mcp_busy_tracker is not None
|
||||
|
||||
# Build the same console app the production path builds and verify /mcp
|
||||
# is mounted (responds 401, not 404) without a bearer token.
|
||||
mcp_token_store = McpTokenStore(config.identity_path.parent / "host_mcp_token.json")
|
||||
mcp_token_store.load_or_create()
|
||||
mcp_busy_tracker = McpBusyTracker(ttl_seconds=20.0)
|
||||
mcp_server = build_mcp_server(
|
||||
manager=application.heartbeat.manager,
|
||||
mcp_busy_tracker=mcp_busy_tracker,
|
||||
status_tracker=status_tracker,
|
||||
)
|
||||
console_app = create_console_app(
|
||||
config=config,
|
||||
manager=application.heartbeat.manager,
|
||||
config_store=config_store,
|
||||
local_account_store=LocalAccountStore(config.local_account_path),
|
||||
identity_store=identity_store,
|
||||
history_store=history_store,
|
||||
status_tracker=status_tracker,
|
||||
session_manager=SessionManager(ttl_seconds=config.console_session_ttl_seconds),
|
||||
enrollment_client=None,
|
||||
host_client=application.client,
|
||||
metadata_store=metadata_store,
|
||||
timeline=timeline,
|
||||
executor=AssignmentExecutor(
|
||||
create_execution_factories(
|
||||
application.heartbeat.manager,
|
||||
metadata_store=metadata_store,
|
||||
timeline=timeline,
|
||||
host_agent_config=config,
|
||||
),
|
||||
mcp_busy_tracker=mcp_busy_tracker,
|
||||
),
|
||||
mcp_server=mcp_server,
|
||||
mcp_token_store=mcp_token_store,
|
||||
mcp_busy_tracker=mcp_busy_tracker,
|
||||
)
|
||||
with TestClient(console_app) as client:
|
||||
resp = client.post("/mcp/")
|
||||
assert resp.status_code == 401 # auth required, not 404
|
||||
|
||||
asyncio.run(application.client.aclose())
|
||||
|
||||
|
||||
def test_lock_released_after_run_async_allows_restart(tmp_path, monkeypatch) -> None:
|
||||
monkeypatch.chdir(tmp_path)
|
||||
identity_path = tmp_path / "host_identity.json"
|
||||
|
||||
@@ -154,7 +154,14 @@ def test_workflow_assignment_maps_cancellation_stop_to_cancelled_status() -> Non
|
||||
return object() if definition_id == "workflow-a" else None
|
||||
|
||||
class FakeWorkflowRunner:
|
||||
def run(self, loaded_definition, device_id: str, *, should_stop=None, stop_reason=None):
|
||||
def run(
|
||||
self,
|
||||
loaded_definition,
|
||||
device_id: str,
|
||||
*,
|
||||
should_stop=None,
|
||||
stop_reason=None,
|
||||
):
|
||||
assert should_stop is not None and should_stop()
|
||||
assert stop_reason is not None
|
||||
return SimpleNamespace(
|
||||
@@ -180,6 +187,63 @@ def test_workflow_assignment_maps_cancellation_stop_to_cancelled_status() -> Non
|
||||
}
|
||||
|
||||
|
||||
def test_execute_fails_fast_when_mcp_session_holds_device() -> None:
|
||||
"""Cloud assignment arriving for a device currently held by an MCP
|
||||
session must fail immediately rather than fight for the device."""
|
||||
from host_agent.mcp_lock import McpBusyTracker
|
||||
|
||||
tracker = McpBusyTracker()
|
||||
tracker.acquire("phone-1", "sess-mcp")
|
||||
executor = AssignmentExecutor(
|
||||
_build_factories(),
|
||||
mcp_busy_tracker=tracker,
|
||||
)
|
||||
assignment = _assignment(device_id="phone-1")
|
||||
result = executor.execute(assignment)
|
||||
assert result.status == "failed"
|
||||
assert "MCP" in (result.failure_reason or "")
|
||||
|
||||
|
||||
def test_execute_skips_check_when_tracker_is_none() -> None:
|
||||
"""Default backward-compat: no tracker → no fail-fast."""
|
||||
executor = AssignmentExecutor(_build_factories())
|
||||
# Without a real workflow store / task runner this test verifies the
|
||||
# entry-point path doesn't raise on the mcp_busy check.
|
||||
# We use a goal + a mock runner factory so execute() runs through.
|
||||
assignment = _assignment()
|
||||
result = executor.execute(assignment)
|
||||
# Should run through normally (not fail on MCP check)
|
||||
assert result.status == "done"
|
||||
|
||||
|
||||
def _build_factories() -> ExecutionFactories:
|
||||
"""Shared factory fixture used by MCP-hold tests."""
|
||||
received: list[Task] = []
|
||||
|
||||
class FakeTaskRunner:
|
||||
def run(self, task: Task) -> Task:
|
||||
received.append(task)
|
||||
task.status = "completed"
|
||||
return task
|
||||
|
||||
class FakeMetadataStore:
|
||||
def create_task(
|
||||
self,
|
||||
task: Task,
|
||||
*,
|
||||
source_task_id: str | None = None,
|
||||
source_attempt: int | None = None,
|
||||
) -> None:
|
||||
pass
|
||||
|
||||
return ExecutionFactories(
|
||||
task_runner_factory=lambda: FakeTaskRunner(), # type: ignore[arg-type,return-value]
|
||||
workflow_runner_factory=lambda: object(), # type: ignore[arg-type,return-value]
|
||||
workflow_store=object(), # type: ignore[arg-type]
|
||||
metadata_store=FakeMetadataStore(), # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
|
||||
def test_unknown_workflow_fails_without_running() -> None:
|
||||
class FakeWorkflowStore:
|
||||
def get_definition(self, definition_id: str):
|
||||
|
||||
@@ -146,3 +146,21 @@ def test_duplicate_instance_exits_with_clear_error(
|
||||
err = capsys.readouterr().err
|
||||
assert "another Host Agent instance" in err
|
||||
assert str(lock_path) in err
|
||||
|
||||
|
||||
def test_mcp_token_subcommand_prints_token(tmp_path, capsys, monkeypatch) -> None:
|
||||
monkeypatch.setenv("HOST_AGENT_IDENTITY_PATH", str(tmp_path / "host_identity.json"))
|
||||
monkeypatch.setenv(
|
||||
"HOST_AGENT_LOCAL_ACCOUNT_PATH", str(tmp_path / "host_local_account.json")
|
||||
)
|
||||
# Also set control plane URL to satisfy config loading
|
||||
monkeypatch.setenv("HOST_AGENT_CONTROL_PLANE_URL", "https://cloud.example")
|
||||
from host_agent.cli import main
|
||||
|
||||
main(["mcp-token"])
|
||||
out = capsys.readouterr().out.strip()
|
||||
assert len(out) >= 40 # token is ~43 chars
|
||||
# Subsequent invocation prints the same token (idempotent).
|
||||
main(["mcp-token"])
|
||||
out2 = capsys.readouterr().out.strip()
|
||||
assert out == out2
|
||||
|
||||
@@ -353,6 +353,60 @@ def test_submit_self_task_does_not_duplicate_when_response_is_lost() -> None:
|
||||
assert attempts == 1
|
||||
|
||||
|
||||
def test_heartbeat_includes_mcp_busy_device_ids_in_payload() -> None:
|
||||
"""When mcp_busy_device_ids is passed, the client sends it in the request."""
|
||||
captured: list[dict[str, object]] = []
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
captured.append(json.loads(request.content))
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"host_id": "host-a",
|
||||
"accepted_devices": 0,
|
||||
"received_at": "2026-07-12T00:00:00Z",
|
||||
},
|
||||
)
|
||||
|
||||
async def scenario() -> None:
|
||||
async with httpx.AsyncClient(
|
||||
transport=httpx.MockTransport(handler),
|
||||
base_url="https://control.example",
|
||||
) as http_client:
|
||||
client = HostAgentClient(_config(), http_client=http_client)
|
||||
await client.heartbeat([], mcp_busy_device_ids=["phone-1"])
|
||||
|
||||
asyncio.run(scenario())
|
||||
assert captured[0]["mcp_busy_device_ids"] == ["phone-1"]
|
||||
|
||||
|
||||
def test_heartbeat_omits_mcp_busy_device_ids_when_empty() -> None:
|
||||
"""Empty list is omitted from the payload (backward compatible)."""
|
||||
captured: list[dict[str, object]] = []
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
captured.append(json.loads(request.content))
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"host_id": "host-a",
|
||||
"accepted_devices": 0,
|
||||
"received_at": "2026-07-12T00:00:00Z",
|
||||
},
|
||||
)
|
||||
|
||||
async def scenario() -> None:
|
||||
async with httpx.AsyncClient(
|
||||
transport=httpx.MockTransport(handler),
|
||||
base_url="https://control.example",
|
||||
) as http_client:
|
||||
client = HostAgentClient(_config(), http_client=http_client)
|
||||
await client.heartbeat([], mcp_busy_device_ids=[])
|
||||
|
||||
asyncio.run(scenario())
|
||||
assert "mcp_busy_device_ids" not in captured[0]
|
||||
|
||||
|
||||
def test_bootstrap_client_directly_enrolls_and_enrolls_device() -> None:
|
||||
requests: list[httpx.Request] = []
|
||||
host_attempts = 0
|
||||
|
||||
@@ -96,6 +96,36 @@ def test_decide_base64_encodes_screenshot() -> None:
|
||||
assert body["screenshot_base64"] == "aGVsbG8="
|
||||
|
||||
|
||||
def test_decide_forwards_planner_history() -> None:
|
||||
seen_requests: list[httpx.Request] = []
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
seen_requests.append(request)
|
||||
return httpx.Response(200, json={"tool_name": "tap", "arguments": {}})
|
||||
|
||||
client = _client(handler)
|
||||
history = [
|
||||
{
|
||||
"user_prompt": "first screen",
|
||||
"tool_name": "tap",
|
||||
"arguments": {"x": 1, "y": 2},
|
||||
"rationale": "Open it.",
|
||||
"tool_result": {"success": True},
|
||||
}
|
||||
]
|
||||
|
||||
client.decide(
|
||||
system_prompt="sp",
|
||||
user_prompt="next screen",
|
||||
screenshot=None,
|
||||
tools=_TOOLS,
|
||||
timeout=10.0,
|
||||
history=history,
|
||||
)
|
||||
|
||||
assert json.loads(seen_requests[0].content)["history"] == history
|
||||
|
||||
|
||||
def test_decide_clamps_legacy_timeout_and_waits_for_cloud_profile_timeout() -> None:
|
||||
seen_requests: list[httpx.Request] = []
|
||||
|
||||
|
||||
@@ -27,6 +27,17 @@ def test_load_host_agent_config_allows_explicit_direct_planner_transport() -> No
|
||||
assert config.ai_planner_transport == "direct"
|
||||
|
||||
|
||||
def test_load_host_agent_config_supports_local_mode() -> None:
|
||||
config = load_host_agent_config({"HOST_AGENT_MODE": "local"})
|
||||
|
||||
assert config.mode == "local"
|
||||
assert config.control_plane_url == ""
|
||||
assert config.enrollment_managed is False
|
||||
assert config.ai_planner_transport == "direct"
|
||||
assert config.dependency_supervisor_enabled is True
|
||||
assert config.appium_supervised is True
|
||||
|
||||
|
||||
def test_load_host_agent_config_parses_poll_and_retry_values() -> None:
|
||||
config = load_host_agent_config(
|
||||
{
|
||||
|
||||
@@ -8,6 +8,7 @@ from cloud.internal_api.models import HostGovernancePolicyModel
|
||||
from device.manager import DeviceManager
|
||||
from host_agent.config import HostAgentConfig
|
||||
from host_agent.heartbeat import HeartbeatSynchronizer, build_device_snapshot
|
||||
from host_agent.mcp_lock import McpBusyTracker
|
||||
from host_agent.policy_cache import HostPolicyCacheStore
|
||||
from host_agent.status import AgentStatusTracker
|
||||
|
||||
@@ -16,6 +17,9 @@ class ConnectableDriver:
|
||||
def connect(self) -> None:
|
||||
return None
|
||||
|
||||
def screenshot(self) -> bytes:
|
||||
return b"ok"
|
||||
|
||||
|
||||
def _config() -> HostAgentConfig:
|
||||
return HostAgentConfig(
|
||||
@@ -60,7 +64,9 @@ def test_heartbeat_synchronizer_runs_at_configured_interval_until_stopped() -> N
|
||||
calls: list[list[str]] = []
|
||||
|
||||
class FakeClient:
|
||||
async def heartbeat(self, devices, *, address=None, policy_revision=0):
|
||||
async def heartbeat(
|
||||
self, devices, *, address=None, policy_revision=0, **kwargs
|
||||
):
|
||||
calls.append([device.device_id for device in devices])
|
||||
return HeartbeatResponse(
|
||||
host_id="host-a",
|
||||
@@ -86,6 +92,33 @@ def test_heartbeat_synchronizer_runs_at_configured_interval_until_stopped() -> N
|
||||
assert manager.status("device-a") == "busy"
|
||||
|
||||
|
||||
def test_offline_device_is_retried_on_next_heartbeat_cycle() -> None:
|
||||
class FlakyDriver(ConnectableDriver):
|
||||
def __init__(self, fail: bool) -> None:
|
||||
self.fail = fail
|
||||
|
||||
def screenshot(self) -> bytes:
|
||||
if self.fail:
|
||||
raise RuntimeError("WDA disconnected")
|
||||
return b"ok"
|
||||
|
||||
instances: list[FlakyDriver] = []
|
||||
|
||||
def factory() -> FlakyDriver:
|
||||
driver = FlakyDriver(not instances)
|
||||
instances.append(driver)
|
||||
return driver
|
||||
|
||||
manager = DeviceManager()
|
||||
manager.register_device("device-a", factory) # type: ignore[arg-type]
|
||||
manager.connect("device-a")
|
||||
sync = HeartbeatSynchronizer(manager, object(), _config()) # type: ignore[arg-type]
|
||||
sync.probe_connected_devices()
|
||||
assert manager.status("device-a") == "offline"
|
||||
sync.connect_devices()
|
||||
assert manager.status("device-a") == "busy"
|
||||
|
||||
|
||||
def test_sync_once_notifies_status_tracker_and_on_sync_with_device_count() -> None:
|
||||
manager = DeviceManager()
|
||||
manager.register_device(
|
||||
@@ -98,7 +131,9 @@ def test_sync_once_notifies_status_tracker_and_on_sync_with_device_count() -> No
|
||||
)
|
||||
|
||||
class FakeClient:
|
||||
async def heartbeat(self, devices, *, address=None, policy_revision=0):
|
||||
async def heartbeat(
|
||||
self, devices, *, address=None, policy_revision=0, **kwargs
|
||||
):
|
||||
return HeartbeatResponse(
|
||||
host_id="host-a",
|
||||
accepted_devices=len(devices),
|
||||
@@ -133,7 +168,9 @@ def test_heartbeat_caches_safe_host_policy_and_reuses_its_revision(tmp_path) ->
|
||||
revisions: list[int] = []
|
||||
|
||||
class UpdatingClient:
|
||||
async def heartbeat(self, devices, *, address=None, policy_revision=0):
|
||||
async def heartbeat(
|
||||
self, devices, *, address=None, policy_revision=0, **kwargs
|
||||
):
|
||||
revisions.append(policy_revision)
|
||||
return HeartbeatResponse(
|
||||
host_id="host-a",
|
||||
@@ -175,6 +212,63 @@ def test_heartbeat_caches_safe_host_policy_and_reuses_its_revision(tmp_path) ->
|
||||
|
||||
asyncio.run(scenario())
|
||||
assert revisions == [0]
|
||||
assert '"token":' not in (
|
||||
tmp_path / "host_policy.json"
|
||||
).read_text(encoding="utf-8")
|
||||
assert '"token":' not in (tmp_path / "host_policy.json").read_text(encoding="utf-8")
|
||||
|
||||
|
||||
def test_sync_once_passes_mcp_busy_device_ids_to_client() -> None:
|
||||
"""When mcp_busy_tracker has a lease, sync_once relays the device_ids."""
|
||||
manager = DeviceManager()
|
||||
tracker = McpBusyTracker()
|
||||
assert tracker.acquire("phone-1", "sess-a")
|
||||
last_kwargs: dict[str, object] = {}
|
||||
|
||||
class FakeClient:
|
||||
async def heartbeat(
|
||||
self, devices, *, address=None, policy_revision=0, **kwargs
|
||||
):
|
||||
last_kwargs.update(kwargs)
|
||||
return HeartbeatResponse(
|
||||
host_id="host-a",
|
||||
accepted_devices=len(devices),
|
||||
received_at=datetime.now(UTC),
|
||||
)
|
||||
|
||||
async def scenario() -> None:
|
||||
sync = HeartbeatSynchronizer(
|
||||
manager,
|
||||
FakeClient(), # type: ignore[arg-type]
|
||||
_config(),
|
||||
mcp_busy_tracker=tracker,
|
||||
)
|
||||
await sync.sync_once()
|
||||
|
||||
asyncio.run(scenario())
|
||||
assert last_kwargs.get("mcp_busy_device_ids") == ["phone-1"]
|
||||
|
||||
|
||||
def test_sync_once_passes_empty_when_tracker_is_none() -> None:
|
||||
"""Default: no tracker → no busy device ids forwarded."""
|
||||
manager = DeviceManager()
|
||||
last_kwargs: dict[str, object] = {}
|
||||
|
||||
class FakeClient:
|
||||
async def heartbeat(
|
||||
self, devices, *, address=None, policy_revision=0, **kwargs
|
||||
):
|
||||
last_kwargs.update(kwargs)
|
||||
return HeartbeatResponse(
|
||||
host_id="host-a",
|
||||
accepted_devices=len(devices),
|
||||
received_at=datetime.now(UTC),
|
||||
)
|
||||
|
||||
async def scenario() -> None:
|
||||
sync = HeartbeatSynchronizer(
|
||||
manager,
|
||||
FakeClient(), # type: ignore[arg-type]
|
||||
_config(),
|
||||
)
|
||||
await sync.sync_once()
|
||||
|
||||
asyncio.run(scenario())
|
||||
assert not last_kwargs.get("mcp_busy_device_ids")
|
||||
|
||||
@@ -0,0 +1,58 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
from starlette.applications import Starlette
|
||||
from starlette.responses import JSONResponse
|
||||
from starlette.testclient import TestClient
|
||||
|
||||
from host_agent.mcp_token import McpTokenStore
|
||||
from host_agent.web.mcp_auth import BearerAuthMiddleware
|
||||
|
||||
|
||||
def _make_client(tmp_path: Path) -> tuple[TestClient, str]:
|
||||
store = McpTokenStore(tmp_path / "host_mcp_token.json")
|
||||
token = store.load_or_create().token
|
||||
|
||||
async def hello(request): # type: ignore[no-untyped-def]
|
||||
return JSONResponse({"ok": True})
|
||||
|
||||
inner = Starlette(routes=[])
|
||||
inner.router.add_route("/", hello, methods=["GET"])
|
||||
wrapped = Starlette()
|
||||
wrapped.add_middleware(BearerAuthMiddleware, token_store=store)
|
||||
wrapped.mount("/", inner)
|
||||
return TestClient(wrapped), token
|
||||
|
||||
|
||||
def test_no_header_returns_401(tmp_path: Path) -> None:
|
||||
client, _ = _make_client(tmp_path)
|
||||
resp = client.get("/")
|
||||
assert resp.status_code == 401
|
||||
assert resp.headers["WWW-Authenticate"] == "Bearer"
|
||||
assert resp.json() == {"error": "invalid token"}
|
||||
|
||||
|
||||
def test_wrong_token_returns_401(tmp_path: Path) -> None:
|
||||
client, _ = _make_client(tmp_path)
|
||||
resp = client.get("/", headers={"Authorization": "Bearer wrong"})
|
||||
assert resp.status_code == 401
|
||||
|
||||
|
||||
def test_correct_token_passes_through(tmp_path: Path) -> None:
|
||||
client, token = _make_client(tmp_path)
|
||||
resp = client.get("/", headers={"Authorization": f"Bearer {token}"})
|
||||
assert resp.status_code == 200
|
||||
assert resp.json() == {"ok": True}
|
||||
|
||||
|
||||
def test_non_bearer_scheme_returns_401(tmp_path: Path) -> None:
|
||||
client, token = _make_client(tmp_path)
|
||||
resp = client.get("/", headers={"Authorization": f"Basic {token}"})
|
||||
assert resp.status_code == 401
|
||||
|
||||
|
||||
def test_header_case_insensitive(tmp_path: Path) -> None:
|
||||
client, token = _make_client(tmp_path)
|
||||
resp = client.get("/", headers={"authorization": f"Bearer {token}"})
|
||||
assert resp.status_code == 200
|
||||
@@ -0,0 +1,175 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
from datetime import UTC, datetime
|
||||
|
||||
from host_agent.mcp_lock import McpBusyTracker
|
||||
|
||||
|
||||
def _tracker_with_now() -> tuple[McpBusyTracker, list[datetime]]:
|
||||
times: list[datetime] = []
|
||||
|
||||
def now() -> datetime:
|
||||
return times[-1] if times else datetime(2026, 1, 1, tzinfo=UTC)
|
||||
|
||||
tracker = McpBusyTracker(ttl_seconds=60.0, now=now)
|
||||
return tracker, times
|
||||
|
||||
|
||||
def test_acquire_succeeds_on_empty() -> None:
|
||||
tracker, _ = _tracker_with_now()
|
||||
assert tracker.acquire("phone-1", "sess-a") is True
|
||||
assert "phone-1" in tracker.busy_device_ids()
|
||||
|
||||
|
||||
def test_acquire_fails_when_held_by_other_session() -> None:
|
||||
tracker, _ = _tracker_with_now()
|
||||
assert tracker.acquire("phone-1", "sess-a") is True
|
||||
assert tracker.acquire("phone-1", "sess-b") is False
|
||||
|
||||
|
||||
def test_acquire_is_idempotent_for_same_session() -> None:
|
||||
tracker, _ = _tracker_with_now()
|
||||
assert tracker.acquire("phone-1", "sess-a") is True
|
||||
# Same session re-acquiring is allowed (acts as renew).
|
||||
assert tracker.acquire("phone-1", "sess-a") is True
|
||||
|
||||
|
||||
def test_renew_refreshes_last_seen() -> None:
|
||||
tracker, times = _tracker_with_now()
|
||||
times.append(datetime(2026, 1, 1, 12, 0, tzinfo=UTC))
|
||||
tracker.acquire("phone-1", "sess-a")
|
||||
initial = tracker.snapshot()[0]
|
||||
times.append(datetime(2026, 1, 1, 12, 0, 30, tzinfo=UTC))
|
||||
assert tracker.renew("phone-1", "sess-a") is True
|
||||
refreshed = tracker.snapshot()[0]
|
||||
assert refreshed.last_seen_at > initial.last_seen_at
|
||||
|
||||
|
||||
def test_renew_fails_when_held_by_other() -> None:
|
||||
tracker, _ = _tracker_with_now()
|
||||
tracker.acquire("phone-1", "sess-a")
|
||||
assert tracker.renew("phone-1", "sess-b") is False
|
||||
|
||||
|
||||
def test_release_returns_freed_device_ids() -> None:
|
||||
tracker, _ = _tracker_with_now()
|
||||
tracker.acquire("phone-1", "sess-a")
|
||||
tracker.acquire("phone-2", "sess-a")
|
||||
freed = tracker.release("sess-a")
|
||||
assert sorted(freed) == ["phone-1", "phone-2"]
|
||||
assert tracker.busy_device_ids() == []
|
||||
|
||||
|
||||
def test_release_only_frees_caller_session() -> None:
|
||||
tracker, _ = _tracker_with_now()
|
||||
tracker.acquire("phone-1", "sess-a")
|
||||
tracker.acquire("phone-1", "sess-b") # fails
|
||||
freed = tracker.release("sess-b")
|
||||
assert freed == []
|
||||
assert "phone-1" in tracker.busy_device_ids()
|
||||
|
||||
|
||||
def test_ttl_sweeps_expired_leases() -> None:
|
||||
tracker, times = _tracker_with_now()
|
||||
times.append(datetime(2026, 1, 1, 12, 0, tzinfo=UTC))
|
||||
tracker.acquire("phone-1", "sess-a")
|
||||
# Advance past TTL without renew.
|
||||
times.append(datetime(2026, 1, 1, 12, 1, 1, tzinfo=UTC)) # 61s later
|
||||
assert tracker.busy_device_ids() == []
|
||||
|
||||
|
||||
def test_renew_after_ttl_tolerates_same_session() -> None:
|
||||
"""Scene 10: lease expired but session_id matches -> re-acquire."""
|
||||
tracker, times = _tracker_with_now()
|
||||
times.append(datetime(2026, 1, 1, 12, 0, tzinfo=UTC))
|
||||
tracker.acquire("phone-1", "sess-a")
|
||||
times.append(datetime(2026, 1, 1, 12, 1, 1, tzinfo=UTC)) # expired
|
||||
# renew from the same session should succeed (re-acquire).
|
||||
assert tracker.renew("phone-1", "sess-a") is True
|
||||
assert "phone-1" in tracker.busy_device_ids()
|
||||
|
||||
|
||||
def test_snapshot_matches_busy_device_ids() -> None:
|
||||
tracker, _ = _tracker_with_now()
|
||||
tracker.acquire("phone-1", "sess-a")
|
||||
tracker.acquire("phone-2", "sess-a")
|
||||
snap = tracker.snapshot()
|
||||
assert {lease.device_id for lease in snap} == set(tracker.busy_device_ids())
|
||||
|
||||
|
||||
def test_wait_until_usable_succeeds_when_free() -> None:
|
||||
tracker, _ = _tracker_with_now()
|
||||
ok = tracker.wait_until_usable("phone-1", "sess-a", timeout=1.0, poll_interval=0.01)
|
||||
assert ok is True
|
||||
assert "phone-1" in tracker.busy_device_ids()
|
||||
|
||||
|
||||
def test_wait_until_usable_returns_false_on_timeout() -> None:
|
||||
tracker, _ = _tracker_with_now()
|
||||
tracker.acquire("phone-1", "sess-a")
|
||||
ok = tracker.wait_until_usable("phone-1", "sess-b", timeout=0.1, poll_interval=0.02)
|
||||
assert ok is False
|
||||
|
||||
|
||||
def test_wait_until_usable_blocks_then_succeeds_when_released() -> None:
|
||||
tracker, _ = _tracker_with_now()
|
||||
tracker.acquire("phone-1", "sess-a")
|
||||
|
||||
def releaser() -> None:
|
||||
import time
|
||||
|
||||
time.sleep(0.05)
|
||||
tracker.release("sess-a")
|
||||
|
||||
t = threading.Thread(target=releaser)
|
||||
t.start()
|
||||
try:
|
||||
ok = tracker.wait_until_usable(
|
||||
"phone-1", "sess-b", timeout=2.0, poll_interval=0.02
|
||||
)
|
||||
assert ok is True
|
||||
finally:
|
||||
t.join()
|
||||
|
||||
|
||||
def test_wait_until_usable_blocks_then_fails_when_cloud_remains_busy() -> None:
|
||||
tracker, _ = _tracker_with_now()
|
||||
ok = tracker.wait_until_usable(
|
||||
"phone-1",
|
||||
"sess-a",
|
||||
timeout=0.1,
|
||||
poll_interval=0.02,
|
||||
cloud_busy_check=lambda: True,
|
||||
)
|
||||
assert ok is False
|
||||
assert tracker.busy_device_ids() == []
|
||||
|
||||
|
||||
def test_default_ttl_is_20_seconds() -> None:
|
||||
"""The McpBusyTracker default TTL is 20s: short enough to recover from a
|
||||
dead MCP session within a heartbeat interval without an explicit release
|
||||
callback (mcp SDK 1.28.1 has no per-session shutdown hook), but long
|
||||
enough that an actively-busy session does not lose its lease during
|
||||
normal operator pauses."""
|
||||
tracker = McpBusyTracker()
|
||||
assert tracker._ttl == 20.0
|
||||
|
||||
|
||||
def test_default_ttl_recovers_dead_session_within_one_window() -> None:
|
||||
"""With the 20s default, a session that never renews its lease is
|
||||
reaped within one TTL window on the next read. This is the
|
||||
concrete fallback behavior for I1 (no FastMCP session-end hook)."""
|
||||
times: list[datetime] = []
|
||||
|
||||
def now() -> datetime:
|
||||
return times[-1] if times else datetime(2026, 1, 1, tzinfo=UTC)
|
||||
|
||||
tracker = McpBusyTracker(now=now) # default 20s TTL
|
||||
times.append(datetime(2026, 1, 1, 12, 0, tzinfo=UTC))
|
||||
assert tracker.acquire("phone-1", "dead-session") is True
|
||||
# No renew: advance 21s. Lease should be swept on next read.
|
||||
times.append(datetime(2026, 1, 1, 12, 0, 21, tzinfo=UTC))
|
||||
assert tracker.busy_device_ids() == []
|
||||
# New session can now acquire cleanly (no stale-busy contamination).
|
||||
assert tracker.acquire("phone-1", "new-session") is True
|
||||
@@ -0,0 +1,107 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
import stat
|
||||
import sys
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
from host_agent.mcp_token import McpToken, McpTokenStore, McpTokenStoreError
|
||||
|
||||
|
||||
def test_load_or_create_generates_when_missing(tmp_path: Path) -> None:
|
||||
store = McpTokenStore(tmp_path / "host_mcp_token.json")
|
||||
token = store.load_or_create()
|
||||
assert token.version == 1
|
||||
assert len(token.token) >= 40 # secrets.token_urlsafe(32) -> ~43 chars
|
||||
assert isinstance(token.created_at, datetime)
|
||||
# File now exists.
|
||||
assert (tmp_path / "host_mcp_token.json").exists()
|
||||
|
||||
|
||||
def test_load_or_create_is_idempotent(tmp_path: Path) -> None:
|
||||
store = McpTokenStore(tmp_path / "host_mcp_token.json")
|
||||
first = store.load_or_create()
|
||||
second = McpTokenStore(tmp_path / "host_mcp_token.json").load_or_create()
|
||||
assert first.token == second.token
|
||||
|
||||
|
||||
def test_load_or_create_writes_json_schema(tmp_path: Path) -> None:
|
||||
path = tmp_path / "host_mcp_token.json"
|
||||
McpTokenStore(path).load_or_create()
|
||||
data = json.loads(path.read_text())
|
||||
assert set(data) == {"version", "token", "created_at"}
|
||||
assert data["version"] == 1
|
||||
assert isinstance(data["token"], str)
|
||||
# created_at is ISO 8601.
|
||||
datetime.fromisoformat(data["created_at"])
|
||||
|
||||
|
||||
@pytest.mark.skipif(sys.platform == "win32", reason="POSIX perms only")
|
||||
def test_load_or_create_sets_posix_permissions(tmp_path: Path) -> None:
|
||||
path = tmp_path / "host_mcp_token.json"
|
||||
McpTokenStore(path).load_or_create()
|
||||
mode = stat.S_IMODE(os.fstat(os.open(path, os.O_RDONLY)).st_mode)
|
||||
assert mode == 0o600
|
||||
|
||||
|
||||
def test_verify_accepts_correct_token(tmp_path: Path) -> None:
|
||||
store = McpTokenStore(tmp_path / "host_mcp_token.json")
|
||||
token = store.load_or_create()
|
||||
assert store.verify(token.token) is True
|
||||
|
||||
|
||||
def test_verify_rejects_wrong_token(tmp_path: Path) -> None:
|
||||
store = McpTokenStore(tmp_path / "host_mcp_token.json")
|
||||
store.load_or_create()
|
||||
assert store.verify("wrong") is False
|
||||
|
||||
|
||||
def test_load_or_create_raises_on_corrupt_json(tmp_path: Path) -> None:
|
||||
path = tmp_path / "host_mcp_token.json"
|
||||
path.write_text("{not valid json")
|
||||
with pytest.raises(McpTokenStoreError):
|
||||
McpTokenStore(path).load_or_create()
|
||||
|
||||
|
||||
@pytest.mark.skipif(sys.platform == "win32", reason="POSIX chmod enforcement only")
|
||||
def test_load_or_create_raises_on_unwritable_dir(tmp_path: Path) -> None:
|
||||
unwritable = tmp_path / "ro"
|
||||
unwritable.mkdir()
|
||||
os.chmod(unwritable, 0o500) # r-x for owner
|
||||
try:
|
||||
with pytest.raises(McpTokenStoreError):
|
||||
McpTokenStore(unwritable / "host_mcp_token.json").load_or_create()
|
||||
finally:
|
||||
os.chmod(unwritable, 0o700) # restore so cleanup works
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
sys.platform == "win32",
|
||||
reason="POSIX atomic-rename semantics only",
|
||||
)
|
||||
def test_load_or_create_concurrent_calls_do_not_corrupt(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
"""Two store instances racing to create: both end up reading the same token."""
|
||||
import threading
|
||||
|
||||
path = tmp_path / "host_mcp_token.json"
|
||||
results: list[McpToken] = []
|
||||
barrier = threading.Barrier(2)
|
||||
|
||||
def worker() -> None:
|
||||
barrier.wait()
|
||||
store = McpTokenStore(path)
|
||||
results.append(store.load_or_create())
|
||||
|
||||
threads = [threading.Thread(target=worker) for _ in range(2)]
|
||||
for t in threads:
|
||||
t.start()
|
||||
for t in threads:
|
||||
t.join()
|
||||
assert len(results) == 2
|
||||
assert results[0].token == results[1].token
|
||||
@@ -6,6 +6,7 @@ from datetime import UTC, datetime
|
||||
from typing import Any
|
||||
|
||||
from fastapi.testclient import TestClient
|
||||
from mcp.server.fastmcp import FastMCP
|
||||
|
||||
from cloud.internal_api.models import AssignmentModel
|
||||
from core.models import Task
|
||||
@@ -15,9 +16,12 @@ from host_agent.config import HostAgentConfig
|
||||
from host_agent.history import ConsoleHistoryStore
|
||||
from host_agent.identity import HostIdentityStore
|
||||
from host_agent.local_account import LocalAccountStore
|
||||
from host_agent.mcp_lock import McpBusyTracker
|
||||
from host_agent.mcp_token import McpTokenStore
|
||||
from host_agent.status import AgentStatusTracker
|
||||
from host_agent.web.app import SESSION_COOKIE_NAME, create_console_app
|
||||
from host_agent.web.auth import SessionManager
|
||||
from host_agent.web.mcp import build_mcp_server
|
||||
from storage.device_config import DeviceConfigStore
|
||||
from storage.task_metadata import TaskMetadataStore
|
||||
|
||||
@@ -34,6 +38,9 @@ def _build_client(
|
||||
submit_self_task: TaskSubmissionCallable | None = None,
|
||||
cancel_task: TaskCancellationCallable | None = None,
|
||||
include_metadata_store: bool = True,
|
||||
mcp_server: FastMCP | None = None,
|
||||
mcp_token_store: McpTokenStore | None = None,
|
||||
mcp_busy_tracker: McpBusyTracker | None = None,
|
||||
) -> tuple[TestClient, dict]:
|
||||
config = HostAgentConfig(
|
||||
control_plane_url="https://control.example",
|
||||
@@ -68,6 +75,9 @@ def _build_client(
|
||||
submit_self_task=submit_self_task,
|
||||
cancel_task=cancel_task,
|
||||
metadata_store=metadata_store,
|
||||
mcp_server=mcp_server,
|
||||
mcp_token_store=mcp_token_store,
|
||||
mcp_busy_tracker=mcp_busy_tracker,
|
||||
)
|
||||
client = TestClient(app)
|
||||
context = {
|
||||
@@ -232,6 +242,147 @@ def test_add_device_appears_in_devices_page_and_manager(tmp_path) -> None:
|
||||
assert [device.id for device in context["manager"].list_devices()] == ["device-a"]
|
||||
|
||||
|
||||
def test_devices_page_captures_screenshot_only_when_button_endpoint_is_called(
|
||||
tmp_path,
|
||||
) -> None:
|
||||
class ScreenshotDriver:
|
||||
def __init__(self) -> None:
|
||||
self.capture_count = 0
|
||||
|
||||
def connect(self) -> None:
|
||||
pass
|
||||
|
||||
def disconnect(self) -> None:
|
||||
pass
|
||||
|
||||
def screenshot(self) -> bytes:
|
||||
self.capture_count += 1
|
||||
return b"fake-png"
|
||||
|
||||
driver = ScreenshotDriver()
|
||||
client, context = _build_client(tmp_path)
|
||||
context["config_store"].add(
|
||||
device_id="device-a",
|
||||
name="Lab iPhone",
|
||||
driver_type="wda",
|
||||
connection_info={},
|
||||
)
|
||||
context["manager"].register_device("device-a", lambda: driver)
|
||||
context["manager"].connect("device-a")
|
||||
csrf_token = _login(client)
|
||||
|
||||
page = client.get("/devices")
|
||||
assert page.status_code == 200
|
||||
assert 'class="screenshot-button"' in page.text
|
||||
assert 'data-device-id="device-a"' in page.text
|
||||
assert driver.capture_count == 0
|
||||
|
||||
response = client.post(
|
||||
"/api/devices/device-a/screenshot",
|
||||
headers={"X-CSRF-Token": csrf_token},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.content == b"fake-png"
|
||||
assert response.headers["content-type"] == "image/png"
|
||||
assert response.headers["cache-control"] == "no-store"
|
||||
assert driver.capture_count == 1
|
||||
|
||||
|
||||
def test_device_screenshot_requires_csrf_and_connected_device(tmp_path) -> None:
|
||||
class ScreenshotDriver:
|
||||
def connect(self) -> None:
|
||||
pass
|
||||
|
||||
def disconnect(self) -> None:
|
||||
pass
|
||||
|
||||
def screenshot(self) -> bytes:
|
||||
return b"fake-png"
|
||||
|
||||
client, context = _build_client(tmp_path)
|
||||
context["manager"].register_device("device-a", ScreenshotDriver)
|
||||
csrf_token = _login(client)
|
||||
|
||||
missing_csrf = client.post("/api/devices/device-a/screenshot")
|
||||
assert missing_csrf.status_code == 403
|
||||
|
||||
offline = client.post(
|
||||
"/api/devices/device-a/screenshot",
|
||||
headers={"X-CSRF-Token": csrf_token},
|
||||
)
|
||||
assert offline.status_code == 503
|
||||
|
||||
|
||||
def test_device_connection_test_connects_checks_health_and_disconnects(
|
||||
tmp_path, monkeypatch
|
||||
) -> None:
|
||||
events: list[str] = []
|
||||
|
||||
class ProbeDriver:
|
||||
def connect(self) -> None:
|
||||
events.append("connect")
|
||||
|
||||
def health_check(self) -> None:
|
||||
events.append("health")
|
||||
|
||||
def disconnect(self) -> None:
|
||||
events.append("disconnect")
|
||||
|
||||
def factory(driver_type, connection_info):
|
||||
assert driver_type == "wda"
|
||||
assert connection_info["udid"] == "ios-udid"
|
||||
return ProbeDriver
|
||||
|
||||
monkeypatch.setattr("host_agent.web.app.build_driver_factory", factory)
|
||||
client, context = _build_client(tmp_path)
|
||||
csrf_token = _login(client)
|
||||
|
||||
response = client.post(
|
||||
"/api/devices/test-connection",
|
||||
json={"driver_type": "wda", "connection_info": {"udid": "ios-udid"}},
|
||||
headers={"X-CSRF-Token": csrf_token},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.json() == {"ok": True, "driver_type": "wda"}
|
||||
assert events == ["connect", "health", "disconnect"]
|
||||
assert context["config_store"].list() == []
|
||||
|
||||
|
||||
def test_ios_discovery_returns_connected_devices_and_unique_ports(
|
||||
tmp_path, monkeypatch
|
||||
) -> None:
|
||||
monkeypatch.setattr(
|
||||
"host_agent.web.app.discover_connected_ios_devices",
|
||||
lambda: [
|
||||
{
|
||||
"udid": "ios-new",
|
||||
"name": "New iPhone",
|
||||
"model": "iPhone 15",
|
||||
"os_version": "18.0",
|
||||
"transport": "wired",
|
||||
}
|
||||
],
|
||||
)
|
||||
client, context = _build_client(tmp_path)
|
||||
context["config_store"].add(
|
||||
device_id="existing",
|
||||
driver_type="wda",
|
||||
connection_info={"udid": "ios-old", "wda_local_port": 8100, "mjpegServerPort": 9100},
|
||||
)
|
||||
_login(client)
|
||||
|
||||
response = client.get("/api/devices/discover-ios")
|
||||
|
||||
assert response.status_code == 200
|
||||
device = response.json()["devices"][0]
|
||||
assert device["udid"] == "ios-new"
|
||||
assert device["configured"] is False
|
||||
assert device["suggested_wda_port"] == 8101
|
||||
assert device["suggested_mjpeg_port"] == 9101
|
||||
|
||||
|
||||
def test_remove_device_unregisters_from_manager(tmp_path) -> None:
|
||||
client, context = _build_client(tmp_path)
|
||||
csrf_token = _login(client)
|
||||
@@ -823,9 +974,7 @@ def _seed_local_task(
|
||||
source_task_id: str | None = "cloud-task-1",
|
||||
) -> str:
|
||||
task = Task(goal="open settings", device_id="dev-1", status=status)
|
||||
metadata_store.create_task(
|
||||
task, source_task_id=source_task_id, source_attempt=1
|
||||
)
|
||||
metadata_store.create_task(task, source_task_id=source_task_id, source_attempt=1)
|
||||
return task.id
|
||||
|
||||
|
||||
@@ -959,3 +1108,131 @@ def test_cancel_task_without_csrf_token_is_rejected(tmp_path) -> None:
|
||||
|
||||
assert response.status_code == 403
|
||||
assert captured == {}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# MCP mount + /api/status fields + dashboard row
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _build_mcp_components(tmp_path) -> tuple[FastMCP, McpTokenStore, McpBusyTracker]:
|
||||
manager = DeviceManager()
|
||||
status_tracker = AgentStatusTracker()
|
||||
tracker = McpBusyTracker()
|
||||
token_store = McpTokenStore(tmp_path / "host_mcp_token.json")
|
||||
server = build_mcp_server(
|
||||
manager=manager,
|
||||
mcp_busy_tracker=tracker,
|
||||
status_tracker=status_tracker,
|
||||
)
|
||||
return server, token_store, tracker
|
||||
|
||||
|
||||
def test_console_app_mounts_mcp_when_all_components_provided(tmp_path) -> None:
|
||||
server, token_store, tracker = _build_mcp_components(tmp_path)
|
||||
client, _ = _build_client(
|
||||
tmp_path,
|
||||
mcp_server=server,
|
||||
mcp_token_store=token_store,
|
||||
mcp_busy_tracker=tracker,
|
||||
)
|
||||
# Without auth, the bearer middleware should respond 401 — not 404.
|
||||
resp = client.post("/mcp/", json={"jsonrpc": "2.0", "method": "ping", "id": 1})
|
||||
assert resp.status_code != 404
|
||||
|
||||
|
||||
def test_console_app_does_not_mount_mcp_when_components_missing(tmp_path) -> None:
|
||||
client, _ = _build_client(tmp_path)
|
||||
resp = client.post("/mcp/", json={"jsonrpc": "2.0", "method": "ping", "id": 1})
|
||||
assert resp.status_code == 404
|
||||
|
||||
|
||||
def test_api_status_includes_mcp_busy_devices(tmp_path) -> None:
|
||||
server, token_store, tracker = _build_mcp_components(tmp_path)
|
||||
# Acquire a lease without going through HTTP — tracker exposes a direct API.
|
||||
tracker.acquire("phone-1", "test-session")
|
||||
client, _ = _build_client(
|
||||
tmp_path,
|
||||
mcp_server=server,
|
||||
mcp_token_store=token_store,
|
||||
mcp_busy_tracker=tracker,
|
||||
)
|
||||
_login(client)
|
||||
|
||||
response = client.get("/api/status")
|
||||
assert response.status_code == 200
|
||||
body = response.json()
|
||||
assert "mcp_busy_devices" in body
|
||||
assert "phone-1" in body["mcp_busy_devices"]
|
||||
assert body["mcp_endpoint"] == "/mcp"
|
||||
|
||||
|
||||
def test_api_status_omits_mcp_fields_when_components_missing(tmp_path) -> None:
|
||||
client, _ = _build_client(tmp_path)
|
||||
_login(client)
|
||||
|
||||
response = client.get("/api/status")
|
||||
assert response.status_code == 200
|
||||
body = response.json()
|
||||
assert body["mcp_busy_devices"] == []
|
||||
assert body["mcp_endpoint"] is None
|
||||
|
||||
|
||||
def test_dashboard_renders_mcp_status_row(tmp_path) -> None:
|
||||
server, token_store, tracker = _build_mcp_components(tmp_path)
|
||||
tracker.acquire("phone-1", "test-session")
|
||||
client, _ = _build_client(
|
||||
tmp_path,
|
||||
mcp_server=server,
|
||||
mcp_token_store=token_store,
|
||||
mcp_busy_tracker=tracker,
|
||||
)
|
||||
_login(client)
|
||||
|
||||
response = client.get("/")
|
||||
assert response.status_code == 200
|
||||
text = response.text
|
||||
assert "<td>MCP</td>" in text
|
||||
assert "/mcp" in text
|
||||
assert "phone-1" in text
|
||||
|
||||
|
||||
def test_dashboard_renders_mcp_not_configured_when_components_missing(
|
||||
tmp_path,
|
||||
) -> None:
|
||||
client, _ = _build_client(tmp_path)
|
||||
_login(client)
|
||||
|
||||
response = client.get("/")
|
||||
assert response.status_code == 200
|
||||
text = response.text
|
||||
assert "<td>MCP</td>" in text
|
||||
assert "not configured" in text
|
||||
|
||||
|
||||
def test_mcp_endpoint_unauthorized_without_bearer_token(tmp_path) -> None:
|
||||
server, token_store, tracker = _build_mcp_components(tmp_path)
|
||||
client, _ = _build_client(
|
||||
tmp_path,
|
||||
mcp_server=server,
|
||||
mcp_token_store=token_store,
|
||||
mcp_busy_tracker=tracker,
|
||||
)
|
||||
resp = client.post("/mcp/", json={"jsonrpc": "2.0", "method": "ping", "id": 1})
|
||||
assert resp.status_code == 401
|
||||
|
||||
|
||||
def test_mcp_endpoint_rejects_invalid_bearer_token(tmp_path) -> None:
|
||||
server, token_store, tracker = _build_mcp_components(tmp_path)
|
||||
client, _ = _build_client(
|
||||
tmp_path,
|
||||
mcp_server=server,
|
||||
mcp_token_store=token_store,
|
||||
mcp_busy_tracker=tracker,
|
||||
)
|
||||
resp = client.post(
|
||||
"/mcp/",
|
||||
headers={"Authorization": "Bearer not-the-real-token"},
|
||||
json={"jsonrpc": "2.0", "method": "ping", "id": 1},
|
||||
)
|
||||
assert resp.status_code == 401
|
||||
|
||||
@@ -0,0 +1,415 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import UTC, datetime
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
from cloud.internal_api.models import AssignmentModel
|
||||
from device.manager import DeviceManager
|
||||
from driver.base import Driver
|
||||
from host_agent.mcp_lock import McpBusyTracker
|
||||
from host_agent.status import AgentStatusTracker
|
||||
from host_agent.web.mcp import (
|
||||
McpDeviceBusyError,
|
||||
_call_tool_sync,
|
||||
_current_session_id,
|
||||
build_mcp_server,
|
||||
)
|
||||
from mcp.server.fastmcp import Context
|
||||
from mcp.shared.context import RequestContext
|
||||
|
||||
|
||||
class _FakeDriver(Driver):
|
||||
"""Minimal driver. connect/screenshot/tap are exercised; remaining abstract
|
||||
methods are stubbed to satisfy Driver's ABC contract."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.taps: list[tuple[float, float]] = []
|
||||
|
||||
def connect(self) -> None:
|
||||
return None
|
||||
|
||||
def disconnect(self) -> None:
|
||||
return None
|
||||
|
||||
def screenshot(self) -> bytes:
|
||||
return b"fake"
|
||||
|
||||
def tap(self, x: float, y: float) -> None:
|
||||
self.taps.append((x, y))
|
||||
|
||||
def long_press(self, x: float, y: float, duration_ms: int = 1200) -> None:
|
||||
return None
|
||||
|
||||
def swipe(
|
||||
self,
|
||||
start_x: float,
|
||||
start_y: float,
|
||||
end_x: float,
|
||||
end_y: float,
|
||||
duration_ms: int = 500,
|
||||
) -> None:
|
||||
return None
|
||||
|
||||
def swipe_path(
|
||||
self, waypoints: list[tuple[float, float]], duration_ms: int
|
||||
) -> None:
|
||||
return None
|
||||
|
||||
def double_tap(self, x: float, y: float, interval_ms: int = 80) -> None:
|
||||
return None
|
||||
|
||||
def input(self, text: str) -> None:
|
||||
return None
|
||||
|
||||
def launch(self, app_id: str) -> None:
|
||||
return None
|
||||
|
||||
def terminate(self, app_id: str) -> None:
|
||||
return None
|
||||
|
||||
def tree(self) -> Any:
|
||||
return None
|
||||
|
||||
def home(self) -> None:
|
||||
return None
|
||||
|
||||
def lock(self) -> None:
|
||||
return None
|
||||
|
||||
def unlock(self) -> None:
|
||||
return None
|
||||
|
||||
|
||||
def _make_manager_with_device(device_id: str = "phone-1") -> DeviceManager:
|
||||
manager = DeviceManager()
|
||||
manager.register_device(
|
||||
device_id=device_id,
|
||||
driver_factory=lambda: _FakeDriver(),
|
||||
name=device_id,
|
||||
)
|
||||
manager.connect(device_id)
|
||||
return manager
|
||||
|
||||
|
||||
def test_build_mcp_server_returns_fastmcp_instance() -> None:
|
||||
from mcp.server.fastmcp import FastMCP
|
||||
|
||||
manager = _make_manager_with_device()
|
||||
tracker = McpBusyTracker()
|
||||
status = AgentStatusTracker()
|
||||
server = build_mcp_server(
|
||||
manager=manager, mcp_busy_tracker=tracker, status_tracker=status
|
||||
)
|
||||
assert isinstance(server, FastMCP)
|
||||
|
||||
|
||||
def test_call_tool_succeeds_when_device_is_free() -> None:
|
||||
manager = _make_manager_with_device()
|
||||
tracker = McpBusyTracker()
|
||||
status = AgentStatusTracker()
|
||||
server = build_mcp_server(
|
||||
manager=manager, mcp_busy_tracker=tracker, status_tracker=status
|
||||
)
|
||||
result = _call_tool_sync(
|
||||
server, "take_screenshot", {"device_id": "phone-1"}, session_id="sess-a"
|
||||
)
|
||||
assert result["ok"] is True
|
||||
assert "phone-1" in tracker.busy_device_ids()
|
||||
|
||||
|
||||
def test_call_tool_fails_when_cloud_uses_device() -> None:
|
||||
"""AgentStatusTracker.current_assignment.device_id matches -> busy."""
|
||||
manager = _make_manager_with_device()
|
||||
tracker = McpBusyTracker()
|
||||
status = AgentStatusTracker()
|
||||
status.mark_assignment_started(
|
||||
AssignmentModel(
|
||||
task_id="t1",
|
||||
attempt=1,
|
||||
lease_id="l1",
|
||||
lease_expires_at=datetime.now(UTC),
|
||||
host_id="h1",
|
||||
device_id="phone-1",
|
||||
goal="cloud task",
|
||||
)
|
||||
)
|
||||
server = build_mcp_server(
|
||||
manager=manager, mcp_busy_tracker=tracker, status_tracker=status
|
||||
)
|
||||
with pytest.raises(McpDeviceBusyError) as exc:
|
||||
_call_tool_sync(
|
||||
server,
|
||||
"take_screenshot",
|
||||
{"device_id": "phone-1"},
|
||||
session_id="sess-a",
|
||||
)
|
||||
assert exc.value.device_id == "phone-1"
|
||||
assert exc.value.busy_owner == "cloud_assignment"
|
||||
|
||||
|
||||
def test_call_tool_fails_when_another_mcp_session_holds_device() -> None:
|
||||
manager = _make_manager_with_device()
|
||||
tracker = McpBusyTracker()
|
||||
status = AgentStatusTracker()
|
||||
# Pre-acquire as a different session.
|
||||
tracker.acquire("phone-1", "sess-other")
|
||||
server = build_mcp_server(
|
||||
manager=manager, mcp_busy_tracker=tracker, status_tracker=status
|
||||
)
|
||||
with pytest.raises(McpDeviceBusyError) as exc:
|
||||
_call_tool_sync(
|
||||
server,
|
||||
"take_screenshot",
|
||||
{"device_id": "phone-1"},
|
||||
session_id="sess-a",
|
||||
)
|
||||
assert exc.value.busy_owner.startswith("mcp_session:")
|
||||
|
||||
|
||||
def test_call_tool_renews_when_same_session_already_holds() -> None:
|
||||
manager = _make_manager_with_device()
|
||||
tracker = McpBusyTracker()
|
||||
status = AgentStatusTracker()
|
||||
server = build_mcp_server(
|
||||
manager=manager, mcp_busy_tracker=tracker, status_tracker=status
|
||||
)
|
||||
_call_tool_sync(
|
||||
server, "take_screenshot", {"device_id": "phone-1"}, session_id="sess-a"
|
||||
)
|
||||
# Second call from the same session should succeed.
|
||||
result = _call_tool_sync(
|
||||
server, "take_screenshot", {"device_id": "phone-1"}, session_id="sess-a"
|
||||
)
|
||||
assert result["ok"] is True
|
||||
|
||||
|
||||
def test_list_devices_uses_display_status() -> None:
|
||||
"""Connected-but-idle devices report as 'connected', not 'busy'."""
|
||||
manager = _make_manager_with_device()
|
||||
tracker = McpBusyTracker()
|
||||
status = AgentStatusTracker()
|
||||
server = build_mcp_server(
|
||||
manager=manager, mcp_busy_tracker=tracker, status_tracker=status
|
||||
)
|
||||
result = _call_tool_sync(server, "list_devices", {}, session_id="sess-a")
|
||||
assert isinstance(result, list)
|
||||
assert result[0]["status"] == "connected"
|
||||
|
||||
|
||||
def test_unknown_device_returns_semantic_error_dict() -> None:
|
||||
"""take_screenshot against an unknown device returns the api-errors semantic
|
||||
error dict (``ok=False, error="device not found"``) rather than raising.
|
||||
|
||||
Note: this test adapts the brief's exception-assertion semantics to the
|
||||
actual behavior of ``call_with_semantic_errors`` in ``api/errors.py`` —
|
||||
the brief's expectation that an exception is raised here is incorrect for
|
||||
the current handler implementation."""
|
||||
manager = _make_manager_with_device()
|
||||
tracker = McpBusyTracker()
|
||||
status = AgentStatusTracker()
|
||||
server = build_mcp_server(
|
||||
manager=manager, mcp_busy_tracker=tracker, status_tracker=status
|
||||
)
|
||||
result = _call_tool_sync(
|
||||
server,
|
||||
"take_screenshot",
|
||||
{"device_id": "does-not-exist"},
|
||||
session_id="sess-a",
|
||||
)
|
||||
assert isinstance(result, dict)
|
||||
assert result["ok"] is False
|
||||
assert "device" in result["error"].lower()
|
||||
|
||||
|
||||
def test_manager_required_for_build_mcp_server() -> None:
|
||||
"""Reinforces D12 — build_mcp_server requires manager as keyword-only."""
|
||||
import inspect
|
||||
|
||||
sig = inspect.signature(build_mcp_server)
|
||||
assert sig.parameters["manager"].kind == inspect.Parameter.KEYWORD_ONLY
|
||||
|
||||
|
||||
def _fake_ctx(session_obj: object) -> Context:
|
||||
"""Build a Context whose ``session`` attribute returns ``session_obj``.
|
||||
|
||||
Context's ``session`` is a property backed by ``request_context.session``;
|
||||
we construct a minimal ``RequestContext`` and set it as the private
|
||||
``_request_context`` field. The pydantic public API doesn't expose a
|
||||
setter for ``session``, so we use ``object.__setattr__`` on the private
|
||||
backing field.
|
||||
"""
|
||||
ctx = Context.model_construct()
|
||||
request_ctx = RequestContext(
|
||||
request_id="req-test",
|
||||
meta=None,
|
||||
session=session_obj,
|
||||
lifespan_context=None,
|
||||
)
|
||||
object.__setattr__(ctx, "_request_context", request_ctx)
|
||||
return ctx
|
||||
|
||||
|
||||
def test_current_session_id_is_stable_across_calls_same_session() -> None:
|
||||
"""Production-path identity: two tool calls from the same MCP session
|
||||
must yield the same session_id so the busy tracker can renew the lease.
|
||||
|
||||
This exercises the ``Context.session`` code path (NOT the
|
||||
``_TEST_SESSION_ID`` fallback used by ``_call_tool_sync``)."""
|
||||
sentinel_session = object()
|
||||
ctx = _fake_ctx(sentinel_session)
|
||||
first = _current_session_id(ctx)
|
||||
second = _current_session_id(ctx)
|
||||
assert first == second
|
||||
assert first.startswith("mcp_session:")
|
||||
# Object identity of the underlying ServerSession is the key — verifies
|
||||
# we use id(ctx.session) rather than e.g. ctx.request_id.
|
||||
assert first == f"mcp_session:{id(sentinel_session)}"
|
||||
|
||||
|
||||
def test_current_session_id_differs_across_sessions() -> None:
|
||||
"""Two different MCP sessions (distinct ServerSession objects) must
|
||||
produce distinct session_ids so the busy tracker can isolate them."""
|
||||
sess_a = object()
|
||||
sess_b = object()
|
||||
assert _current_session_id(_fake_ctx(sess_a)) != _current_session_id(
|
||||
_fake_ctx(sess_b)
|
||||
)
|
||||
|
||||
|
||||
def test_current_session_id_falls_back_when_no_context() -> None:
|
||||
"""When no Context is available (e.g. outside a FastMCP request lifecycle,
|
||||
or via ``_call_tool_sync`` which omits the ctx kwarg), the test
|
||||
contextvars override provides the session_id."""
|
||||
token = None
|
||||
try:
|
||||
from host_agent.web import mcp as mcp_mod
|
||||
|
||||
token = mcp_mod._TEST_SESSION_ID.set("test-session-xyz")
|
||||
assert _current_session_id(None) == "test-session-xyz"
|
||||
finally:
|
||||
if token is not None:
|
||||
from host_agent.web import mcp as mcp_mod
|
||||
|
||||
mcp_mod._TEST_SESSION_ID.reset(token)
|
||||
|
||||
|
||||
def test_wrapped_tool_accepts_context_kwarg() -> None:
|
||||
"""The wrapper registered on FastMCP must declare a ``ctx`` parameter so
|
||||
FastMCP injects the live Context (and ``tool.context_kwarg`` is set to
|
||||
``"ctx"``). Without this, FastMCP never injects context and we fall
|
||||
back to the empty test default — the production bug this PR fixes."""
|
||||
manager = _make_manager_with_device()
|
||||
tracker = McpBusyTracker()
|
||||
status = AgentStatusTracker()
|
||||
server = build_mcp_server(
|
||||
manager=manager, mcp_busy_tracker=tracker, status_tracker=status
|
||||
)
|
||||
tool_manager = server._tool_manager # type: ignore[attr-defined]
|
||||
tool = tool_manager.get_tool("take_screenshot")
|
||||
assert tool is not None
|
||||
assert tool.context_kwarg == "ctx"
|
||||
|
||||
|
||||
def test_busy_error_wire_shape_is_calltoolresult_iserror() -> None:
|
||||
"""Regression test for spec §7 — busy errors must be visible on the wire.
|
||||
|
||||
mcp SDK 1.28.1's ``Tool.run`` wraps every non-``UrlElicitationRequiredError``
|
||||
exception (including ``McpError`` and our ``McpDeviceBusyError``) into
|
||||
``ToolError`` (see ``mcp/server/fastmcp/tools/base.py``). The lowlevel
|
||||
``call_tool`` handler then builds a ``CallToolResult(isError=True,
|
||||
content=[TextContent(...)])`` (see
|
||||
``mcp/server/lowlevel/server.py::_make_error_result``). There is no public
|
||||
path that surfaces JSON-RPC ``-32000`` + structured ``data.busy_owner`` from
|
||||
a tool call site — the SDK's wire contract for tool errors is the
|
||||
``isError=true`` flag plus text content. This test pins the wire shape so
|
||||
any future SDK upgrade that exposes a true JSON-RPC error path is caught."""
|
||||
import asyncio
|
||||
|
||||
import mcp.types as types
|
||||
from mcp.server.fastmcp import FastMCP
|
||||
from mcp.server.fastmcp.exceptions import ToolError
|
||||
|
||||
manager = _make_manager_with_device()
|
||||
tracker = McpBusyTracker()
|
||||
status = AgentStatusTracker()
|
||||
tracker.acquire("phone-1", "sess-other") # different session holds the device
|
||||
|
||||
server: FastMCP = build_mcp_server(
|
||||
manager=manager, mcp_busy_tracker=tracker, status_tracker=status
|
||||
)
|
||||
tool = server._tool_manager.get_tool("take_screenshot") # type: ignore[attr-defined]
|
||||
assert tool is not None
|
||||
|
||||
sentinel_session = object()
|
||||
ctx = _fake_ctx(sentinel_session)
|
||||
|
||||
with pytest.raises(ToolError) as tool_exc:
|
||||
asyncio.run(tool.run({"device_id": "phone-1"}, context=ctx))
|
||||
|
||||
# ToolError text carries the original exception message verbatim,
|
||||
# which is what the lowlevel handler copies into TextContent.
|
||||
message = str(tool_exc.value)
|
||||
assert "phone-1" in message
|
||||
assert "busy" in message
|
||||
assert "mcp_session:sess-oth" in message # truncated busy_owner
|
||||
|
||||
# The lowlevel handler converts any exception into a CallToolResult
|
||||
# with isError=True (mcp SDK 1.28.1 — not a JSON-RPC error envelope).
|
||||
# We invoke the SDK helper directly to lock the wire contract.
|
||||
from mcp.server.lowlevel.server import Server as LowlevelServer
|
||||
|
||||
lowlevel = LowlevelServer("test-lowlevel")
|
||||
error_result = lowlevel._make_error_result(message) # type: ignore[attr-defined]
|
||||
inner = error_result.root
|
||||
assert isinstance(inner, types.CallToolResult)
|
||||
assert inner.isError is True
|
||||
assert len(inner.content) == 1
|
||||
text_block = inner.content[0]
|
||||
assert isinstance(text_block, types.TextContent)
|
||||
assert text_block.text == message
|
||||
# And confirm the wire shape is NOT a JSON-RPC error envelope — that
|
||||
# would require code=-32000 + data.busy_owner, which is not exposed
|
||||
# in mcp SDK 1.28.1 for tool-call errors.
|
||||
assert not hasattr(inner, "code")
|
||||
assert inner.structuredContent is None
|
||||
|
||||
|
||||
def test_busy_error_text_includes_cloud_assignment_owner() -> None:
|
||||
"""Same wire-shape test for the cloud_assignment branch — verifies the
|
||||
human-readable busy_owner value (the only place to surface it given the
|
||||
SDK forces tool errors into CallToolResult.isError=true) is correct."""
|
||||
import asyncio
|
||||
|
||||
from mcp.server.fastmcp.exceptions import ToolError
|
||||
|
||||
manager = _make_manager_with_device()
|
||||
tracker = McpBusyTracker()
|
||||
status = AgentStatusTracker()
|
||||
status.mark_assignment_started(
|
||||
AssignmentModel(
|
||||
task_id="t1",
|
||||
attempt=1,
|
||||
lease_id="l1",
|
||||
lease_expires_at=datetime.now(UTC),
|
||||
host_id="h1",
|
||||
device_id="phone-1",
|
||||
goal="cloud task",
|
||||
)
|
||||
)
|
||||
|
||||
server = build_mcp_server(
|
||||
manager=manager, mcp_busy_tracker=tracker, status_tracker=status
|
||||
)
|
||||
tool = server._tool_manager.get_tool("take_screenshot") # type: ignore[attr-defined]
|
||||
assert tool is not None
|
||||
|
||||
sentinel_session = object()
|
||||
ctx = _fake_ctx(sentinel_session)
|
||||
|
||||
with pytest.raises(ToolError) as tool_exc:
|
||||
asyncio.run(tool.run({"device_id": "phone-1"}, context=ctx))
|
||||
assert "cloud_assignment" in str(tool_exc.value)
|
||||
assert "phone-1" in str(tool_exc.value)
|
||||
assert "busy" in str(tool_exc.value)
|
||||
@@ -122,6 +122,26 @@ class DeviceManager:
|
||||
self._drivers.pop(device_id, None)
|
||||
self._set_status(device_id, "offline" if offline else "error")
|
||||
|
||||
def probe(self, device_id: str) -> bool:
|
||||
"""Check whether an established driver is still reachable."""
|
||||
with self._lock:
|
||||
self._device(device_id)
|
||||
driver = self._drivers.get(device_id)
|
||||
if driver is None:
|
||||
return False
|
||||
try:
|
||||
health_check = getattr(driver, "health_check", None)
|
||||
if callable(health_check):
|
||||
health_check()
|
||||
else:
|
||||
# Compatibility for drivers implemented before health_check
|
||||
# existed. Built-in drivers use the non-screen health check.
|
||||
driver.screenshot()
|
||||
except Exception:
|
||||
self.mark_error(device_id, offline=True)
|
||||
return False
|
||||
return True
|
||||
|
||||
def active_driver(self, device_id: str | None = None) -> Driver:
|
||||
with self._lock:
|
||||
if device_id is None:
|
||||
|
||||
@@ -527,6 +527,20 @@ Appium server 默认使用 4723;WDA 通常使用 8100。多设备必须为每
|
||||
input、launch 和 UI tree 验证基础控制,再单独处理 PaddleOCR/PaddlePaddle 的 macOS
|
||||
wheel 与 Apple Silicon 兼容性。
|
||||
|
||||
## MCP server (Hermes Agent integration)
|
||||
|
||||
Host-agent now exposes an MCP server on the same port as the local
|
||||
console (`127.0.0.1:8765/mcp`). To drive your iPhone from Hermes Agent
|
||||
or any MCP-compatible client:
|
||||
|
||||
1. Start host-agent normally.
|
||||
2. Get the bearer token: `device-host-agent mcp-token`.
|
||||
3. Configure Hermes per `docs/MCP_INTEGRATION.md`.
|
||||
|
||||
The MCP path reuses the same WDA session that the cloud worker uses.
|
||||
Per-device locking prevents both sides from driving the same device at
|
||||
once; see `docs/MCP_INTEGRATION.md` for the full concurrency model.
|
||||
|
||||
## 12. 完成检查表
|
||||
|
||||
- [ ] Xcode 能看到已解锁的 iPhone。
|
||||
|
||||
@@ -0,0 +1,125 @@
|
||||
# Host-Agent MCP Server Integration
|
||||
|
||||
The host-agent process exposes a Streamable HTTP MCP server on the same
|
||||
port as the local console (default `127.0.0.1:8765`), at path `/mcp`. This
|
||||
lets any MCP-compatible client — Hermes Agent, Claude Desktop, custom
|
||||
scripts using the `mcp` Python SDK — drive devices directly through the
|
||||
same `DeviceManager` the cloud worker uses.
|
||||
|
||||
## Prerequisites
|
||||
|
||||
- Host-agent built from this repo (see `docs/MACOS_IPHONE_SETUP.md`).
|
||||
- An MCP client that supports the Streamable HTTP transport (mcp SDK
|
||||
1.20+ on the client side).
|
||||
|
||||
## Get the bearer token
|
||||
|
||||
The first time host-agent starts after this feature ships, it generates
|
||||
a random bearer token and writes it to:
|
||||
|
||||
<identity_path.parent>/host_mcp_token.json
|
||||
|
||||
(Default: `tasks/host_mcp_token.json` next to `host_identity.json`.)
|
||||
|
||||
To print it for copy/paste:
|
||||
|
||||
device-host-agent mcp-token
|
||||
|
||||
To rotate: delete the file and restart host-agent. Old tokens stop
|
||||
working immediately.
|
||||
|
||||
## Hermes Agent configuration
|
||||
|
||||
Add to `~/.hermes/config.yaml`:
|
||||
|
||||
```yaml
|
||||
mcp_servers:
|
||||
apex_device:
|
||||
url: "http://127.0.0.1:8765/mcp"
|
||||
headers:
|
||||
Authorization: "Bearer <paste-token-here>"
|
||||
```
|
||||
|
||||
Start (or restart) Hermes. Verify by asking Hermes to list devices:
|
||||
|
||||
> Use the apex_device MCP to list connected devices.
|
||||
|
||||
## Tools exposed
|
||||
|
||||
All 11 device tools from `api/mcp.py`:
|
||||
|
||||
- `take_screenshot(device_id?)`
|
||||
- `tap(x, y, device_id?)`
|
||||
- `swipe(start_x, start_y, end_x, end_y, duration_ms?, device_id?)`
|
||||
- `input_text(text, device_id?)`
|
||||
- `launch_app(app_id, device_id?)`
|
||||
- `find_text(query, device_id?)`
|
||||
- `find_icon(name, device_id?)`
|
||||
- `get_ui_tree(device_id?, include_app_info?)`
|
||||
- `describe_screen(device_id?)`
|
||||
- `list_devices()`
|
||||
- `device_status(device_id)`
|
||||
|
||||
## Concurrency model
|
||||
|
||||
- The cloud worker and MCP clients share the same `DeviceManager`.
|
||||
- Per-device, session-level locking: the first caller (cloud or MCP) to
|
||||
touch a device holds it; the other side sees a busy error.
|
||||
- MCP sessions hold their lock until **20 seconds of inactivity**
|
||||
(the `McpBusyTracker` default TTL). The mcp SDK 1.28.1 does not expose
|
||||
a per-session shutdown callback, so a clean Hermes disconnect is also
|
||||
recovered via the 20s TTL sweep — see the implementation note in
|
||||
spec §6.5. Cloud assignments hold theirs until the assignment
|
||||
terminates.
|
||||
- The cloud scheduler is told about MCP-held devices via the heartbeat
|
||||
`mcp_busy_device_ids` field, so it normally won't even try to dispatch
|
||||
to them. A 30-second window exists between an MCP acquire and the next
|
||||
heartbeat; during that window cloud may dispatch, and the host-agent
|
||||
will fail-fast the assignment with `failure_reason="device held by an
|
||||
active MCP session"`.
|
||||
|
||||
## Network binding
|
||||
|
||||
The MCP endpoint is bound to the same address as the local console. By
|
||||
default this is `127.0.0.1` (loopback only). To expose on a different
|
||||
interface, set `HOST_AGENT_CONSOLE_BIND_HOST` AND
|
||||
`HOST_AGENT_CONSOLE_ALLOW_NON_LOOPBACK=true` — both are required. This
|
||||
is the same escape hatch the local console uses; there is no MCP-only
|
||||
override.
|
||||
|
||||
## Error responses
|
||||
|
||||
The mcp SDK 1.28.1 forces tool errors into `CallToolResult(isError=true,
|
||||
content=[TextContent(message)])` — there is no public path that surfaces
|
||||
JSON-RPC `-32000` with a structured `data.busy_owner` field from a tool
|
||||
call site. The busy-owner value lives inside the text content (full
|
||||
string for `cloud_assignment`, truncated session_id prefix for
|
||||
`mcp_session:` collisions).
|
||||
|
||||
| Condition | JSON-RPC envelope | `result.content[0].text` |
|
||||
|---|---|---|
|
||||
| Missing/wrong bearer token | HTTP 401 (transport-level) | `{"error": "invalid token"}` + `WWW-Authenticate: Bearer` |
|
||||
| Device busy (cloud) | `result.isError = true` | `"device <X> is busy (held by cloud assignment)"` |
|
||||
| Device busy (other MCP) | `result.isError = true` | `"device <X> is busy (held by mcp_session:<8-char-prefix>)"` |
|
||||
| Unknown device | `result.isError = false` | JSON `{"ok": false, "error": "device not found: <X>"}` |
|
||||
| Tool error | `result.isError = true` | `"Error executing tool <name>: <original-message>"` |
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
- **`list_devices` returns `[]`**: no devices registered. Use the local
|
||||
console at `http://127.0.0.1:8765/` to add one (Login → Devices).
|
||||
- **`device X is busy` even when cloud console says device is idle**:
|
||||
check whether another MCP session is holding it. The local console
|
||||
dashboard shows active MCP sessions and held device_ids.
|
||||
- **Token verification fails after restart**: confirm you copied the
|
||||
token from the current `host_mcp_token.json`, not an older one.
|
||||
Rotation = delete file + restart.
|
||||
|
||||
## Out of scope (current version)
|
||||
|
||||
- `wait_until_usable` MCP tool: implemented internally but not exposed.
|
||||
MVP callers must handle busy errors themselves.
|
||||
- MCP call history in the local console: only current state is surfaced,
|
||||
not a call log.
|
||||
- Token rotation CLI: use delete-and-restart for now.
|
||||
- Non-loopback binding without explicit opt-in.
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,459 @@
|
||||
# Host-Agent MCP Server — Design Spec
|
||||
|
||||
- **Date**: 2026-07-21
|
||||
- **Status**: Draft, pending user review
|
||||
- **Owner**: Jerry Yan
|
||||
- **Target package**: `apps/device-host-agent`(主)+ `packages/cloud-platform`(schema/scheduler 扩展)+ `api/mcp.py`(一处签名收紧)
|
||||
|
||||
## 1. Overview
|
||||
|
||||
在 host-agent 进程内挂载 Streamable HTTP MCP server,让 Hermes Agent(或任意 MCP 客户端)作为外部大脑直接驱动 host-agent 管理的设备,与既有的 Cloud Control Plane 长轮询 worker 路径并存。
|
||||
|
||||
**一句话**:Hermes 当大脑,host-agent 当手;同一个 host-agent 进程同时服务两条职责,per-device 互斥,先到先得。
|
||||
|
||||
## 2. Background & Motivation
|
||||
|
||||
- 当前 host-agent 只能通过 Cloud Control Plane 派发任务驱动;操作员想直接用本地 Hermes Agent 临时操控设备时,必须先在 cloud 侧创建 task,路径长、延迟高、依赖网络。
|
||||
- `api/mcp.py::create_mcp_server()` 已经用 FastMCP 暴露了 11 个设备工具(`take_screenshot`/`tap`/`swipe`/`input_text`/`launch_app`/`find_text`/`find_icon`/`get_ui_tree`/`describe_screen`/`list_devices`/`device_status`),但只服务于 Runtime 进程(port 8000),且没有可运行入口被 host-agent 复用。
|
||||
- Hermes Agent 官方支持 HTTP MCP client(见 [Hermes MCP docs](https://hermes-agent.nousresearch.com/docs/user-guide/features/mcp)),配置长这样:
|
||||
```yaml
|
||||
mcp_servers:
|
||||
apex_device:
|
||||
url: "http://127.0.0.1:8765/mcp"
|
||||
headers:
|
||||
Authorization: "Bearer <token>"
|
||||
```
|
||||
- uv.lock 已锁定 `mcp==1.28.1`,支持 Streamable HTTP transport。
|
||||
- 既有的 `host-agent-local-console` 已立先例:host-agent 进程内 always-on 跑 FastAPI + uvicorn(默认 port 8765,loopback only),本次 MCP server 沿用同一 server、同一端口,仅新增一个 `/mcp` mount。
|
||||
|
||||
## 3. Locked Decisions(来自 grill 阶段)
|
||||
|
||||
| # | 决策 | 备注 |
|
||||
|---|---|---|
|
||||
| D1 | Transport = Streamable HTTP | mcp SDK 1.28.1 支持 |
|
||||
| D2 | Mount 路径 = `/mcp`,与 console 同端口(默认 8765) | FastAPI `app.mount()` |
|
||||
| D3 | 网络 = loopback only,沿用 `HOST_AGENT_CONSOLE_ALLOW_NON_LOOPBACK` 双保险 | 不新增 bind 配置 |
|
||||
| D4 | 鉴权 = 独立 MCP bearer token,存 `host_mcp_token.json` | 文件路径 = `config.identity_path.parent / "host_mcp_token.json"` |
|
||||
| D5 | 并发锁 = per-device,先到先得 + fail-fast | 不做 wait 队列 |
|
||||
| D6 | 设备状态映射 = 走 `_device_display_status()` 同款逻辑 | 避免"所有连上的设备看起来都 busy" |
|
||||
| D7 | Cloud worker 与 MCP server 在同一 host-agent 进程并存 | 不互斥,共享 `DeviceManager` |
|
||||
| D8 | Cloud ↔ MCP 协调 = 心跳上报 `mcp_busy_device_ids`,cloud scheduler 跳过 | 心跳 schema 扩展,cloud 侧 _matches() 一处改动 |
|
||||
| D9 | MCP session 级 lazy acquire 锁,20s TTL 兜底(mcp SDK 1.28.1 无 session-end callback,见 §6.5) | session_id 来自 FastMCP 上下文 |
|
||||
| D10 | Skill catalog 工具 MVP 不暴露,保留 `create_mcp_server(skill_catalog_store=...)` 参数化挂载点 | 未来可加 mutating 工具 |
|
||||
| D11 | `wait_until_usable` 方法实现 + 单测,但调用方不接入 | 预留能力,MVP 全部 fail-fast |
|
||||
| D12 | `tool_handlers(manager)` 改为必传 | 已 grep 确认无调用方依赖 None 默认,根治 `DeviceNotFoundError` 类静默回退地雷 |
|
||||
|
||||
## 4. Architecture
|
||||
|
||||
### 4.1 总览
|
||||
|
||||
```
|
||||
┌──────────────────────────── host-agent 进程 ────────────────────────────┐
|
||||
│ │
|
||||
│ run_async() 主循环 │
|
||||
│ ├─ HeartbeatSynchronizer ──reads──→ McpBusyTracker ──┐ │
|
||||
│ │ (已有) payload.mcp_busy_device_ids │ │
|
||||
│ ├─ AssignmentProcessor ──reads──→ AgentStatusTracker │ │
|
||||
│ │ (已有,cloud 路径) .current_assignment │
|
||||
│ │ └─ AssignmentExecutor.execute() 启动前 fail-fast │ │
|
||||
│ │ 检查 mcp_busy_tracker.busy_device_ids() │ │
|
||||
│ └─ Console uvicorn server(已有 port 8765) │ │
|
||||
│ └─ FastAPI app │ │
|
||||
│ ├─ /login, /devices, /tasks, ... (已有) │ │
|
||||
│ └─ Mount("/mcp") │ │
|
||||
│ └─ FastMCP.streamable_http_app() │ │
|
||||
│ ├─ BearerAuthMiddleware ──verify──→ McpTokenStore │
|
||||
│ └─ 11 tools(包装层) │ │
|
||||
│ ├─ busy check ────────────────────┘ │
|
||||
│ │ (cloud 占用? McpBusyTracker 占用?) │
|
||||
│ ├─ acquire device lock(首次调用时) │
|
||||
│ ├─ tool_handlers(manager=<host_agent_manager>)[name]│
|
||||
│ └─ release on session end / TTL │
|
||||
│ │
|
||||
└─────────────────────────────────────────────────────────────────────────┘
|
||||
▲ ▲
|
||||
│ HTTP + Bearer │ HTTP heartbeat
|
||||
│ │ + mcp_busy_device_ids
|
||||
┌───────────┴───────────┐ ┌─────────┴─────────┐
|
||||
│ Hermes Agent (client) │ │ Cloud Control │
|
||||
│ ~/.hermes/config.yaml │ │ Plane (server) │
|
||||
│ mcp_servers.apex.url │ │ scheduler skips │
|
||||
│ = http://127.0.0.1 │ │ mcp-busy devices │
|
||||
│ :8765/mcp │ │ │
|
||||
└────────────────────────┘ └────────────────────┘
|
||||
```
|
||||
|
||||
### 4.2 关键不变量
|
||||
|
||||
- `DeviceManager` 单例;cloud 路径和 MCP 路径都通过显式参数注入。包装层在 `manager is None` 时抛错而非回退 `DEFAULT_MANAGER`(D12)。
|
||||
- 同一设备同一时刻只能被一边占用:cloud 占用 → MCP busy 错误;MCP 占用 → 心跳上报 → cloud scheduler 跳过;窗口期内冲突由 AssignmentExecutor 启动前 fail-fast 兜底。
|
||||
- `runtime/` / `api/` 包边界不破坏:所有新代码在 `host_agent/`;`api/mcp.py` 只做一处签名收紧。
|
||||
|
||||
## 5. Components
|
||||
|
||||
### 5.1 `host_agent/mcp_lock.py`(新)— `McpBusyTracker`
|
||||
|
||||
```python
|
||||
@dataclass(frozen=True)
|
||||
class McpDeviceLease:
|
||||
device_id: str
|
||||
session_id: str
|
||||
acquired_at: datetime
|
||||
last_seen_at: datetime
|
||||
|
||||
class McpBusyTracker:
|
||||
def __init__(self, *, ttl_seconds: float = 20.0, now=None) -> None: ...
|
||||
def acquire(self, device_id: str, session_id: str) -> bool: ...
|
||||
def renew(self, device_id: str, session_id: str) -> bool: ...
|
||||
def release(self, session_id: str) -> list[str]: ...
|
||||
def release_device(self, device_id: str, session_id: str) -> bool: ...
|
||||
def busy_device_ids(self) -> list[str]: ... # lazy sweep
|
||||
def snapshot(self) -> list[McpDeviceLease]: ... # lazy sweep, for UI
|
||||
def wait_until_usable(
|
||||
self, device_id: str, session_id: str, *,
|
||||
timeout: float, poll_interval: float = 1.0,
|
||||
cloud_busy_check: Callable[[], bool] | None = None,
|
||||
) -> bool: ... # 预留,MVP 不被调用
|
||||
```
|
||||
|
||||
- 进程内单实例,`threading.Lock` 保护。
|
||||
- TTL sweep 在 `busy_device_ids()` / `snapshot()` 调用时 lazy 执行。
|
||||
- Cloud 占用检查**不放这里**——保持单一职责;调用方在 acquire 前查 `AgentStatusTracker`。
|
||||
|
||||
### 5.2 `host_agent/mcp_token.py`(新)— `McpTokenStore`
|
||||
|
||||
```python
|
||||
@dataclass(frozen=True)
|
||||
class McpToken:
|
||||
version: int # =1
|
||||
token: str # secrets.token_urlsafe(32)
|
||||
created_at: datetime
|
||||
|
||||
class McpTokenStore:
|
||||
def __init__(self, path: Path, *, now=None) -> None: ...
|
||||
def load_or_create(self) -> McpToken: ... # atomic, 0o600
|
||||
def verify(self, presented: str) -> bool: ...
|
||||
```
|
||||
|
||||
- 文件路径:`config.identity_path.parent / "host_mcp_token.json"`。
|
||||
- JSON 格式:`{"version": 1, "token": "...", "created_at": "2026-07-21T..."}`。
|
||||
- 写文件用 `tempfile + os.replace` 原子重命名;权限 0o600。
|
||||
- 文件损坏或不可写 → `McpTokenStoreError`,不静默重新生成。
|
||||
|
||||
### 5.3 `host_agent/web/mcp_auth.py`(新)— Bearer Auth Middleware
|
||||
|
||||
```python
|
||||
class BearerAuthMiddleware(BaseHTTPMiddleware):
|
||||
def __init__(self, app, token_store: McpTokenStore) -> None: ...
|
||||
# 401 + WWW-Authenticate: Bearer + JSON {"error":"invalid token"} on fail
|
||||
```
|
||||
|
||||
- 挂在 FastMCP `streamable_http_app()` 外层(不是 console FastAPI 外层)。
|
||||
- 不写失败日志(避免暴力枚举刷屏);首次成功鉴权写 INFO。
|
||||
|
||||
### 5.4 `host_agent/web/mcp.py`(新)— `build_mcp_server()`
|
||||
|
||||
```python
|
||||
def build_mcp_server(
|
||||
*,
|
||||
manager: DeviceManager,
|
||||
mcp_busy_tracker: McpBusyTracker,
|
||||
status_tracker: AgentStatusTracker,
|
||||
) -> FastMCP:
|
||||
"""Wraps api.mcp.tool_handlers(manager=manager) with:
|
||||
- busy check (cloud OR mcp-busy → JSON-RPC -32000)
|
||||
- lazy device lock acquire on first call
|
||||
- lease renew on every call
|
||||
- status mapping for list_devices / device_status (display_status)"""
|
||||
```
|
||||
|
||||
- 单一 decorator 包装所有 11 个工具,避免重复。
|
||||
- `session_id` 来自 FastMCP streamable HTTP 上下文。
|
||||
- `list_devices` / `device_status` 不走 busy 检查(无 device_id 参数),但走 display status 映射。
|
||||
|
||||
### 5.5 改动的现有模块
|
||||
|
||||
| 模块 | 改动 |
|
||||
|---|---|
|
||||
| `host_agent/app.py::create_application()` | 构造 `McpTokenStore` / `McpBusyTracker` / `build_mcp_server()`,传给 `create_console_app()`;InstanceLock acquire 后、heartbeat 启动前完成 token 生成 |
|
||||
| `host_agent/web/app.py::create_console_app()` | 新参数 `mcp_server` / `mcp_token_store` / `mcp_busy_tracker`;`app.mount("/mcp", auth_wrapped(server.streamable_http_app()))`;`/api/status` 加 `mcp_busy_devices` 字段 |
|
||||
| `host_agent/web/templates/dashboard.html` | 一行 MCP 状态显示 |
|
||||
| `host_agent/cli.py` | 新子命令 `mcp-token` 打印当前 token(不存在则生成) |
|
||||
| `host_agent/heartbeat.py` | payload 加 `mcp_busy_device_ids` 字段,值来自 `mcp_busy_tracker.busy_device_ids()` |
|
||||
| `host_agent/assignment.py::AssignmentExecutor` | 构造参数加 `mcp_busy_tracker: McpBusyTracker \| None = None`(None 时跳过检查,保持向后兼容;host-agent 装配时必传);`execute()` 启动前(`_execute_goal` / `_execute_workflow` 实际跑之前)fail-fast 检查 `mcp_busy_tracker.busy_device_ids()`,命中则返回 `status="failed"` + `failure_reason="device held by active MCP session"` |
|
||||
| `cloud/internal_api/models.py` | Heartbeat payload schema 加 `mcp_busy_device_ids: list[str] = []` |
|
||||
| `cloud/scheduler.py::_matches()` | 心跳里的 `mcp_busy_device_ids` 内的 device 视为非 idle |
|
||||
| `api/mcp.py::tool_handlers` | 签名改为 `tool_handlers(*, manager: DeviceManager)`(必传,D12 根治) |
|
||||
|
||||
## 6. Data Flows
|
||||
|
||||
### 6.1 启动
|
||||
|
||||
```
|
||||
create_application()
|
||||
├─ InstanceLock acquire(已有)
|
||||
├─ resolve_host_identity()(已有)
|
||||
├─ DeviceManager 构造 + 设备注册(已有)
|
||||
├─ McpTokenStore(identity_path.parent / "host_mcp_token.json").load_or_create()
|
||||
│ └─ 首次:生成 token、atomic 写文件、INFO 日志 "MCP token generated at <path>"
|
||||
├─ McpBusyTracker(ttl_seconds=60)
|
||||
├─ build_mcp_server(manager=manager, mcp_busy_tracker=..., status_tracker=...)
|
||||
└─ create_console_app(..., mcp_server=server, mcp_token_store=..., mcp_busy_tracker=...)
|
||||
└─ app.mount("/mcp", auth_wrapped(server.streamable_http_app()))
|
||||
|
||||
run_async() 主循环启动(行为不变):
|
||||
├─ heartbeat_task: 每 30s 读 mcp_busy_tracker.busy_device_ids() → payload
|
||||
├─ claim loop: 不变(仍会 claim cloud 任务)
|
||||
└─ console_server: 现在同时服务 / * (HTML) 和 /mcp (Streamable HTTP)
|
||||
```
|
||||
|
||||
### 6.2 Hermes 第一次 MCP 调用(lazy acquire)
|
||||
|
||||
```
|
||||
1. Hermes → POST /mcp (JSON-RPC tools/call, name=take_screenshot, args={device_id: "phone-1"})
|
||||
2. BearerAuthMiddleware: verify token → pass
|
||||
3. FastMCP route → 包装层 decorator:
|
||||
a. session_id = ctx.session_id
|
||||
b. busy check:
|
||||
- cloud 占用? → status_tracker.snapshot()["current_assignment"]["device_id"] == "phone-1"? NO
|
||||
- mcp 占用? → mcp_busy_tracker.busy_device_ids() 包含 "phone-1"? NO
|
||||
c. acquire: mcp_busy_tracker.acquire("phone-1", session_id) → True
|
||||
d. tool_handlers(manager=...)["take_screenshot"](device_id="phone-1") → screenshot
|
||||
e. renew: mcp_busy_tracker.renew("phone-1", session_id)
|
||||
f. return screenshot_base64
|
||||
4. 下一次心跳(≤30s): payload.mcp_busy_device_ids = ["phone-1"]
|
||||
5. Cloud scheduler 收到 → 把 phone-1 视为非 idle → 不派任务给它
|
||||
```
|
||||
|
||||
### 6.3 同 session 第二次调用(锁已持有)
|
||||
|
||||
```
|
||||
1. Hermes → POST /mcp (tap, device_id: "phone-1")
|
||||
2-3a. 同上
|
||||
3b. mcp_busy_tracker.busy_device_ids() 包含 "phone-1",但持有者就是当前 session_id → 通过
|
||||
3c. acquire 已持有,no-op(或 assert)
|
||||
3d-f. 同上
|
||||
```
|
||||
|
||||
### 6.4 Cloud 任务抢占的预防
|
||||
|
||||
```
|
||||
场景:Hermes 正在用 phone-1,cloud 这时派任务给 phone-1
|
||||
|
||||
正常路径:
|
||||
1. Cloud scheduler 选设备:phone-1 在最近心跳的 mcp_busy_device_ids 内 → 视为非 idle → 跳过
|
||||
2. Cloud 选 phone-2 或等其他设备 idle
|
||||
|
||||
窗口期兜底(心跳 30s 内 cloud 基于旧心跳派任务):
|
||||
1. AssignmentProcessor.process(assignment) → AssignmentExecutor.execute(assignment)
|
||||
2. execute() 启动前检查:
|
||||
if assignment.device_id in mcp_busy_tracker.busy_device_ids():
|
||||
return AssignmentExecutionResult(
|
||||
status="failed",
|
||||
failure_reason="device held by active MCP session"
|
||||
)
|
||||
3. client.report_result(...) 上报 fail → cloud attempt+1 或派给别的 host
|
||||
```
|
||||
|
||||
### 6.5 Session 结束 / TTL 过期
|
||||
|
||||
```
|
||||
正常:Hermes 主动断开 → mcp SDK 1.28.1 没有 per-session shutdown hook
|
||||
→ 该 session 的 lease 进入 TTL 倒计时
|
||||
→ 20s TTL 到期 → 下次 busy_device_ids() 或 snapshot() 调用时 lazy sweep
|
||||
→ lease 清理 → 下次心跳 payload 不再包含 → cloud 重新视为 idle
|
||||
|
||||
异常:Hermes 崩溃 / 网络断 → 同上,无 shutdown callback
|
||||
→ 20s TTL 到期 → 下次 busy_device_ids() 或 snapshot() 调用时 lazy sweep
|
||||
→ lease 清理 → 下次心跳 payload 不再包含 → cloud 重新视为 idle
|
||||
```
|
||||
|
||||
**Implementation note (2026-07-21 fix wave):** mcp SDK 1.28.1 exposes
|
||||
only a server-level `lifespan` hook; `ServerSession.__aexit__` and
|
||||
`StreamableHTTPSessionManager` do not surface a per-session
|
||||
shutdown callback. The spec originally described a release-on-clean-
|
||||
disconnect path that the SDK cannot deliver today. The fallback is
|
||||
the 20-second TTL sweep — short enough that a normal heartbeat
|
||||
interval (30s) catches the recovery before the cloud scheduler
|
||||
notices, long enough that an actively-busy session does not lose its
|
||||
lease during normal operator pauses. Explicit release on session end
|
||||
remains a future enhancement if/when the SDK exposes the hook.
|
||||
|
||||
## 7. Error Handling Matrix
|
||||
|
||||
| # | 触发条件 | 返回语义 | 备注 |
|
||||
|---|---|---|---|
|
||||
| 1 | `Authorization` 缺失/不匹配 | HTTP 401 + `WWW-Authenticate: Bearer` + JSON `{"error":"invalid token"}` | 不写失败日志;首次成功鉴权写 INFO |
|
||||
| 2 | Cloud 占用目标设备 | `CallToolResult(isError=true, content=[TextContent("device phone-1 is busy (held by cloud assignment)")])` | DEBUG 日志 |
|
||||
| 3 | 另一 MCP session 占用 | `CallToolResult(isError=true, content=[TextContent("device phone-1 is busy (held by mcp_session:<8-char-prefix>)")])` | DEBUG 日志 |
|
||||
| 4 | `device_id` 不存在 | `CallToolResult(isError=false)` + JSON `{"ok": false, "error": "device not found: phone-1"}` | 复用 `call_with_semantic_errors`;语义错误不抛 |
|
||||
| 5 | 工具底层异常 | `CallToolResult(isError=true, content=[TextContent("Error executing tool <name>: <orig-msg>")])` | WARNING + exc_info |
|
||||
| 6 | Cloud assignment 启动前命中 MCP 占用 | `AssignmentExecutionResult(status="failed", failure_reason="device held by active MCP session")` | INFO 一次 |
|
||||
| 7 | Token 文件损坏 JSON | host-agent 启动失败,stderr 提示 | 不静默重新生成 |
|
||||
| 8 | Token 文件不可写 | host-agent 启动失败 | 同上 |
|
||||
| 9 | MCP session 异常断开 | lease 进入 TTL 倒计时 | 20s 后 lazy sweep |
|
||||
| 10 | TTL 过期瞬间 Hermes 重连 | renew 容忍边界:session_id 匹配 → 重新 acquire 而非报错 | 无感 |
|
||||
| 11 | 同 session 并发不同设备 | 各自独立 acquire | per-device 设计 |
|
||||
| 12 | 同 session 并发同一设备 | 第一个 acquire;第二个 renew(同 session_id) | 并发 safe |
|
||||
| 13 | FastMCP 提取不到 session_id | `CallToolResult(isError=true, content=[TextContent("cannot determine MCP session")])` | ERROR + exc_info |
|
||||
| 14 | host-agent 关停时有 active MCP session | lease 随进程退出消失 | 不需显式清理 |
|
||||
|
||||
### 错误返回格式
|
||||
|
||||
```json
|
||||
{
|
||||
"jsonrpc": "2.0",
|
||||
"id": "<request-id>",
|
||||
"result": {
|
||||
"content": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "device phone-1 is busy (held by cloud assignment)"
|
||||
}
|
||||
],
|
||||
"isError": true
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
**Implementation note (2026-07-21 fix wave):** mcp SDK 1.28.1's
|
||||
`Tool.run` wraps every non-`UrlElicitationRequiredError` exception
|
||||
(including `McpError` with a typed `ErrorData`) into `ToolError`.
|
||||
The lowlevel `call_tool` handler then serializes any exception as
|
||||
`CallToolResult(isError=true, content=[TextContent(message)])` via
|
||||
`_make_error_result`. There is no public path that surfaces JSON-RPC
|
||||
`-32000` with a structured `data.busy_owner` field from a tool call
|
||||
site — the SDK's wire contract for tool errors is the `isError=true`
|
||||
flag plus text content. The busy-owner value lives in the text
|
||||
content (truncated session_id for `mcp_session:` collisions, full
|
||||
string for `cloud_assignment`). Unknown-device and other semantic
|
||||
errors are returned as normal `ok=False` payloads inside a successful
|
||||
`CallToolResult(isError=false)` (see `api/mcp.py::call_with_semantic_errors`).
|
||||
|
||||
## 8. Testing Strategy
|
||||
|
||||
### 8.1 单元测试
|
||||
|
||||
- **`McpBusyTracker`**(~10 用例):acquire/renew/release/busy_device_ids/snapshot/wait_until_usable 的成功/冲突/并发/TTL sweep
|
||||
- **`McpTokenStore`**(~6 用例):首次创建/二次复用/并发 atomic/verify timing-safe/损坏 JSON 抛错/不可写抛错
|
||||
- **`BearerAuthMiddleware`**(~5 用例):无 header/错误 token/正确 token/非 Bearer scheme/大小写
|
||||
- **包装层 decorator**(~10 用例,**重点**):
|
||||
- 显式 `manager=` 传入 → 工具调用成功(**回归 D12 地雷**)
|
||||
- `manager=None` → 包装层断言失败(**新约束**)
|
||||
- cloud 占用 → `-32000` busy(mock AgentStatusTracker)
|
||||
- 另一 MCP session 占用 → `-32000` busy
|
||||
- 同 session 已持有 → renew 后通过
|
||||
- acquire/renew 调用顺序
|
||||
- 底层 `DeviceNotFoundError` → `-32602`
|
||||
- 底层 `DriverError` → `-32000`
|
||||
- `list_devices` 不 busy check 但走 display status
|
||||
- `device_status` 同上
|
||||
|
||||
### 8.2 集成测试
|
||||
|
||||
- **Mount wiring**(~4 用例):`/mcp` mount 存在、不继承 cookie session、`/api/status` 加字段、dashboard HTML 含状态行
|
||||
- **心跳 payload**(~3 用例):空 / 含 phone-1 / TTL 过期后空
|
||||
- **AssignmentExecutor fail-fast**(~2 用例):命中 MCP 占用立即 fail、不命中正常执行
|
||||
- **Cloud 侧**(~4 用例):Heartbeat model 向后兼容空 list、含 device、`_matches()` 跳过 mcp_busy、旧 host-agent 不发字段时行为不变
|
||||
|
||||
### 8.3 端到端(1 个,可选)
|
||||
|
||||
- MCP 客户端模拟(`mcp` SDK client + ASGI TestClient)→ host-agent 进程内 → FakeDriver
|
||||
- 覆盖:鉴权 → list_devices → screenshot → tap → 中途 cloud assignment fail-fast → session 断开 release
|
||||
- 如果 SDK client + TestClient 组合有坑,降级为直接调用 FastMCP tool registry
|
||||
|
||||
### 8.4 已有测试不破坏
|
||||
|
||||
- `test_app.py` 9 处 `create_application()` 调用:新组件在 `create_application()` 内部构造,外部 API 不变
|
||||
- `test_runtime_owned_packages_do_not_import_host_or_cloud_concerns`:所有新代码在 `host_agent/`,不动 `runtime/` / `api/` 包边界(`api/mcp.py` 签名收紧不算破坏)
|
||||
- 根非集成测试套件:跑一遍无回归
|
||||
|
||||
### 8.5 人工/真机(不在自动化覆盖)
|
||||
|
||||
- 真实 Hermes Agent CLI 接入(写 SOUL.md、跑真实对话)
|
||||
- 真机 iPhone/Android MCP 工具调用
|
||||
|
||||
## 9. Scope
|
||||
|
||||
### 9.1 In Scope
|
||||
|
||||
见第 5 节 Components 表。共:
|
||||
- 4 个新模块(`mcp_lock.py` / `mcp_token.py` / `web/mcp_auth.py` / `web/mcp.py`)
|
||||
- 8 个改动的现有模块
|
||||
- 2 个文档(`docs/MCP_INTEGRATION.md` 新增 + `docs/MACOS_IPHONE_SETUP.md` 加小节)
|
||||
- 单元 + 集成测试全覆盖;e2e 视成本而定
|
||||
|
||||
### 9.2 Non-Goals
|
||||
|
||||
- `wait_until_usable` 的调用方接入(方法实现 + 单测,但 AssignmentExecutor 和 MCP 工具包装层都不调用)
|
||||
- 暴露 `wait_until_device_usable` 成独立 MCP 工具
|
||||
- `device-host-agent mcp-rotate-token` CLI(MVP 轮换走"删文件 + 重启")
|
||||
- MCP 调用进 `ConsoleHistoryStore`(MVP 不记录调用历史)
|
||||
- Skill catalog 工具暴露给 Hermes(保留参数化挂载点)
|
||||
- Token rate limiting / 失败计数
|
||||
- 多 Hermes 实例协调(技术允许但不专门测试)
|
||||
- 远程访问(非 loopback)支持
|
||||
- Hermes 侧 SOUL.md / 配置自动化
|
||||
- MCP over WebSocket
|
||||
- Cloud scheduler 改造为支持"MCP 优先级"/队列
|
||||
|
||||
## 10. Open Questions
|
||||
|
||||
- **Q1(已答)**:`tool_handlers(manager)` 是否改为必传?→ **是**。Grep 全仓库确认所有调用方(`api/mcp.py:110`、`tests/test_mcp.py:13,56`、`tests/test_skill_catalog_e2e.py:287`)都已显式传 `manager=manager`,改动面为 0。
|
||||
- **Q2**:`mcp-token` CLI 是否需要鉴权?→ MVP 不鉴权,假定能访问宿主机的操作者可信(与 `setup` 子命令同款)。后续可加 `--password` 校验 local_account。
|
||||
- **Q3**:MCP 工具调用是否需要 wall-clock 超时?→ MVP 不加,依赖底层超时;如有"Hermes 调用挂死"报告再加。
|
||||
- **Q4**:心跳扩展是 Alembic migration 还是仅 schema 字段?→ MVP 仅 transient 字段(heartbeat 接收 → scheduler 用 → 丢弃),无需 migration。
|
||||
- **Q5(已答,2026-07-21 fix wave)**:TTL 从 60s 降到 20s,因为 mcp SDK 1.28.1 没有 per-session shutdown callback(见 §6.5)。20s 仍能容许正常 operator 暂停,但能在一次 30s 心跳窗口内回收崩盘 session 的锁;新 `test_default_ttl_is_20_seconds` 和 `test_default_ttl_recovers_dead_session_within_one_window` 锁定该值。如有"Hermes 长操作横跨 20s 静默"报告再调高。
|
||||
|
||||
## 11. Risks
|
||||
|
||||
- **R1**:Cloud 心跳窗口期(30s)冲突。缓解:AssignmentExecutor 启动前 fail-fast。残留:cloud 可能基于过期心跳派任务、host-agent fail、cloud 重试——浪费 attempt 配额。**接受**。
|
||||
- **R2**:Hermes 长时间占用设备导致 cloud 任务反复 fail。MVP 无自动缓解;用户手动管控;未来启用 `wait_until_usable`。**接受**。
|
||||
- **R3**:Hermes 崩溃后 ≤20s 设备不可用(TTL 兜底,但窗口存在)。窗口仍 ≤30s 心跳间隔,cloud 侧下次心跳能学到。**接受,记入 Q5**。
|
||||
- **R4(已消除)**:`tool_handlers` 签名改动影响面。Grep 确认 0 调用方依赖 None 默认。
|
||||
- **R5**:FastMCP `session_id` 提取依赖 `mcp` SDK 内部 API。缓解:e2e 测试覆盖;SDK 升级 CI 能及时暴露。
|
||||
- **R6**:`mcp` SDK 需作为 device-host-agent 直接依赖(目前通过 Runtime 传递)。需加进 `apps/device-host-agent/pyproject.toml`(与 `filelock` 直接化先例一致)。
|
||||
- **R7**:Token 文件首次生成的竞态。缓解:InstanceLock 前置保护;`tempfile + os.replace` 原子重命名。
|
||||
|
||||
## 12. Coordination with Existing Changes
|
||||
|
||||
- **`host-agent-single-instance-lock`**:本次依赖 InstanceLock,token 文件生成在 InstanceLock acquire 之后,安全。无需修改。
|
||||
- **`host-agent-dependency-supervisor`**:本次不引入新的外部进程依赖(FastMCP 是库)。不冲突。
|
||||
- **`task-execution-progress-visibility`**:本次不改 `TaskMetadataStore` / `Timeline`。不冲突。
|
||||
- **`host-agent-local-console`**:本次复用其 FastAPI + uvicorn 设施,新增一个 mount 点。Console 现有 cookie session 鉴权**不**继承到 `/mcp`(鉴权走独立 bearer middleware)。
|
||||
|
||||
## 13. Hermes 侧配置示例(参考,非 host-agent 代码范围)
|
||||
|
||||
`~/.hermes/config.yaml`:
|
||||
```yaml
|
||||
mcp_servers:
|
||||
apex_device:
|
||||
url: "http://127.0.0.1:8765/mcp"
|
||||
headers:
|
||||
Authorization: "Bearer <token-from-host_mcp_token.json>"
|
||||
```
|
||||
|
||||
首次启动 host-agent 后,从 `tasks/host_mcp_token.json`(或 `HOST_AGENT_IDENTITY_PATH` 同目录)读 token,或跑 `device-host-agent mcp-token` 打印。
|
||||
|
||||
Hermes 的 `SOUL.md`(profile 级别)建议补充:
|
||||
- 工具调用前先 `list_devices` 看可用设备
|
||||
- 设备 `busy` 时等待或换设备
|
||||
- iOS 设备坐标是 points(非 pixels),Android 是 pixels
|
||||
- OCR/UI 树 bounds 与 screenshot 像素已对齐(perception 层已处理)
|
||||
|
||||
## 14. Implementation Order(建议,writing-plans 阶段细化)
|
||||
|
||||
1. `api/mcp.py::tool_handlers` 签名收紧(D12)+ 同步更新调用方 docstring/类型
|
||||
2. `McpTokenStore` + 单测
|
||||
3. `McpBusyTracker`(含 `wait_until_usable`)+ 单测
|
||||
4. `BearerAuthMiddleware` + 单测
|
||||
5. `build_mcp_server` 包装层 + 单测(覆盖 D12 回归)
|
||||
6. `create_console_app` mount wiring + 集成测试
|
||||
7. `create_application` 装配
|
||||
8. 心跳 payload 扩展 + cloud scheduler 改动 + 集成测试
|
||||
9. `AssignmentExecutor` fail-fast 检查 + 测试
|
||||
10. CLI `mcp-token` 子命令
|
||||
11. Dashboard 状态行
|
||||
12. 文档(`MCP_INTEGRATION.md` + `MACOS_IPHONE_SETUP.md`)
|
||||
13. e2e 测试(可选)
|
||||
14. 全量非集成测试回归 + ruff + compileall + openspec strict validation
|
||||
@@ -81,6 +81,13 @@ class AndroidDriver(Driver):
|
||||
except Exception as exc:
|
||||
raise DriverError("screenshot failed") from exc
|
||||
|
||||
def health_check(self) -> None:
|
||||
client = self._require_client()
|
||||
try:
|
||||
client.get_status()
|
||||
except Exception as exc:
|
||||
raise DriverError("health check failed") from exc
|
||||
|
||||
def tap(self, x: float, y: float) -> None:
|
||||
client = self._require_client()
|
||||
try:
|
||||
|
||||
@@ -25,6 +25,15 @@ class Driver(ABC):
|
||||
def screenshot(self) -> bytes:
|
||||
"""Return the current screen as image bytes."""
|
||||
|
||||
def health_check(self) -> None:
|
||||
"""Verify the live session without reading the device screen.
|
||||
|
||||
Drivers with a transport-level status endpoint should override this
|
||||
method. The default is a no-op for legacy drivers that do not expose
|
||||
a separate health check.
|
||||
"""
|
||||
return None
|
||||
|
||||
@abstractmethod
|
||||
def tap(self, x: float, y: float) -> None:
|
||||
"""Tap the screen at the given coordinates."""
|
||||
|
||||
@@ -71,6 +71,13 @@ class WDADriver(Driver):
|
||||
except Exception as exc:
|
||||
raise DriverError("screenshot failed") from exc
|
||||
|
||||
def health_check(self) -> None:
|
||||
client = self._require_client()
|
||||
try:
|
||||
client.get_status()
|
||||
except Exception as exc:
|
||||
raise DriverError("health check failed") from exc
|
||||
|
||||
def tap(self, x: float, y: float) -> None:
|
||||
client = self._require_client()
|
||||
try:
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from sqlalchemy import (
|
||||
Boolean,
|
||||
ForeignKey,
|
||||
Index,
|
||||
Integer,
|
||||
@@ -86,6 +87,7 @@ class PooledDeviceRow(Base):
|
||||
status: Mapped[str] = mapped_column(String, nullable=False)
|
||||
capability_tags_json: Mapped[str] = mapped_column(Text, nullable=False)
|
||||
synced_at: Mapped[str | None] = mapped_column(String, nullable=True)
|
||||
mcp_busy: Mapped[bool] = mapped_column(Boolean, nullable=False, default=False)
|
||||
|
||||
|
||||
class ScheduledTaskRow(Base):
|
||||
|
||||
@@ -6,6 +6,7 @@ import json
|
||||
import logging
|
||||
from collections.abc import Awaitable, Callable
|
||||
from datetime import timedelta
|
||||
from inspect import Parameter, signature
|
||||
from time import monotonic
|
||||
from typing import TYPE_CHECKING
|
||||
from uuid import uuid4
|
||||
@@ -185,6 +186,7 @@ def create_internal_router(
|
||||
address=payload.address,
|
||||
allow_device_takeover=allow_device_takeover,
|
||||
planner_transport=payload.planner_transport,
|
||||
mcp_busy_device_ids=payload.mcp_busy_device_ids,
|
||||
)
|
||||
policy = pool.store.get_host_governance_policy(host_id)
|
||||
policy_revision = policy.revision if policy is not None else 0
|
||||
@@ -504,13 +506,21 @@ def create_internal_router(
|
||||
|
||||
started_at = monotonic()
|
||||
try:
|
||||
decision = client.decide(
|
||||
system_prompt=payload.system_prompt,
|
||||
user_prompt=payload.user_prompt,
|
||||
screenshot=screenshot,
|
||||
tools=tools,
|
||||
timeout=planner_timeout,
|
||||
)
|
||||
decision_kwargs = {
|
||||
"system_prompt": payload.system_prompt,
|
||||
"user_prompt": payload.user_prompt,
|
||||
"screenshot": screenshot,
|
||||
"tools": tools,
|
||||
"timeout": planner_timeout,
|
||||
}
|
||||
parameters = signature(client.decide).parameters.values()
|
||||
if any(
|
||||
parameter.name == "history"
|
||||
or parameter.kind == Parameter.VAR_KEYWORD
|
||||
for parameter in parameters
|
||||
):
|
||||
decision_kwargs["history"] = payload.history
|
||||
decision = client.decide(**decision_kwargs)
|
||||
except ToolCallUnavailable as exc:
|
||||
logger.info(
|
||||
"planner-decision request failed",
|
||||
|
||||
@@ -40,6 +40,7 @@ class HeartbeatRequest(BaseModel):
|
||||
devices: list[DeviceSnapshotModel] = Field(default_factory=list)
|
||||
policy_revision: int = Field(default=0, ge=0)
|
||||
planner_transport: Literal["direct", "cloud"] = "direct"
|
||||
mcp_busy_device_ids: list[str] = Field(default_factory=list)
|
||||
|
||||
|
||||
class HostGovernancePolicyModel(BaseModel):
|
||||
@@ -142,6 +143,7 @@ class PlannerDecisionRequest(BaseModel):
|
||||
host_id: str = Field(min_length=1)
|
||||
system_prompt: str
|
||||
user_prompt: str
|
||||
history: list[dict[str, Any]] = Field(default_factory=list)
|
||||
screenshot_base64: str | None = None
|
||||
tools: list[PlannerToolSpecModel] = Field(default_factory=list)
|
||||
timeout_seconds: float = Field(default=30.0, gt=0, le=120)
|
||||
|
||||
@@ -0,0 +1,28 @@
|
||||
"""Add mcp_busy flag column to pooled_devices."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
|
||||
revision = "0014_pooled_device_mcp_busy"
|
||||
down_revision = "0013_task_cancellation"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.add_column(
|
||||
"pooled_devices",
|
||||
sa.Column(
|
||||
"mcp_busy",
|
||||
sa.Boolean(),
|
||||
nullable=False,
|
||||
server_default=sa.false(),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_column("pooled_devices", "mcp_busy")
|
||||
@@ -47,6 +47,7 @@ class PooledDevice:
|
||||
status: PooledDeviceStatus
|
||||
capability_tags: list[str] = field(default_factory=list)
|
||||
synced_at: datetime | None = None
|
||||
mcp_busy: bool = False
|
||||
|
||||
|
||||
class DevicePool:
|
||||
@@ -64,6 +65,7 @@ class DevicePool:
|
||||
address: str | None = None,
|
||||
planner_transport: Literal["direct", "cloud"] = "direct",
|
||||
allow_device_takeover: bool = False,
|
||||
mcp_busy_device_ids: list[str] | None = None,
|
||||
) -> None:
|
||||
"""Push a host's current device snapshot into the pool.
|
||||
|
||||
@@ -78,7 +80,11 @@ class DevicePool:
|
||||
last_seen_at=now,
|
||||
planner_transport=planner_transport,
|
||||
)
|
||||
devices = [self._to_pooled(device, host_id, now) for device in snapshot]
|
||||
busy_set = set(mcp_busy_device_ids or [])
|
||||
devices = [
|
||||
self._to_pooled(device, host_id, now, mcp_busy=device.id in busy_set)
|
||||
for device in snapshot
|
||||
]
|
||||
if allow_device_takeover:
|
||||
self.store.replace_host_devices(
|
||||
host_id,
|
||||
@@ -119,6 +125,8 @@ class DevicePool:
|
||||
device: Device,
|
||||
host_id: str,
|
||||
synced_at: datetime,
|
||||
*,
|
||||
mcp_busy: bool = False,
|
||||
) -> PooledDevice:
|
||||
raw_status = (
|
||||
device.status if device.status in _HOST_REPORTED_STATUSES else "idle"
|
||||
@@ -131,6 +139,7 @@ class DevicePool:
|
||||
status=raw_status, # type: ignore[arg-type]
|
||||
capability_tags=tags,
|
||||
synced_at=synced_at,
|
||||
mcp_busy=mcp_busy,
|
||||
)
|
||||
|
||||
def _is_stale(self, host: HostRegistration, now: datetime) -> bool:
|
||||
|
||||
@@ -204,6 +204,8 @@ class TaskScheduler:
|
||||
|
||||
|
||||
def _matches(device: "PooledDevice", constraints: TaskConstraints) -> bool:
|
||||
if device.mcp_busy:
|
||||
return False
|
||||
if constraints.target_host_id and device.host_id != constraints.target_host_id:
|
||||
return False
|
||||
if (
|
||||
|
||||
@@ -9,7 +9,7 @@ from alembic.runtime.migration import MigrationContext
|
||||
from cloud.database import create_database_engine, normalize_database_url
|
||||
|
||||
|
||||
HEAD_REVISION = "0013_task_cancellation"
|
||||
HEAD_REVISION = "0014_pooled_device_mcp_busy"
|
||||
|
||||
|
||||
class SchemaVersionError(RuntimeError):
|
||||
|
||||
@@ -342,6 +342,7 @@ class SQLAlchemyCloudRepository:
|
||||
ensure_ascii=False,
|
||||
),
|
||||
synced_at=_iso(device.synced_at) if device.synced_at else None,
|
||||
mcp_busy=getattr(device, "mcp_busy", False),
|
||||
)
|
||||
for device in devices
|
||||
]
|
||||
@@ -2232,6 +2233,7 @@ def _device_from_row(row: PooledDeviceRow) -> Any:
|
||||
status=row.status,
|
||||
capability_tags=tags,
|
||||
synced_at=_parse_dt(row.synced_at),
|
||||
mcp_busy=bool(getattr(row, "mcp_busy", False)),
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,20 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from cloud.internal_api.models import HeartbeatRequest
|
||||
|
||||
|
||||
def test_heartbeat_request_defaults_mcp_busy_device_ids_to_empty() -> None:
|
||||
req = HeartbeatRequest(host_id="h1")
|
||||
assert req.mcp_busy_device_ids == []
|
||||
|
||||
|
||||
def test_heartbeat_request_accepts_mcp_busy_device_ids() -> None:
|
||||
req = HeartbeatRequest(host_id="h1", mcp_busy_device_ids=["phone-1"])
|
||||
assert req.mcp_busy_device_ids == ["phone-1"]
|
||||
|
||||
|
||||
def test_heartbeat_request_omitting_field_is_backward_compatible() -> None:
|
||||
"""Old host-agents that don't send the field must still validate."""
|
||||
raw = {"host_id": "h1", "devices": []}
|
||||
req = HeartbeatRequest.model_validate(raw)
|
||||
assert req.mcp_busy_device_ids == []
|
||||
@@ -0,0 +1,68 @@
|
||||
"""Tests for the cloud DevicePool, focused on the mcp_busy flag plumbing."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
from cloud.config import CloudConfig
|
||||
from cloud.pool import DevicePool
|
||||
from cloud.store import CloudStore
|
||||
from core.models import Device
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def pool() -> DevicePool:
|
||||
"""Build a fresh DevicePool backed by a temporary SQLite file."""
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
store = CloudStore(Path(tmp) / "cloud.sqlite3")
|
||||
try:
|
||||
yield DevicePool(store=store, config=CloudConfig())
|
||||
finally:
|
||||
store.close()
|
||||
|
||||
|
||||
def _device(device_id: str, *, status: str = "idle") -> Device:
|
||||
return Device(
|
||||
id=device_id,
|
||||
driver_type="wda",
|
||||
status=status, # type: ignore[arg-type]
|
||||
capability_tags=[],
|
||||
)
|
||||
|
||||
|
||||
def test_sync_host_devices_marks_mcp_busy_devices(pool: DevicePool) -> None:
|
||||
"""When a host reports device-1 as MCP-busy, the pool PooledDevice for
|
||||
device-1 has mcp_busy=True."""
|
||||
pool.sync_host_devices(
|
||||
"host-1",
|
||||
[_device("device-1", status="idle")],
|
||||
mcp_busy_device_ids=["device-1"],
|
||||
)
|
||||
devices = pool.list_devices()
|
||||
busy = [d for d in devices if d.device_id == "device-1"]
|
||||
assert len(busy) == 1
|
||||
assert busy[0].mcp_busy is True
|
||||
|
||||
|
||||
def test_sync_host_devices_default_mcp_busy_is_false(pool: DevicePool) -> None:
|
||||
pool.sync_host_devices("host-1", [_device("device-1", status="idle")])
|
||||
devices = pool.list_devices()
|
||||
assert devices[0].mcp_busy is False
|
||||
|
||||
|
||||
def test_sync_host_devices_clears_mcp_busy_on_next_sync(
|
||||
pool: DevicePool,
|
||||
) -> None:
|
||||
"""MCP releases device -> next heartbeat without device in
|
||||
mcp_busy_device_ids -> pool reflects mcp_busy=False."""
|
||||
pool.sync_host_devices(
|
||||
"host-1",
|
||||
[_device("device-1", status="idle")],
|
||||
mcp_busy_device_ids=["device-1"],
|
||||
)
|
||||
pool.sync_host_devices("host-1", [_device("device-1", status="idle")])
|
||||
devices = pool.list_devices()
|
||||
assert devices[0].mcp_busy is False
|
||||
@@ -0,0 +1,53 @@
|
||||
"""Tests for the cloud TaskScheduler, focused on skipping MCP-busy devices."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
from cloud.config import CloudConfig
|
||||
from cloud.pool import DevicePool
|
||||
from cloud.scheduler import TaskConstraints, TaskScheduler
|
||||
from cloud.store import CloudStore
|
||||
from core.models import Device
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def pool() -> DevicePool:
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
store = CloudStore(Path(tmp) / "cloud.sqlite3")
|
||||
try:
|
||||
yield DevicePool(store=store, config=CloudConfig())
|
||||
finally:
|
||||
store.close()
|
||||
|
||||
|
||||
def _device(device_id: str, *, status: str = "idle") -> Device:
|
||||
return Device(
|
||||
id=device_id,
|
||||
driver_type="wda",
|
||||
status=status, # type: ignore[arg-type]
|
||||
capability_tags=[],
|
||||
)
|
||||
|
||||
|
||||
def test_mcp_busy_device_is_skipped_by_scheduler(pool: DevicePool) -> None:
|
||||
"""A device with mcp_busy=True is not selected for assignment.
|
||||
|
||||
Two devices exist: dev-busy (idle status, mcp_busy=True) and dev-idle
|
||||
(idle status, mcp_busy=False). One task is submitted with no
|
||||
constraints, so both are candidates before the mcp_busy filter.
|
||||
The scheduler must pick dev-idle.
|
||||
"""
|
||||
pool.sync_host_devices(
|
||||
"host-1",
|
||||
[_device("dev-busy"), _device("dev-idle")],
|
||||
mcp_busy_device_ids=["dev-busy"],
|
||||
)
|
||||
scheduler = TaskScheduler(pool=pool, store=pool.store, config=CloudConfig())
|
||||
scheduler.submit(goal="test", constraints=TaskConstraints())
|
||||
assignments = scheduler.assign()
|
||||
assert len(assignments) == 1
|
||||
assert assignments[0].device_id == "dev-idle"
|
||||
+87
-24
@@ -1,5 +1,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from inspect import Parameter, signature
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from core.errors import TaskFailedError
|
||||
@@ -23,9 +25,11 @@ class AIPlanner(Planner):
|
||||
*,
|
||||
client: ToolCallingClient | None = None,
|
||||
config: PlannerConfig | None = None,
|
||||
event_logger: Callable[[dict[str, Any]], None] | None = None,
|
||||
) -> None:
|
||||
self.config = config or load_config()
|
||||
self.client = client or build_client(self.config)
|
||||
self.event_logger = event_logger
|
||||
|
||||
def plan(
|
||||
self,
|
||||
@@ -36,25 +40,58 @@ class AIPlanner(Planner):
|
||||
world: "WorldState | None" = None,
|
||||
screenshot: bytes | None = None,
|
||||
) -> list[PlannedStep]:
|
||||
_sync_tool_results(context)
|
||||
scene_json = _without_ocr(scene.to_dict()) if self.config.multimodal else scene.to_dict()
|
||||
user_prompt = planner_user_prompt(
|
||||
goal=goal,
|
||||
scene_json=scene.to_dict(),
|
||||
history_summary=_history_summary(world),
|
||||
scene_json=scene_json,
|
||||
device_platform=context.device_platform,
|
||||
)
|
||||
decision = self.client.decide(
|
||||
system_prompt=PLANNER_SYSTEM_PROMPT,
|
||||
user_prompt=user_prompt,
|
||||
screenshot=screenshot,
|
||||
tools=ALL_TOOL_SPECS,
|
||||
timeout=self.config.timeout,
|
||||
)
|
||||
if context.step_results and not context.step_results[-1].success:
|
||||
user_prompt += (
|
||||
"\n\nThe previous tool call failed. Diagnose the failure and choose a "
|
||||
"corrected action or finish the task if it cannot proceed.\n"
|
||||
f"Previous failure: {context.step_results[-1].error or 'unknown error'}"
|
||||
)
|
||||
self._log({"type": "llm_request", "task_id": context.task_id, "goal": goal,
|
||||
"system_prompt": PLANNER_SYSTEM_PROMPT, "user_prompt": user_prompt,
|
||||
"has_screenshot": screenshot is not None})
|
||||
try:
|
||||
kwargs = {
|
||||
"system_prompt": PLANNER_SYSTEM_PROMPT,
|
||||
"user_prompt": user_prompt,
|
||||
"screenshot": screenshot,
|
||||
"tools": ALL_TOOL_SPECS,
|
||||
"timeout": self.config.timeout,
|
||||
}
|
||||
if _accepts_history(self.client.decide):
|
||||
kwargs["history"] = context.planner_history[
|
||||
-self.config.history_max_turns :
|
||||
]
|
||||
decision = self.client.decide(**kwargs)
|
||||
except Exception as exc:
|
||||
self._log({"type": "agent_error", "task_id": context.task_id, "error": str(exc)})
|
||||
raise
|
||||
self._log({"type": "llm_response", "task_id": context.task_id,
|
||||
"content": decision.text_output, "thinking": decision.thinking,
|
||||
"tool_name": decision.tool_name, "arguments": decision.arguments,
|
||||
"purpose": decision.purpose, "expected_outcome": decision.expected_outcome})
|
||||
|
||||
if decision.tool_name == FINISH_TASK_TOOL:
|
||||
if decision.arguments.get("success"):
|
||||
return []
|
||||
raise TaskFailedError(decision.arguments.get("reason") or "task failed")
|
||||
|
||||
context.planner_history.append(
|
||||
{
|
||||
"user_prompt": user_prompt,
|
||||
"tool_name": decision.tool_name,
|
||||
"arguments": _conversation_arguments(decision),
|
||||
"rationale": decision.text_output,
|
||||
"tool_result": None,
|
||||
}
|
||||
)
|
||||
|
||||
return [
|
||||
PlannedStep(
|
||||
action=decision.tool_name,
|
||||
@@ -78,19 +115,45 @@ class AIPlanner(Planner):
|
||||
# (mapped to an empty plan above), never via this hook.
|
||||
return False
|
||||
|
||||
def _log(self, event: dict[str, Any]) -> None:
|
||||
if self.event_logger is not None:
|
||||
try:
|
||||
self.event_logger(event)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def _history_summary(world: "WorldState | None") -> list[dict[str, Any]]:
|
||||
if world is None:
|
||||
return []
|
||||
return [
|
||||
{
|
||||
"page": event.page,
|
||||
"action": event.action,
|
||||
"arguments": dict(event.arguments),
|
||||
"rationale": event.rationale,
|
||||
"purpose": event.purpose,
|
||||
"expected_outcome": event.expected_outcome,
|
||||
"success": event.success,
|
||||
}
|
||||
for event in world.history
|
||||
]
|
||||
|
||||
def _without_ocr(scene_json: dict[str, Any]) -> dict[str, Any]:
|
||||
cleaned = dict(scene_json)
|
||||
elements = cleaned.get("elements")
|
||||
if isinstance(elements, list):
|
||||
cleaned["elements"] = [
|
||||
{key: value for key, value in element.items() if key not in {"source", "confidence", "foreground_color", "background_color"}}
|
||||
for element in elements
|
||||
if isinstance(element, dict) and element.get("source") != "ocr"
|
||||
]
|
||||
cleaned.pop("ocr_elements", None)
|
||||
return cleaned
|
||||
|
||||
|
||||
def _sync_tool_results(context: TaskContext) -> None:
|
||||
for turn, result in zip(context.planner_history, context.step_results, strict=False):
|
||||
if turn.get("tool_result") is None:
|
||||
turn["tool_result"] = result.to_dict()
|
||||
|
||||
|
||||
def _conversation_arguments(decision: Any) -> dict[str, Any]:
|
||||
arguments = dict(decision.arguments)
|
||||
if decision.purpose is not None:
|
||||
arguments["purpose"] = decision.purpose
|
||||
if decision.expected_outcome is not None:
|
||||
arguments["expected_outcome"] = decision.expected_outcome
|
||||
return arguments
|
||||
|
||||
|
||||
def _accepts_history(method: Any) -> bool:
|
||||
parameters = signature(method).parameters.values()
|
||||
return any(
|
||||
parameter.name == "history" or parameter.kind == Parameter.VAR_KEYWORD
|
||||
for parameter in parameters
|
||||
)
|
||||
|
||||
+2
-1
@@ -1,7 +1,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import TYPE_CHECKING
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from core.models import Scene
|
||||
|
||||
@@ -18,6 +18,7 @@ class TaskContext:
|
||||
scenes: list[Scene] = field(default_factory=list)
|
||||
step_results: list["StepResult"] = field(default_factory=list)
|
||||
world: "WorldState | None" = None
|
||||
planner_history: list[dict[str, Any]] = field(default_factory=list)
|
||||
|
||||
def add_scene(self, scene: Scene) -> None:
|
||||
self.scenes.append(scene)
|
||||
|
||||
+5
-1
@@ -124,7 +124,7 @@ def default_tool_registry(
|
||||
from tools.tap import tap
|
||||
from tools.ui_tree import get_ui_tree
|
||||
|
||||
return {
|
||||
registry = {
|
||||
"take_screenshot": _bind_manager(take_screenshot, manager),
|
||||
"screenshot": _bind_manager(take_screenshot, manager),
|
||||
"tap": _bind_manager(tap, manager),
|
||||
@@ -143,6 +143,10 @@ def default_tool_registry(
|
||||
"find_icon": find_icon,
|
||||
"find_icon_on_screen": _bind_manager(find_icon_on_screen, manager),
|
||||
}
|
||||
if manager is not None:
|
||||
registry["list_devices"] = lambda: [device.to_dict() for device in manager.list_devices()]
|
||||
registry["device_status"] = lambda device_id: {"device_id": device_id, "status": manager.status(device_id)}
|
||||
return registry
|
||||
|
||||
|
||||
def _bind_manager(func: ToolCallable, manager: DeviceManager | None) -> ToolCallable:
|
||||
|
||||
@@ -8,14 +8,20 @@ DEFAULT_PROVIDER = "anthropic"
|
||||
DEFAULT_MODEL_BY_PROVIDER = {
|
||||
"anthropic": "claude-sonnet-5",
|
||||
"openai": "gpt-5.6",
|
||||
"openai_compatible": "local-model",
|
||||
}
|
||||
DEFAULT_TIMEOUT_SECONDS = 30.0
|
||||
DEFAULT_HISTORY_MAX_TURNS = 20
|
||||
|
||||
ENABLED_ENV = "AI_PLANNER_ENABLED"
|
||||
PROVIDER_ENV = "AI_PLANNER_PROVIDER"
|
||||
MODEL_ENV = "AI_PLANNER_MODEL"
|
||||
TIMEOUT_ENV = "AI_PLANNER_TIMEOUT_SECONDS"
|
||||
THINKING_BUDGET_ENV = "AI_PLANNER_THINKING_BUDGET_TOKENS"
|
||||
API_KEY_ENV = "AI_PLANNER_API_KEY"
|
||||
BASE_URL_ENV = "AI_PLANNER_BASE_URL"
|
||||
MULTIMODAL_ENV = "AI_PLANNER_MULTIMODAL"
|
||||
HISTORY_MAX_TURNS_ENV = "AI_PLANNER_HISTORY_MAX_TURNS"
|
||||
|
||||
SUPPORTED_PROVIDERS = frozenset(DEFAULT_MODEL_BY_PROVIDER)
|
||||
|
||||
@@ -27,6 +33,10 @@ class PlannerConfig:
|
||||
model: str = ""
|
||||
timeout: float = DEFAULT_TIMEOUT_SECONDS
|
||||
thinking_budget_tokens: int | None = None
|
||||
api_key: str | None = None
|
||||
base_url: str | None = None
|
||||
multimodal: bool = False
|
||||
history_max_turns: int = DEFAULT_HISTORY_MAX_TURNS
|
||||
|
||||
def resolved_model(self) -> str:
|
||||
return self.model or DEFAULT_MODEL_BY_PROVIDER[self.provider]
|
||||
@@ -40,6 +50,13 @@ def load_config(env: Mapping[str, str] | None = None) -> PlannerConfig:
|
||||
model=values.get(MODEL_ENV) or "",
|
||||
timeout=_parse_timeout(values.get(TIMEOUT_ENV)),
|
||||
thinking_budget_tokens=_parse_thinking_budget(values.get(THINKING_BUDGET_ENV)),
|
||||
api_key=values.get(API_KEY_ENV) or _provider_key(values),
|
||||
base_url=values.get(BASE_URL_ENV) or None,
|
||||
multimodal=_parse_bool(values.get(MULTIMODAL_ENV), default=False),
|
||||
history_max_turns=_parse_positive_int(
|
||||
values.get(HISTORY_MAX_TURNS_ENV),
|
||||
default=DEFAULT_HISTORY_MAX_TURNS,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@@ -53,6 +70,8 @@ def _parse_provider(value: str | None) -> str:
|
||||
if value is None:
|
||||
return DEFAULT_PROVIDER
|
||||
provider = value.strip().lower()
|
||||
if provider in {"openai-compatible", "openai_compatible", "local"}:
|
||||
return "openai_compatible"
|
||||
return provider if provider in SUPPORTED_PROVIDERS else DEFAULT_PROVIDER
|
||||
|
||||
|
||||
@@ -74,3 +93,22 @@ def _parse_thinking_budget(value: str | None) -> int | None:
|
||||
except ValueError:
|
||||
return None
|
||||
return budget if budget > 0 else None
|
||||
|
||||
|
||||
def _parse_positive_int(value: str | None, *, default: int) -> int:
|
||||
if value is None:
|
||||
return default
|
||||
try:
|
||||
parsed = int(value)
|
||||
except ValueError:
|
||||
return default
|
||||
return parsed if parsed > 0 else default
|
||||
|
||||
|
||||
def _provider_key(values: Mapping[str, str]) -> str | None:
|
||||
provider = (values.get(PROVIDER_ENV) or DEFAULT_PROVIDER).strip().lower()
|
||||
if provider == "openai":
|
||||
return values.get("OPENAI_API_KEY") or None
|
||||
if provider == "anthropic":
|
||||
return values.get("ANTHROPIC_API_KEY") or None
|
||||
return None
|
||||
|
||||
@@ -66,7 +66,6 @@ def planner_user_prompt(
|
||||
*,
|
||||
goal: str,
|
||||
scene_json: dict[str, Any],
|
||||
history_summary: list[dict[str, Any]],
|
||||
device_platform: str | None = None,
|
||||
now: datetime | None = None,
|
||||
) -> str:
|
||||
@@ -83,8 +82,6 @@ def planner_user_prompt(
|
||||
f"{goal}\n\n"
|
||||
"Current Scene (JSON):\n"
|
||||
f"{json.dumps(scene_json, ensure_ascii=False, sort_keys=True)}\n\n"
|
||||
"Recent history, oldest first (JSON):\n"
|
||||
f"{json.dumps(history_summary, ensure_ascii=False, sort_keys=True)}\n\n"
|
||||
"Call exactly one tool for this turn."
|
||||
)
|
||||
|
||||
|
||||
+12
-7
@@ -187,13 +187,18 @@ class TaskRunner:
|
||||
"failed",
|
||||
result.error or "step failed",
|
||||
)
|
||||
self._update_task(
|
||||
task,
|
||||
status="failed",
|
||||
completed=True,
|
||||
failure_reason=result.error or "step failed",
|
||||
)
|
||||
return task
|
||||
# AI planners can use the structured failure feedback to
|
||||
# correct malformed arguments or choose another action.
|
||||
# Deterministic planners retain their fail-fast behavior.
|
||||
if self.planner.__class__.__name__ != "AIPlanner":
|
||||
self._update_task(
|
||||
task,
|
||||
status="failed",
|
||||
completed=True,
|
||||
failure_reason=result.error or "step failed",
|
||||
)
|
||||
return task
|
||||
break
|
||||
|
||||
self._emit_step_progress(
|
||||
len(context.step_results),
|
||||
|
||||
@@ -51,6 +51,7 @@ class ToolCallingClient(Protocol):
|
||||
screenshot: bytes | None,
|
||||
tools: list[ToolSpec],
|
||||
timeout: float,
|
||||
history: list[dict[str, Any]] | None = None,
|
||||
) -> ToolCallDecision: ...
|
||||
|
||||
|
||||
@@ -80,6 +81,7 @@ class AnthropicToolCallingClient:
|
||||
screenshot: bytes | None,
|
||||
tools: list[ToolSpec],
|
||||
timeout: float,
|
||||
history: list[dict[str, Any]] | None = None,
|
||||
) -> ToolCallDecision:
|
||||
try:
|
||||
response = self._create_message(
|
||||
@@ -87,6 +89,7 @@ class AnthropicToolCallingClient:
|
||||
user_prompt,
|
||||
screenshot,
|
||||
tools,
|
||||
history=history,
|
||||
timeout=timeout,
|
||||
forced=False,
|
||||
)
|
||||
@@ -103,6 +106,7 @@ class AnthropicToolCallingClient:
|
||||
user_prompt,
|
||||
screenshot,
|
||||
tools,
|
||||
history=history,
|
||||
timeout=timeout,
|
||||
forced=True,
|
||||
)
|
||||
@@ -123,6 +127,7 @@ class AnthropicToolCallingClient:
|
||||
screenshot: bytes | None,
|
||||
tools: list[ToolSpec],
|
||||
*,
|
||||
history: list[dict[str, Any]] | None,
|
||||
timeout: float,
|
||||
forced: bool,
|
||||
) -> Any:
|
||||
@@ -146,10 +151,8 @@ class AnthropicToolCallingClient:
|
||||
}
|
||||
],
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": _anthropic_content(user_prompt, screenshot),
|
||||
}
|
||||
*_anthropic_history(history or []),
|
||||
{"role": "user", "content": _anthropic_content(user_prompt, screenshot)},
|
||||
],
|
||||
"tools": [_anthropic_tool(spec) for spec in tools],
|
||||
"tool_choice": {
|
||||
@@ -207,6 +210,7 @@ class OpenAIToolCallingClient:
|
||||
screenshot: bytes | None,
|
||||
tools: list[ToolSpec],
|
||||
timeout: float,
|
||||
history: list[dict[str, Any]] | None = None,
|
||||
) -> ToolCallDecision:
|
||||
try:
|
||||
response = self._create_completion(
|
||||
@@ -214,6 +218,7 @@ class OpenAIToolCallingClient:
|
||||
user_prompt,
|
||||
screenshot,
|
||||
tools,
|
||||
history=history,
|
||||
timeout=timeout,
|
||||
forced=False,
|
||||
)
|
||||
@@ -227,6 +232,7 @@ class OpenAIToolCallingClient:
|
||||
user_prompt,
|
||||
screenshot,
|
||||
tools,
|
||||
history=history,
|
||||
timeout=timeout,
|
||||
forced=True,
|
||||
)
|
||||
@@ -247,6 +253,7 @@ class OpenAIToolCallingClient:
|
||||
screenshot: bytes | None,
|
||||
tools: list[ToolSpec],
|
||||
*,
|
||||
history: list[dict[str, Any]] | None,
|
||||
timeout: float,
|
||||
forced: bool,
|
||||
) -> Any:
|
||||
@@ -257,6 +264,7 @@ class OpenAIToolCallingClient:
|
||||
"timeout": timeout,
|
||||
"messages": [
|
||||
{"role": "system", "content": system_prompt},
|
||||
*_openai_history(history or []),
|
||||
{"role": "user", "content": _openai_content(user_prompt, screenshot)},
|
||||
],
|
||||
"tools": [_openai_tool(spec) for spec in tools],
|
||||
@@ -288,10 +296,14 @@ class OpenAIToolCallingClient:
|
||||
|
||||
def build_client(config: PlannerConfig) -> ToolCallingClient:
|
||||
model = config.resolved_model()
|
||||
if config.provider == "openai":
|
||||
return OpenAIToolCallingClient(model=model)
|
||||
if config.provider in {"openai", "openai_compatible"}:
|
||||
return OpenAIToolCallingClient(
|
||||
model=model, api_key=config.api_key, base_url=config.base_url
|
||||
)
|
||||
return AnthropicToolCallingClient(
|
||||
model=model,
|
||||
api_key=config.api_key,
|
||||
base_url=config.base_url,
|
||||
thinking_budget_tokens=config.thinking_budget_tokens,
|
||||
)
|
||||
|
||||
@@ -316,6 +328,44 @@ def _anthropic_content(
|
||||
return content
|
||||
|
||||
|
||||
def _anthropic_history(history: list[dict[str, Any]]) -> list[dict[str, Any]]:
|
||||
messages: list[dict[str, Any]] = []
|
||||
for index, turn in enumerate(history):
|
||||
result = turn.get("tool_result")
|
||||
if result is None:
|
||||
continue
|
||||
tool_use_id = f"planner-turn-{index}"
|
||||
assistant_content: list[dict[str, Any]] = []
|
||||
rationale = turn.get("rationale")
|
||||
if isinstance(rationale, str) and rationale:
|
||||
assistant_content.append({"type": "text", "text": rationale})
|
||||
assistant_content.append(
|
||||
{
|
||||
"type": "tool_use",
|
||||
"id": tool_use_id,
|
||||
"name": turn["tool_name"],
|
||||
"input": dict(turn.get("arguments") or {}),
|
||||
}
|
||||
)
|
||||
messages.extend(
|
||||
[
|
||||
{"role": "user", "content": turn["user_prompt"]},
|
||||
{"role": "assistant", "content": assistant_content},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "tool_result",
|
||||
"tool_use_id": tool_use_id,
|
||||
"content": json.dumps(result, ensure_ascii=False, default=str),
|
||||
}
|
||||
],
|
||||
},
|
||||
]
|
||||
)
|
||||
return messages
|
||||
|
||||
|
||||
def _anthropic_tool(spec: ToolSpec) -> dict[str, Any]:
|
||||
return {
|
||||
"name": spec.name,
|
||||
@@ -389,6 +439,43 @@ def _openai_content(
|
||||
]
|
||||
|
||||
|
||||
def _openai_history(history: list[dict[str, Any]]) -> list[dict[str, Any]]:
|
||||
messages: list[dict[str, Any]] = []
|
||||
for index, turn in enumerate(history):
|
||||
result = turn.get("tool_result")
|
||||
if result is None:
|
||||
continue
|
||||
tool_call_id = f"planner-turn-{index}"
|
||||
messages.extend(
|
||||
[
|
||||
{"role": "user", "content": turn["user_prompt"]},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": turn.get("rationale"),
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": tool_call_id,
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": turn["tool_name"],
|
||||
"arguments": json.dumps(
|
||||
turn.get("arguments") or {},
|
||||
ensure_ascii=False,
|
||||
),
|
||||
},
|
||||
}
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": tool_call_id,
|
||||
"content": json.dumps(result, ensure_ascii=False, default=str),
|
||||
},
|
||||
]
|
||||
)
|
||||
return messages
|
||||
|
||||
|
||||
def _openai_tool(spec: ToolSpec) -> dict[str, Any]:
|
||||
return {
|
||||
"type": "function",
|
||||
|
||||
@@ -133,4 +133,24 @@ class TaskMetadataStore:
|
||||
def _connect(self) -> sqlite3.Connection:
|
||||
connection = sqlite3.connect(self.db_path)
|
||||
connection.row_factory = sqlite3.Row
|
||||
# A previous interrupted startup can leave a zero-byte SQLite file
|
||||
# behind before ``__init__`` reaches ``_ensure_schema``. Ensure the
|
||||
# table exists on every connection so the console can recover without
|
||||
# manual deletion or database repair.
|
||||
connection.execute(
|
||||
"""
|
||||
create table if not exists tasks (
|
||||
id text primary key,
|
||||
goal text not null,
|
||||
device_id text not null,
|
||||
status text not null,
|
||||
created_at text not null,
|
||||
updated_at text not null,
|
||||
completed_at text,
|
||||
failure_reason text,
|
||||
source_task_id text,
|
||||
source_attempt integer
|
||||
)
|
||||
"""
|
||||
)
|
||||
return connection
|
||||
|
||||
+57
-53
@@ -8,6 +8,7 @@ from core.errors import TaskFailedError
|
||||
from core.models import Bounds, Scene, SceneElement
|
||||
from runtime.ai_planner import AIPlanner
|
||||
from runtime.context import TaskContext
|
||||
from runtime.executor import StepResult
|
||||
from runtime.planner_config import PlannerConfig
|
||||
from runtime.tool_calling_client import ToolCallDecision
|
||||
from runtime.tool_specs import ALL_TOOL_SPECS
|
||||
@@ -26,6 +27,7 @@ class FakeToolCallingClient:
|
||||
screenshot: bytes | None,
|
||||
tools: list[Any],
|
||||
timeout: float,
|
||||
history: list[dict[str, Any]] | None = None,
|
||||
) -> ToolCallDecision:
|
||||
self.calls.append(
|
||||
{
|
||||
@@ -34,6 +36,7 @@ class FakeToolCallingClient:
|
||||
"screenshot": screenshot,
|
||||
"tools": tools,
|
||||
"timeout": timeout,
|
||||
"history": list(history) if history is not None else None,
|
||||
}
|
||||
)
|
||||
return self.decision
|
||||
@@ -242,63 +245,64 @@ def test_ai_planner_propagates_rationale_and_thinking_to_planned_step() -> None:
|
||||
assert steps[0].expected_outcome == "The account settings page is visible."
|
||||
|
||||
|
||||
def test_history_summary_returns_compact_format() -> None:
|
||||
from collections import deque
|
||||
from runtime.ai_planner import _history_summary
|
||||
from world.models import WorldEvent, WorldState
|
||||
|
||||
state = WorldState(
|
||||
history=deque(
|
||||
[
|
||||
WorldEvent(
|
||||
action="tap",
|
||||
success=True,
|
||||
rationale="Opened settings.",
|
||||
arguments={"x": 1, "y": 2},
|
||||
purpose="Open settings.",
|
||||
expected_outcome="Settings is visible.",
|
||||
page="Home",
|
||||
),
|
||||
WorldEvent(
|
||||
action="swipe",
|
||||
success=False,
|
||||
rationale=None,
|
||||
arguments={"start_y": 700, "end_y": 200},
|
||||
page="Settings",
|
||||
),
|
||||
]
|
||||
def test_ai_planner_carries_completed_turn_into_the_next_llm_call() -> None:
|
||||
client = FakeToolCallingClient(
|
||||
ToolCallDecision(
|
||||
tool_name="tap",
|
||||
arguments={"x": 1, "y": 2},
|
||||
text_output="Opening the send control.",
|
||||
purpose="Open the send control.",
|
||||
expected_outcome="The composer is focused.",
|
||||
)
|
||||
)
|
||||
planner = AIPlanner(client=client)
|
||||
context = _context()
|
||||
|
||||
summary = _history_summary(state)
|
||||
first_step = planner.plan(goal=context.goal, scene=_scene(), context=context)[0]
|
||||
context.add_step_result(
|
||||
StepResult(
|
||||
step=first_step,
|
||||
success=True,
|
||||
attempts=1,
|
||||
result={"ok": True},
|
||||
)
|
||||
)
|
||||
planner.plan(goal=context.goal, scene=_scene(), context=context)
|
||||
|
||||
assert summary == [
|
||||
assert client.calls[0]["history"] == []
|
||||
history = client.calls[1]["history"]
|
||||
assert history is not None
|
||||
assert history[0]["tool_name"] == "tap"
|
||||
assert history[0]["arguments"] == {
|
||||
"x": 1,
|
||||
"y": 2,
|
||||
"purpose": "Open the send control.",
|
||||
"expected_outcome": "The composer is focused.",
|
||||
}
|
||||
assert history[0]["tool_result"]["success"] is True
|
||||
assert history[0]["tool_result"]["result"] == {"ok": True}
|
||||
|
||||
|
||||
def test_ai_planner_limits_history_sent_to_the_llm() -> None:
|
||||
client = FakeToolCallingClient(
|
||||
ToolCallDecision(tool_name="tap", arguments={"x": 1, "y": 2})
|
||||
)
|
||||
planner = AIPlanner(client=client, config=PlannerConfig(history_max_turns=2))
|
||||
context = _context()
|
||||
context.planner_history.extend(
|
||||
{
|
||||
"page": "Home",
|
||||
"action": "tap",
|
||||
"arguments": {"x": 1, "y": 2},
|
||||
"rationale": "Opened settings.",
|
||||
"purpose": "Open settings.",
|
||||
"expected_outcome": "Settings is visible.",
|
||||
"success": True,
|
||||
},
|
||||
{
|
||||
"page": "Settings",
|
||||
"action": "swipe",
|
||||
"arguments": {"start_y": 700, "end_y": 200},
|
||||
"user_prompt": f"turn-{index}",
|
||||
"tool_name": "tap",
|
||||
"arguments": {},
|
||||
"rationale": None,
|
||||
"purpose": None,
|
||||
"expected_outcome": None,
|
||||
"success": False,
|
||||
},
|
||||
"tool_result": {"success": True},
|
||||
}
|
||||
for index in range(3)
|
||||
)
|
||||
|
||||
planner.plan(goal=context.goal, scene=_scene(), context=context)
|
||||
|
||||
assert [turn["user_prompt"] for turn in client.calls[0]["history"]] == [
|
||||
"turn-1",
|
||||
"turn-2",
|
||||
]
|
||||
# Must not contain scene element data
|
||||
for entry in summary:
|
||||
assert "scene_summary" not in entry
|
||||
assert "elements" not in entry
|
||||
|
||||
|
||||
def test_history_summary_returns_empty_for_none_world() -> None:
|
||||
from runtime.ai_planner import _history_summary
|
||||
|
||||
assert _history_summary(None) == []
|
||||
|
||||
@@ -30,3 +30,45 @@ def test_device_manager_marks_unreachable_device_offline() -> None:
|
||||
manager.connect("iphone-1", max_retries=2, retry_backoff_seconds=0)
|
||||
|
||||
assert manager.status("iphone-1") == "offline"
|
||||
|
||||
|
||||
def test_probe_marks_connected_device_offline_when_driver_is_unreachable() -> None:
|
||||
class BrokenDriver:
|
||||
def connect(self) -> None:
|
||||
return None
|
||||
|
||||
def screenshot(self) -> bytes:
|
||||
raise RuntimeError("WDA disconnected")
|
||||
|
||||
manager = DeviceManager()
|
||||
manager.register_device("iphone-1", lambda: BrokenDriver()) # type: ignore[arg-type]
|
||||
manager.connect("iphone-1")
|
||||
|
||||
assert manager.probe("iphone-1") is False
|
||||
assert manager.status("iphone-1") == "offline"
|
||||
|
||||
|
||||
def test_probe_uses_health_check_without_capturing_screen() -> None:
|
||||
class HealthCheckedDriver:
|
||||
def __init__(self) -> None:
|
||||
self.health_checks = 0
|
||||
self.screenshots = 0
|
||||
|
||||
def connect(self) -> None:
|
||||
return None
|
||||
|
||||
def screenshot(self) -> bytes:
|
||||
self.screenshots += 1
|
||||
return b"screen"
|
||||
|
||||
def health_check(self) -> None:
|
||||
self.health_checks += 1
|
||||
|
||||
driver = HealthCheckedDriver()
|
||||
manager = DeviceManager()
|
||||
manager.register_device("iphone-1", lambda: driver) # type: ignore[arg-type]
|
||||
manager.connect("iphone-1")
|
||||
|
||||
assert manager.probe("iphone-1") is True
|
||||
assert driver.health_checks == 1
|
||||
assert driver.screenshots == 0
|
||||
|
||||
@@ -64,3 +64,12 @@ def test_mcp_ui_tree_can_include_active_app_info() -> None:
|
||||
"activity": ".MainActivity",
|
||||
}
|
||||
assert isinstance(response["nodes"], list)
|
||||
|
||||
|
||||
def test_tool_handlers_requires_manager() -> None:
|
||||
"""D12: tool_handlers must not silently fall back to DEFAULT_MANAGER."""
|
||||
from api.mcp import tool_handlers
|
||||
import pytest
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
tool_handlers() # type: ignore[call-arg]
|
||||
|
||||
@@ -9,7 +9,6 @@ def test_planner_user_prompt_includes_time_zone_and_configured_device_type() ->
|
||||
prompt = planner_user_prompt(
|
||||
goal="open settings",
|
||||
scene_json={"screen": {"width": 1, "height": 1}, "elements": []},
|
||||
history_summary=[],
|
||||
device_platform="ios",
|
||||
now=datetime(
|
||||
2026,
|
||||
@@ -35,8 +34,16 @@ def test_planner_user_prompt_uses_scene_platform_when_context_is_unavailable() -
|
||||
"elements": [],
|
||||
"app": {"platform": "android"},
|
||||
},
|
||||
history_summary=[],
|
||||
now=datetime(2026, 7, 16, tzinfo=timezone.utc),
|
||||
)
|
||||
|
||||
assert "Device type: android" in prompt
|
||||
|
||||
|
||||
def test_planner_user_prompt_does_not_duplicate_conversation_history() -> None:
|
||||
prompt = planner_user_prompt(
|
||||
goal="open settings",
|
||||
scene_json={"screen": {"width": 1, "height": 1}, "elements": []},
|
||||
)
|
||||
|
||||
assert "Recent history" not in prompt
|
||||
|
||||
@@ -316,8 +316,10 @@ def test_create_mcp_server_registers_skill_tools_when_store_provided(seeded_stor
|
||||
"""Wire-up: api.mcp.create_mcp_server must register skill tools when
|
||||
skill_catalog_store is provided."""
|
||||
from api.mcp import create_mcp_server
|
||||
from device.manager import DeviceManager
|
||||
|
||||
server = create_mcp_server(
|
||||
manager=DeviceManager(),
|
||||
skill_catalog_store=seeded_store,
|
||||
skill_active_subscriptions={"sub-a"},
|
||||
)
|
||||
@@ -329,8 +331,9 @@ def test_create_mcp_server_registers_skill_tools_when_store_provided(seeded_stor
|
||||
def test_create_mcp_server_omits_skill_tools_when_no_store():
|
||||
"""Wire-up must not break existing behavior when no store is provided."""
|
||||
from api.mcp import create_mcp_server
|
||||
from device.manager import DeviceManager
|
||||
|
||||
server = create_mcp_server()
|
||||
server = create_mcp_server(manager=DeviceManager())
|
||||
names = _fastmcp_tool_names(server)
|
||||
assert "list_skills" not in names
|
||||
assert "tap" in names # existing device tools present
|
||||
|
||||
@@ -0,0 +1,10 @@
|
||||
from storage.task_metadata import TaskMetadataStore
|
||||
|
||||
|
||||
def test_store_recovers_from_preexisting_empty_database(tmp_path) -> None:
|
||||
path = tmp_path / "task_progress.sqlite3"
|
||||
path.touch()
|
||||
|
||||
store = TaskMetadataStore(path)
|
||||
|
||||
assert store.list_tasks() == []
|
||||
@@ -480,6 +480,7 @@ dependencies = [
|
||||
{ name = "filelock" },
|
||||
{ name = "httpx" },
|
||||
{ name = "jinja2" },
|
||||
{ name = "mcp" },
|
||||
{ name = "uvicorn", extra = ["standard"] },
|
||||
]
|
||||
|
||||
@@ -491,6 +492,7 @@ requires-dist = [
|
||||
{ name = "filelock", specifier = ">=3.0" },
|
||||
{ name = "httpx", specifier = ">=0.27.0" },
|
||||
{ name = "jinja2", specifier = ">=3.1" },
|
||||
{ name = "mcp", specifier = ">=1.28,<2" },
|
||||
{ name = "uvicorn", extras = ["standard"], specifier = ">=0.30.0" },
|
||||
]
|
||||
|
||||
|
||||
Reference in New Issue
Block a user