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>
310 lines
10 KiB
Python
310 lines
10 KiB
Python
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" |