diff --git a/api/mcp.py b/api/mcp.py index f2861e9..d9c3f85 100644 --- a/api/mcp.py +++ b/api/mcp.py @@ -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") diff --git a/apps/device-host-agent/host_agent/app.py b/apps/device-host-agent/host_agent/app.py index 8dd8053..0ad401c 100644 --- a/apps/device-host-agent/host_agent/app.py +++ b/apps/device-host-agent/host_agent/app.py @@ -1,6 +1,7 @@ from __future__ import annotations import asyncio +import logging from contextlib import suppress from dataclasses import dataclass @@ -21,6 +22,8 @@ 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.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 +31,7 @@ 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 storage.artifact_store import ArtifactStore from storage.device_config import DeviceConfigStore from storage.task_metadata import TaskMetadataStore @@ -201,13 +205,28 @@ 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) executor = AssignmentExecutor( create_execution_factories( resolved_manager, metadata_store=metadata_store, timeline=timeline, host_agent_config=resolved_config, - ) + ), + mcp_busy_tracker=mcp_busy_tracker, + ) + mcp_server = build_mcp_server( + manager=resolved_manager, + mcp_busy_tracker=mcp_busy_tracker, + status_tracker=status_tracker, ) console_app = create_console_app( config=resolved_config, @@ -225,6 +244,9 @@ 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, ) console_server = _EmbeddedConsoleServer( uvicorn.Config( @@ -240,6 +262,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 ), diff --git a/apps/device-host-agent/host_agent/assignment.py b/apps/device-host-agent/host_agent/assignment.py index 42d3146..cdea87b 100644 --- a/apps/device-host-agent/host_agent/assignment.py +++ b/apps/device-host-agent/host_agent/assignment.py @@ -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 diff --git a/apps/device-host-agent/host_agent/cli.py b/apps/device-host-agent/host_agent/cli.py index 4fba57b..67650da 100644 --- a/apps/device-host-agent/host_agent/cli.py +++ b/apps/device-host-agent/host_agent/cli.py @@ -10,6 +10,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): @@ -20,8 +21,16 @@ def main(argv: Sequence[str] | None = None) -> None: 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 +63,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) diff --git a/apps/device-host-agent/host_agent/client.py b/apps/device-host-agent/host_agent/client.py index 76bf447..42a9f10 100644 --- a/apps/device-host-agent/host_agent/client.py +++ b/apps/device-host-agent/host_agent/client.py @@ -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()) diff --git a/apps/device-host-agent/host_agent/heartbeat.py b/apps/device-host-agent/host_agent/heartbeat.py index ecdf2eb..152229b 100644 --- a/apps/device-host-agent/host_agent/heartbeat.py +++ b/apps/device-host-agent/host_agent/heartbeat.py @@ -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,6 +53,7 @@ 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: @@ -57,10 +61,16 @@ class HeartbeatSynchronizer: async def sync_once(self) -> HeartbeatResponse: 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: diff --git a/apps/device-host-agent/host_agent/mcp_lock.py b/apps/device-host-agent/host_agent/mcp_lock.py new file mode 100644 index 0000000..477ddd5 --- /dev/null +++ b/apps/device-host-agent/host_agent/mcp_lock.py @@ -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] diff --git a/apps/device-host-agent/host_agent/mcp_token.py b/apps/device-host-agent/host_agent/mcp_token.py new file mode 100644 index 0000000..3238eb1 --- /dev/null +++ b/apps/device-host-agent/host_agent/mcp_token.py @@ -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 diff --git a/apps/device-host-agent/host_agent/web/app.py b/apps/device-host-agent/host_agent/web/app.py index 6815795..93d2b2c 100644 --- a/apps/device-host-agent/host_agent/web/app.py +++ b/apps/device-host-agent/host_agent/web/app.py @@ -5,7 +5,7 @@ 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 @@ -24,6 +24,8 @@ 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.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 +33,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 +229,9 @@ 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, ) -> FastAPI: app = FastAPI(title="Host Agent Console") cookie_secure = config.console_bind_host not in _LOOPBACK_BIND_HOSTS @@ -231,6 +240,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 +361,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 +372,8 @@ def create_console_app( identity=identity, devices=devices, config=config, + mcp_endpoint=mcp_endpoint, + mcp_busy_devices=mcp_busy_devices, **texts, ) @@ -375,7 +402,18 @@ 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.get("/devices", response_class=HTMLResponse) async def devices_page( diff --git a/apps/device-host-agent/host_agent/web/mcp.py b/apps/device-host-agent/host_agent/web/mcp.py new file mode 100644 index 0000000..f41c91f --- /dev/null +++ b/apps/device-host-agent/host_agent/web/mcp.py @@ -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) diff --git a/apps/device-host-agent/host_agent/web/mcp_auth.py b/apps/device-host-agent/host_agent/web/mcp_auth.py new file mode 100644 index 0000000..e97d94a --- /dev/null +++ b/apps/device-host-agent/host_agent/web/mcp_auth.py @@ -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"}, + ) diff --git a/apps/device-host-agent/host_agent/web/templates/dashboard.html b/apps/device-host-agent/host_agent/web/templates/dashboard.html index 8108d3c..7922c36 100644 --- a/apps/device-host-agent/host_agent/web/templates/dashboard.html +++ b/apps/device-host-agent/host_agent/web/templates/dashboard.html @@ -20,6 +20,24 @@

