diff --git a/apps/device-host-agent/host_agent/web/mcp.py b/apps/device-host-agent/host_agent/web/mcp.py index 56c2ffe..1846403 100644 --- a/apps/device-host-agent/host_agent/web/mcp.py +++ b/apps/device-host-agent/host_agent/web/mcp.py @@ -22,12 +22,17 @@ 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 FastMCP +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): @@ -40,6 +45,11 @@ class McpDeviceBusyError(Exception): 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 @@ -50,26 +60,23 @@ _TEST_SESSION_ID: contextvars.ContextVar[str] = contextvars.ContextVar( ) -def _current_session_id() -> str: - """Extract session_id from the current FastMCP tool-call context. +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 does not expose a stable ``session_id`` on - ``Context``; the closest analogue is the per-request ``request_id`` - (always a string) plus the long-lived ``session`` object. We try those - first, then fall back to a ``contextvars``-based test override. + 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). """ - try: - from mcp.server.fastmcp import get_context - - ctx = get_context() - request_id = getattr(ctx, "request_id", None) - if isinstance(request_id, str) and request_id: - return request_id - session_id = getattr(ctx, "session_id", None) - if isinstance(session_id, str) and session_id: - return session_id - except Exception: - pass + 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("") @@ -101,7 +108,19 @@ def build_mcp_server( server._tool_manager.add_tool( # type: ignore[attr-defined] raw_handler, name=tool_name ) - server._tool_manager._tools[tool_name].fn = wrapped # type: ignore[attr-defined] + 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 @@ -114,7 +133,8 @@ def _wrap_tool( status_tracker: AgentStatusTracker, ) -> Callable[..., Any]: def wrapped(*args: Any, **kwargs: Any) -> Any: - session_id = _current_session_id() + ctx = kwargs.pop(_CONTEXT_KWARG, None) + session_id = _current_session_id(ctx) device_id = kwargs.get("device_id") if tool_name in _STATUS_TOOLS: @@ -209,7 +229,8 @@ def _call_tool_sync( session_id: str, ) -> Any: """Test helper: invoke a registered tool synchronously with a forced - ``session_id``. Bypasses the HTTP/MCP transport layer to keep tests fast. + ``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 diff --git a/apps/device-host-agent/tests/test_web_mcp.py b/apps/device-host-agent/tests/test_web_mcp.py index 9b958ce..140729a 100644 --- a/apps/device-host-agent/tests/test_web_mcp.py +++ b/apps/device-host-agent/tests/test_web_mcp.py @@ -13,8 +13,11 @@ 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): @@ -223,4 +226,85 @@ def test_manager_required_for_build_mcp_server() -> None: import inspect sig = inspect.signature(build_mcp_server) - assert sig.parameters["manager"].kind == inspect.Parameter.KEYWORD_ONLY \ No newline at end of file + 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" \ No newline at end of file