Compare commits

...
38 Commits
Author SHA1 Message Date
showtan001 9076f8ddb0 feat: discover and add connected iOS devices
Tests / Test apps.device-host-agent.tests.test_mcp_token.test_load_or_create_concurrent_calls_do_not_corrupt failed
2026-09-07 18:33:49 +08:00
showtan001 8c99dc015a feat: add on-demand device screenshots
Tests / Test apps.device-host-agent.tests.test_mcp_token.test_load_or_create_concurrent_calls_do_not_corrupt failed
2026-08-31 10:43:37 +08:00
showtan001 60ee157e97 feat: preserve planner context across task steps
Tests / Test apps.device-host-agent.tests.test_mcp_token.test_load_or_create_concurrent_calls_do_not_corrupt failed
2026-08-30 22:59:19 +08:00
showtan001 dd8df33910 Log task planner conversations locally
Tests / Test apps.device-host-agent.tests.test_mcp_token.test_load_or_create_concurrent_calls_do_not_corrupt failed
2026-08-30 22:38:28 +08:00
showtan001 5458f3b8a4 Support Anthropic conversation logging and retries
Tests / Test apps.device-host-agent.tests.test_mcp_token.test_load_or_create_concurrent_calls_do_not_corrupt failed
2026-08-30 22:29:19 +08:00
showtan001 44e1a6651a Support multimodal planner scenes without OCR 2026-08-30 21:56:11 +08:00
showtan001 5b8daab457 Recover task metadata store schema on startup 2026-08-30 21:50:57 +08:00
showtan001 697e54427b Bind chat agent sessions to individual devices 2026-08-30 21:48:39 +08:00
showtan001 fdaca7539b Show local conversation activity in web console 2026-08-30 21:46:50 +08:00
showtan001 3c9e65c78e Record local agent reasoning and tool activity
Tests / Test apps.device-host-agent.tests.test_mcp_token.test_load_or_create_concurrent_calls_do_not_corrupt failed
2026-08-29 21:19:30 +08:00
showtan001 71fd182f50 Support multimodal images in local chat agent 2026-08-29 21:17:20 +08:00
showtan001 a315c62f3a Improve local agent chat and device health detection 2026-08-29 16:03:15 +08:00
showtan001 050d1329c4 Add local host mode and configurable LLM providers 2026-08-24 08:49:18 +08:00
q792602257 433ab41f95 Merge branch 'worktree-host-agent-mcp-server' — Host-Agent MCP Server
Tests / Test apps.device-host-agent.tests.test_mcp_token.test_load_or_create_concurrent_calls_do_not_corrupt failed
Adds a Streamable HTTP MCP server (mount /mcp, port 8765) to the
device-host-agent process so Hermes Agent (or any MCP client) can
drive devices directly, coexisting with the Cloud Control Plane
worker path. Per-device session-level locking with 20s TTL,
independent bearer-token auth, and bidirectional cloud ↔ MCP
coordination via a new heartbeat field.

Implementation:
- 4 new modules (mcp_token, mcp_lock, web/mcp_auth, web/mcp)
- Console mount at /mcp with bearer auth sub-app
- Cloud heartbeat payload + scheduler skip MCP-busy devices
- AssignmentExecutor fail-fast reverse check
- CLI mcp-token subcommand
- docs/MCP_INTEGRATION.md + MACOS_IPHONE_SETUP.md section

Spec: docs/superpowers/specs/2026-07-21-host-agent-mcp-server-design.md
Plan: docs/superpowers/plans/2026-07-21-host-agent-mcp-server.md

18 implementation commits ( Tasks 1-15 + final fix wave).
Spec/plan cherry-picks are detected as already-applied via patch-id.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>

# Conflicts:
#	docs/superpowers/specs/2026-07-21-host-agent-mcp-server-design.md
2026-07-21 17:00:29 +08:00
q792602257andClaude Opus 4.6 70e0624a47 fix(host-agent): align MCP integration with mcp SDK 1.28.1 realities
Three final-review deviations closed:

I1 (session-end release): mcp SDK 1.28.1 exposes no per-session
shutdown callback (only a server-level lifespan). Lower the
McpBusyTracker default TTL from 60s to 20s and update spec §6.5,
Q5/R3, D9, and docs/MCP_INTEGRATION.md concurrency section to
document the TTL-only recovery path. 20s is short enough to recover
within one 30s heartbeat interval but long enough that an active
session does not lose its lease during normal operator pauses.

I2 (JSON-RPC error shape): FastMCP Tool.run wraps every non-
UrlElicitationRequiredError exception (including McpError with typed
ErrorData) into ToolError, which the lowlevel call_tool handler
serializes as CallToolResult(isError=true, content=[TextContent(...)]).
There is no public path that surfaces JSON-RPC -32000 with structured
data.busy_owner from a tool call site. Update spec §7 error matrix
and docs/MCP_INTEGRATION.md error table to document the actual wire
shape; busy_owner now lives in the text content.

I3 (typing): mcp_server: Any = None -> FastMCP | None = None via
TYPE_CHECKING, keeping the mcp import lazy (matches precedent
elsewhere in the codebase) while adding static type checking at the
create_console_app boundary.

Tests added (4):
- test_default_ttl_is_20_seconds — locks I1's new default TTL
- test_default_ttl_recovers_dead_session_within_one_window — locks
  I1's recovery semantics (lease sweeped on next read after 20s)
- test_busy_error_wire_shape_is_calltoolresult_iserror — pins I2's
  wire envelope via Tool.run + lowlevel Server._make_error_result
- test_busy_error_text_includes_cloud_assignment_owner — same for
  the cloud_assignment busy_owner branch