{{ assignment_text }}

{{ progress_text }}

+
+

MCP

+ + + + + + + +
MCP + {% if mcp_endpoint %} + endpoint {{ mcp_endpoint }}; + {% if mcp_busy_devices %}busy: {{ mcp_busy_devices|join(", ") }}{% else %}idle{% endif %} + {% else %} + not configured + {% endif %} +
+

Devices

diff --git a/apps/device-host-agent/pyproject.toml b/apps/device-host-agent/pyproject.toml index 73392fb..d40a444 100644 --- a/apps/device-host-agent/pyproject.toml +++ b/apps/device-host-agent/pyproject.toml @@ -10,6 +10,7 @@ dependencies = [ "filelock>=3.0", "httpx>=0.27.0", "jinja2>=3.1", + "mcp>=1.28,<2", "uvicorn[standard]>=0.30.0", ] diff --git a/apps/device-host-agent/tests/test_app.py b/apps/device-host-agent/tests/test_app.py index 9b425eb..901981c 100644 --- a/apps/device-host-agent/tests/test_app.py +++ b/apps/device-host-agent/tests/test_app.py @@ -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" diff --git a/apps/device-host-agent/tests/test_assignment.py b/apps/device-host-agent/tests/test_assignment.py index 9818ce2..c1bdeaa 100644 --- a/apps/device-host-agent/tests/test_assignment.py +++ b/apps/device-host-agent/tests/test_assignment.py @@ -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): diff --git a/apps/device-host-agent/tests/test_cli.py b/apps/device-host-agent/tests/test_cli.py index bdfcea0..a590c44 100644 --- a/apps/device-host-agent/tests/test_cli.py +++ b/apps/device-host-agent/tests/test_cli.py @@ -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 diff --git a/apps/device-host-agent/tests/test_client.py b/apps/device-host-agent/tests/test_client.py index e8c70bd..9ec0be7 100644 --- a/apps/device-host-agent/tests/test_client.py +++ b/apps/device-host-agent/tests/test_client.py @@ -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 diff --git a/apps/device-host-agent/tests/test_heartbeat.py b/apps/device-host-agent/tests/test_heartbeat.py index 72e75e3..08a8a0f 100644 --- a/apps/device-host-agent/tests/test_heartbeat.py +++ b/apps/device-host-agent/tests/test_heartbeat.py @@ -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 @@ -60,7 +61,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", @@ -98,7 +101,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 +138,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 +182,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") diff --git a/apps/device-host-agent/tests/test_mcp_auth.py b/apps/device-host-agent/tests/test_mcp_auth.py new file mode 100644 index 0000000..987a6bc --- /dev/null +++ b/apps/device-host-agent/tests/test_mcp_auth.py @@ -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 diff --git a/apps/device-host-agent/tests/test_mcp_lock.py b/apps/device-host-agent/tests/test_mcp_lock.py new file mode 100644 index 0000000..15e09d3 --- /dev/null +++ b/apps/device-host-agent/tests/test_mcp_lock.py @@ -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 diff --git a/apps/device-host-agent/tests/test_mcp_token.py b/apps/device-host-agent/tests/test_mcp_token.py new file mode 100644 index 0000000..edcb617 --- /dev/null +++ b/apps/device-host-agent/tests/test_mcp_token.py @@ -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 diff --git a/apps/device-host-agent/tests/test_web_app.py b/apps/device-host-agent/tests/test_web_app.py index e609ac8..a73fea6 100644 --- a/apps/device-host-agent/tests/test_web_app.py +++ b/apps/device-host-agent/tests/test_web_app.py @@ -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 = { @@ -823,9 +833,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 +967,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 "" 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 "" 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 diff --git a/apps/device-host-agent/tests/test_web_mcp.py b/apps/device-host-agent/tests/test_web_mcp.py new file mode 100644 index 0000000..074fc45 --- /dev/null +++ b/apps/device-host-agent/tests/test_web_mcp.py @@ -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) diff --git a/docs/MACOS_IPHONE_SETUP.md b/docs/MACOS_IPHONE_SETUP.md index 102d93c..8580c13 100644 --- a/docs/MACOS_IPHONE_SETUP.md +++ b/docs/MACOS_IPHONE_SETUP.md @@ -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。 diff --git a/docs/MCP_INTEGRATION.md b/docs/MCP_INTEGRATION.md new file mode 100644 index 0000000..e566e47 --- /dev/null +++ b/docs/MCP_INTEGRATION.md @@ -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: + + /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 " +``` + +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 is busy (held by cloud assignment)"` | +| Device busy (other MCP) | `result.isError = true` | `"device is busy (held by mcp_session:<8-char-prefix>)"` | +| Unknown device | `result.isError = false` | JSON `{"ok": false, "error": "device not found: "}` | +| Tool error | `result.isError = true` | `"Error executing tool : "` | + +## 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. diff --git a/docs/superpowers/specs/2026-07-21-host-agent-mcp-server-design.md b/docs/superpowers/specs/2026-07-21-host-agent-mcp-server-design.md index fbe9065..5a707e9 100644 --- a/docs/superpowers/specs/2026-07-21-host-agent-mcp-server-design.md +++ b/docs/superpowers/specs/2026-07-21-host-agent-mcp-server-design.md @@ -38,7 +38,7 @@ | 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 锁,60s TTL 兜底 | session_id 来自 FastMCP 上下文 | +| 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` 类静默回退地雷 | @@ -102,7 +102,7 @@ class McpDeviceLease: last_seen_at: datetime class McpBusyTracker: - def __init__(self, *, ttl_seconds: float = 60.0, now=None) -> None: ... + 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]: ... @@ -258,32 +258,44 @@ run_async() 主循环启动(行为不变): ### 6.5 Session 结束 / TTL 过期 ``` -正常:Hermes 主动断开 → FastMCP 触发 session shutdown callback - → mcp_busy_tracker.release(session_id) → 该 session 持有的所有 lease 释放 - → 下次心跳 payload 不再包含这些 device_id → cloud 重新视为 idle +正常: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 - → 60s TTL 到期 → 下次 busy_device_ids() 或 snapshot() 调用时 lazy sweep +异常: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 占用目标设备 | JSON-RPC `-32000` + `"device X is busy (held by cloud assignment)"` + `data.busy_owner = "cloud_assignment"` | DEBUG 日志 | -| 3 | 另一 MCP session 占用 | JSON-RPC `-32000` + `"device X is busy (held by another MCP session)"` + `data.busy_owner = "mcp_session:"` | DEBUG 日志 | -| 4 | `device_id` 不存在 | JSON-RPC `-32602` + `"unknown device: X"` | 复用 `call_with_semantic_errors` | -| 5 | 工具底层异常 | JSON-RPC `-32000` + 原异常 message | WARNING + exc_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 : ")])` | 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 倒计时 | 60s 后 lazy sweep | +| 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 | JSON-RPC `-32001` + `"cannot determine MCP session"` | ERROR + exc_info | +| 13 | FastMCP 提取不到 session_id | `CallToolResult(isError=true, content=[TextContent("cannot determine MCP session")])` | ERROR + exc_info | | 14 | host-agent 关停时有 active MCP session | lease 随进程退出消失 | 不需显式清理 | ### 错误返回格式 @@ -292,17 +304,32 @@ run_async() 主循环启动(行为不变): { "jsonrpc": "2.0", "id": "", - "error": { - "code": -32000, - "message": "device phone-1 is busy (held by cloud assignment)", - "data": { - "device_id": "phone-1", - "busy_owner": "cloud_assignment" - } + "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 单元测试 @@ -376,13 +403,13 @@ run_async() 主循环启动(行为不变): - **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(待实测)**:60s TTL 是否合适?→ 实测后调整。太短容易误释放正常 session 的锁;太长 Hermes 崩溃后设备不可用窗口大。 +- **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 崩溃后 60s 内设备不可用。TTL 兜底但窗口存在。**接受,记入 Q5**。 +- **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` 直接化先例一致)。 diff --git a/packages/cloud-platform/cloud/db_models.py b/packages/cloud-platform/cloud/db_models.py index a33c099..e6acbe2 100644 --- a/packages/cloud-platform/cloud/db_models.py +++ b/packages/cloud-platform/cloud/db_models.py @@ -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): diff --git a/packages/cloud-platform/cloud/internal_api/api.py b/packages/cloud-platform/cloud/internal_api/api.py index 057bafa..fb19fdf 100644 --- a/packages/cloud-platform/cloud/internal_api/api.py +++ b/packages/cloud-platform/cloud/internal_api/api.py @@ -185,6 +185,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 diff --git a/packages/cloud-platform/cloud/internal_api/models.py b/packages/cloud-platform/cloud/internal_api/models.py index b72049c..87bf7b3 100644 --- a/packages/cloud-platform/cloud/internal_api/models.py +++ b/packages/cloud-platform/cloud/internal_api/models.py @@ -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): diff --git a/packages/cloud-platform/cloud/migrations/versions/0014_pooled_device_mcp_busy.py b/packages/cloud-platform/cloud/migrations/versions/0014_pooled_device_mcp_busy.py new file mode 100644 index 0000000..3fcce4c --- /dev/null +++ b/packages/cloud-platform/cloud/migrations/versions/0014_pooled_device_mcp_busy.py @@ -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") diff --git a/packages/cloud-platform/cloud/pool.py b/packages/cloud-platform/cloud/pool.py index dfac7c4..a9ddabd 100644 --- a/packages/cloud-platform/cloud/pool.py +++ b/packages/cloud-platform/cloud/pool.py @@ -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: diff --git a/packages/cloud-platform/cloud/scheduler.py b/packages/cloud-platform/cloud/scheduler.py index 2a5a293..83c3626 100644 --- a/packages/cloud-platform/cloud/scheduler.py +++ b/packages/cloud-platform/cloud/scheduler.py @@ -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 ( diff --git a/packages/cloud-platform/cloud/schema.py b/packages/cloud-platform/cloud/schema.py index 6a5acef..0fb55be 100644 --- a/packages/cloud-platform/cloud/schema.py +++ b/packages/cloud-platform/cloud/schema.py @@ -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): diff --git a/packages/cloud-platform/cloud/sql_repository.py b/packages/cloud-platform/cloud/sql_repository.py index 402d876..9ad8ab0 100644 --- a/packages/cloud-platform/cloud/sql_repository.py +++ b/packages/cloud-platform/cloud/sql_repository.py @@ -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)), ) diff --git a/packages/cloud-platform/tests/test_internal_api_models.py b/packages/cloud-platform/tests/test_internal_api_models.py new file mode 100644 index 0000000..b73e26c --- /dev/null +++ b/packages/cloud-platform/tests/test_internal_api_models.py @@ -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 == [] diff --git a/packages/cloud-platform/tests/test_pool.py b/packages/cloud-platform/tests/test_pool.py new file mode 100644 index 0000000..7755306 --- /dev/null +++ b/packages/cloud-platform/tests/test_pool.py @@ -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 diff --git a/packages/cloud-platform/tests/test_scheduler.py b/packages/cloud-platform/tests/test_scheduler.py new file mode 100644 index 0000000..ef78447 --- /dev/null +++ b/packages/cloud-platform/tests/test_scheduler.py @@ -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" diff --git a/tests/test_mcp.py b/tests/test_mcp.py index 049f3a3..313e2e5 100644 --- a/tests/test_mcp.py +++ b/tests/test_mcp.py @@ -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] diff --git a/tests/test_skill_catalog_mcp.py b/tests/test_skill_catalog_mcp.py index 4857da0..881f62c 100644 --- a/tests/test_skill_catalog_mcp.py +++ b/tests/test_skill_catalog_mcp.py @@ -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 diff --git a/uv.lock b/uv.lock index 43221a1..93dde05 100644 --- a/uv.lock +++ b/uv.lock @@ -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" }, ]
MCPMCP