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"