Full non-integration suite: 697 passed / 54 deselected (was 693 / 54).

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
2026-07-21 16:25:35 +08:00
q792602257andClaude Opus 4.6 e69cea0245 style: ruff format after MCP server integration
Reformat the files touched by Tasks 1-14 of the host-agent MCP server
plan. No semantic changes; pre-existing format issues in unrelated
files (test_templates, test_skill_sync_wiring, 0010_skill_management,
test_skill_catalog_mcp) left untouched for a separate housekeeping
pass.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
2026-07-21 15:52:32 +08:00
q792602257andClaude Opus 4.6 6d9237a592 test: align skill catalog and migration tests with new MCP API
Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
2026-07-21 15:31:51 +08:00
q792602257 6241fb9d6d docs: add MCP integration guide 2026-07-21 15:24:16 +08:00
q792602257andClaude Opus 4.6 2d0c740c88 feat(host-agent): add mcp-token CLI subcommand
Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
2026-07-21 15:21:32 +08:00
q792602257 ce64c4eb47 feat(host-agent): wire MCP server into create_application 2026-07-21 15:17:28 +08:00
q792602257 ce2469616e feat(host-agent): mount /mcp + surface MCP status in console 2026-07-21 15:09:41 +08:00
q792602257andClaude Opus 4.6 dcb4798408 feat(host-agent): fail-fast cloud assignment when MCP holds device
Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
2026-07-21 15:00:54 +08:00
q792602257 ab15218b27 feat(host-agent): include mcp_busy_device_ids in heartbeat payload 2026-07-21 14:56:21 +08:00
q792602257 b0932dd398 feat(cloud): skip MCP-busy devices in scheduler 2026-07-21 14:37:21 +08:00
q792602257 9e3007e7f6 feat(cloud): accept mcp_busy_device_ids in heartbeat payload 2026-07-21 14:32:44 +08:00
q792602257andClaude Opus 4.6 98089b6748 fix(host-agent): use stable ServerSession id for MCP lock identity
The previous _current_session_id() implementation tried to import a
non-existent get_context() helper, so the production code path always
fell through to the empty _TEST_SESSION_ID ContextVar — meaning every
MCP client shared the empty-string identity and there was no per-session
isolation in production.

Use Context.session (the long-lived ServerSession object) as the source
of identity. id(ctx.session) is stable across every tool call the same
client makes within a Streamable HTTP session, which is exactly what the
busy tracker needs to renew leases.

Wire FastMCP to inject the Context into the wrapper by setting
tool.context_kwarg = "ctx" after swapping tool.fn; wrap the swap in a
defensive try/except that surfaces a FastMcpSdkIncompatibilityError on
future SDK layout drift.

Add 4 tests covering the production path: stability across calls in the
same session, isolation between sessions, fallback to _TEST_SESSION_ID
when no Context is supplied, and verification that the registered tool
declares context_kwarg="ctx".

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
2026-07-21 14:28:50 +08:00
q792602257 b73db01626 feat(host-agent): wrap tool_handlers with busy check + status mapping 2026-07-21 14:09:42 +08:00
q792602257 cf8affe4d7 feat(host-agent): add BearerAuthMiddleware for MCP server 2026-07-21 14:02:20 +08:00
q792602257 61c923b92b feat(host-agent): add McpBusyTracker for per-device session locks 2026-07-21 13:58:37 +08:00
q792602257 c7faee8da3 feat(host-agent): add McpTokenStore for MCP bearer token 2026-07-21 13:55:24 +08:00
q792602257 29b9a8c39a refactor(api): make tool_handlers require a DeviceManager
Eliminates the silent fallback to DEFAULT_MANAGER that produced the
DeviceNotFoundError incident. All existing callers already pass
manager explicitly.
2026-07-21 13:52:01 +08:00
q792602257 d1b0fffabb build(host-agent): add mcp as direct dependency 2026-07-21 13:49:17 +08:00
q792602257andClaude Opus 4.6 47eac0f2a7 docs(superpowers): add host-agent MCP server implementation plan
15-task TDD plan implementing the spec committed in 3b62195. Covers
the four new host-agent modules (mcp_token, mcp_lock, web/mcp_auth,
web/mcp), cloud heartbeat + scheduler coordination, console mount
wiring, CLI subcommand, and documentation.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
2026-07-21 13:47:15 +08:00
q792602257andClaude Opus 4.6 2f8f4a36c6 docs(superpowers): add host-agent MCP server design spec
Design for mounting a Streamable HTTP MCP server inside the host-agent
process so Hermes Agent (or any MCP client) can drive devices directly.
Reuses the existing console FastAPI + uvicorn on port 8765, adds bearer-
token auth, per-device session-level locks with 60s TTL, and cloud
coordination via a new heartbeat field. Cloud scheduler skips devices
reported as MCP-busy.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
2026-07-21 13:47:15 +08:00
q792602257andClaude Opus 4.6 28dccb908c docs(superpowers): add host-agent MCP server implementation plan
15-task TDD plan implementing the spec committed in 3b62195. Covers
the four new host-agent modules (mcp_token, mcp_lock, web/mcp_auth,
web/mcp), cloud heartbeat + scheduler coordination, console mount
wiring, CLI subcommand, and documentation.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
2026-07-21 13:46:27 +08:00
q792602257andClaude Opus 4.6 3b62195a36 docs(superpowers): add host-agent MCP server design spec
Design for mounting a Streamable HTTP MCP server inside the host-agent
process so Hermes Agent (or any MCP client) can drive devices directly.
Reuses the existing console FastAPI + uvicorn on port 8765, adds bearer-
token auth, per-device session-level locks with 60s TTL, and cloud
coordination via a new heartbeat field. Cloud scheduler skips devices
reported as MCP-busy.

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