|
|
|
|
@@ -8,6 +8,7 @@ two dispatch tools ``describe_mcp`` and ``call_mcp``.
|
|
|
|
|
from __future__ import annotations
|
|
|
|
|
|
|
|
|
|
import asyncio
|
|
|
|
|
import contextlib
|
|
|
|
|
import json
|
|
|
|
|
import re
|
|
|
|
|
from typing import TYPE_CHECKING, Any
|
|
|
|
|
@@ -29,6 +30,7 @@ from strix.tools.mcp import (
|
|
|
|
|
McpConnectionConfig,
|
|
|
|
|
McpConnectionRequest,
|
|
|
|
|
McpRegistry,
|
|
|
|
|
SupervisedMcpSession,
|
|
|
|
|
attach_mcp_requests,
|
|
|
|
|
call_mcp,
|
|
|
|
|
describe_mcp,
|
|
|
|
|
@@ -38,9 +40,11 @@ from strix.tools.mcp import (
|
|
|
|
|
resolve_mcp_call,
|
|
|
|
|
)
|
|
|
|
|
from strix.tools.mcp import client as mcp_client
|
|
|
|
|
from strix.tools.mcp import session as mcp_session_mod
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
if TYPE_CHECKING:
|
|
|
|
|
from collections.abc import Callable
|
|
|
|
|
from pathlib import Path
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@@ -171,6 +175,13 @@ def _ctx(registry: McpRegistry | None) -> ToolContext[dict[str, Any]]:
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
async def _aclose_all(connections: list[Any]) -> None:
|
|
|
|
|
"""Close every supervised session a connect/attach test opened, so no
|
|
|
|
|
supervising task leaks into the event loop's teardown."""
|
|
|
|
|
for connection in connections:
|
|
|
|
|
await connection.session.aclose()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@pytest.fixture(autouse=True)
|
|
|
|
|
def _clear_mcp_env(monkeypatch: pytest.MonkeyPatch) -> None:
|
|
|
|
|
"""Hide any MCP settings the developer has exported in their own shell."""
|
|
|
|
|
@@ -298,6 +309,8 @@ async def test_connect_returns_sessions_without_registering_agent_tools(
|
|
|
|
|
assert [(c.name, c.tool_count) for c in connections] == [("fs", 2), ("db", 1)]
|
|
|
|
|
assert list(factory.registered_agent_tools()) == before
|
|
|
|
|
|
|
|
|
|
await _aclose_all(connections)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
|
|
|
async def test_tool_count_honors_the_allowlist(monkeypatch: pytest.MonkeyPatch) -> None:
|
|
|
|
|
@@ -308,6 +321,8 @@ async def test_tool_count_honors_the_allowlist(monkeypatch: pytest.MonkeyPatch)
|
|
|
|
|
|
|
|
|
|
assert connections[0].tool_count == 1
|
|
|
|
|
|
|
|
|
|
await _aclose_all(connections)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
|
|
|
async def test_connection_notes_ride_on_the_connection(
|
|
|
|
|
@@ -326,6 +341,8 @@ async def test_connection_notes_ride_on_the_connection(
|
|
|
|
|
|
|
|
|
|
assert connections[0].notes == "Staging analytics DB; read-only."
|
|
|
|
|
|
|
|
|
|
await _aclose_all(connections)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
# --- server build branch -----------------------------------------------------
|
|
|
|
|
|
|
|
|
|
@@ -399,11 +416,18 @@ async def test_list_mcps_returns_connections_with_ids_and_descriptions() -> None
|
|
|
|
|
out = await list_mcps.on_invoke_tool(_ctx(registry), "{}")
|
|
|
|
|
|
|
|
|
|
# ``id`` is the exact connection name describe_mcp/call_mcp accept;
|
|
|
|
|
# ``description`` is the summary's purpose; no tool schemas are included.
|
|
|
|
|
# ``description`` is the summary's purpose; ``dead`` is the connection's live
|
|
|
|
|
# health (both healthy here); no tool schemas are included.
|
|
|
|
|
assert out == {
|
|
|
|
|
"connections": [
|
|
|
|
|
{"id": "fs", "name": "fs", "description": "local files", "tool_count": 2},
|
|
|
|
|
{"id": "db", "name": "db", "description": None, "tool_count": 1},
|
|
|
|
|
{
|
|
|
|
|
"id": "fs",
|
|
|
|
|
"name": "fs",
|
|
|
|
|
"description": "local files",
|
|
|
|
|
"tool_count": 2,
|
|
|
|
|
"dead": False,
|
|
|
|
|
},
|
|
|
|
|
{"id": "db", "name": "db", "description": None, "tool_count": 1, "dead": False},
|
|
|
|
|
]
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
@@ -801,37 +825,82 @@ def test_loader_exclude_selection_drops_named(
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
|
|
|
async def test_connect_cleans_up_when_cancelled_mid_connect(
|
|
|
|
|
async def test_connect_skips_a_connection_whose_connect_is_cancelled(
|
|
|
|
|
monkeypatch: pytest.MonkeyPatch,
|
|
|
|
|
) -> None:
|
|
|
|
|
# Each connection now connects on its own supervising task. A cancellation of
|
|
|
|
|
# one session's connect (the transport scope dying mid-connect) is contained
|
|
|
|
|
# to that task: the connection is skipped and cleaned up, and the run's attach
|
|
|
|
|
# keeps going rather than being cancelled.
|
|
|
|
|
cleaned: list[str] = []
|
|
|
|
|
|
|
|
|
|
class _Tracking(FakeMCPServer):
|
|
|
|
|
def __init__(self, name: str, *, fail_connect: bool = False) -> None:
|
|
|
|
|
def __init__(self, name: str, *, cancel_connect: bool = False) -> None:
|
|
|
|
|
super().__init__(name, [_mcp_tool("t")])
|
|
|
|
|
self._fail_connect = fail_connect
|
|
|
|
|
self._cancel_connect = cancel_connect
|
|
|
|
|
|
|
|
|
|
async def connect(self) -> None:
|
|
|
|
|
if self._fail_connect:
|
|
|
|
|
if self._cancel_connect:
|
|
|
|
|
raise asyncio.CancelledError
|
|
|
|
|
|
|
|
|
|
async def cleanup(self) -> None:
|
|
|
|
|
cleaned.append(self._name)
|
|
|
|
|
|
|
|
|
|
servers = {"good": _Tracking("good"), "bad": _Tracking("bad", fail_connect=True)}
|
|
|
|
|
servers = {"good": _Tracking("good"), "bad": _Tracking("bad", cancel_connect=True)}
|
|
|
|
|
monkeypatch.setattr(mcp_client, "_build_server", lambda config: servers[config.name])
|
|
|
|
|
|
|
|
|
|
configs = [
|
|
|
|
|
McpConnectionConfig(name="good", url="https://mcp.example.com", allowed_tools=["t"]),
|
|
|
|
|
McpConnectionConfig(name="bad", url="https://mcp.example.com", allowed_tools=["t"]),
|
|
|
|
|
]
|
|
|
|
|
configs = [_config("good", ["t"]), _config("bad", ["t"])]
|
|
|
|
|
|
|
|
|
|
connections = await mcp_client.connect_mcp_servers(configs)
|
|
|
|
|
|
|
|
|
|
# The cancelled connect is skipped and cleaned up; the good one is returned.
|
|
|
|
|
assert [c.name for c in connections] == ["good"]
|
|
|
|
|
assert "bad" in cleaned
|
|
|
|
|
|
|
|
|
|
await _aclose_all(connections)
|
|
|
|
|
assert "good" in cleaned
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
|
|
|
async def test_connect_cleans_up_started_sessions_when_attach_is_cancelled(
|
|
|
|
|
monkeypatch: pytest.MonkeyPatch,
|
|
|
|
|
) -> None:
|
|
|
|
|
# If the attach coroutine itself is cancelled (the run going down) while a
|
|
|
|
|
# later connection is still connecting, every session started so far is closed
|
|
|
|
|
# on its own task before the cancellation is re-raised, so nothing is orphaned.
|
|
|
|
|
cleaned: list[str] = []
|
|
|
|
|
|
|
|
|
|
class _Tracking(FakeMCPServer):
|
|
|
|
|
def __init__(self, name: str, *, block_connect: bool = False) -> None:
|
|
|
|
|
super().__init__(name, [_mcp_tool("t")])
|
|
|
|
|
self._block_connect = block_connect
|
|
|
|
|
|
|
|
|
|
async def connect(self) -> None:
|
|
|
|
|
if self._block_connect:
|
|
|
|
|
await asyncio.Event().wait() # never completes
|
|
|
|
|
|
|
|
|
|
async def cleanup(self) -> None:
|
|
|
|
|
cleaned.append(self._name)
|
|
|
|
|
|
|
|
|
|
servers = {"good": _Tracking("good"), "slow": _Tracking("slow", block_connect=True)}
|
|
|
|
|
monkeypatch.setattr(mcp_client, "_build_server", lambda config: servers[config.name])
|
|
|
|
|
|
|
|
|
|
async def _attach() -> list[Any]:
|
|
|
|
|
# Connect "good" first, then hang forever connecting "slow".
|
|
|
|
|
return await mcp_client.connect_mcp_servers(
|
|
|
|
|
[_config("good", ["t"]), _config("slow", ["t"])]
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
task = asyncio.create_task(_attach())
|
|
|
|
|
# Give the loop time to connect good and reach slow's hanging connect.
|
|
|
|
|
for _ in range(100):
|
|
|
|
|
await asyncio.sleep(0)
|
|
|
|
|
task.cancel()
|
|
|
|
|
with pytest.raises(asyncio.CancelledError):
|
|
|
|
|
await mcp_client.connect_mcp_servers(configs)
|
|
|
|
|
await task
|
|
|
|
|
|
|
|
|
|
# The server being connected when cancelled, and the one already connected,
|
|
|
|
|
# are both cleaned up rather than orphaned.
|
|
|
|
|
assert cleaned == ["bad", "good"]
|
|
|
|
|
# The already-connected "good" session was cleaned up, not orphaned.
|
|
|
|
|
assert "good" in cleaned
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
# --- reading a tool call back to the server it went out to -------------------
|
|
|
|
|
@@ -916,6 +985,8 @@ async def test_attach_populates_registry_with_provider_and_transform(
|
|
|
|
|
assert entry.purpose == "Customer DB"
|
|
|
|
|
assert entry.result_transform is transform
|
|
|
|
|
|
|
|
|
|
await _aclose_all(connections)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
|
|
|
async def test_attach_bare_request_matches_the_command_line_shape(
|
|
|
|
|
@@ -933,7 +1004,7 @@ async def test_attach_bare_request_matches_the_command_line_shape(
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
registry = McpRegistry()
|
|
|
|
|
await attach_mcp_requests([McpConnectionRequest(config=config)], registry)
|
|
|
|
|
connections = await attach_mcp_requests([McpConnectionRequest(config=config)], registry)
|
|
|
|
|
|
|
|
|
|
entry = registry.get("db")
|
|
|
|
|
assert entry is not None
|
|
|
|
|
@@ -941,6 +1012,8 @@ async def test_attach_bare_request_matches_the_command_line_shape(
|
|
|
|
|
assert entry.result_transform is None
|
|
|
|
|
assert entry.purpose == "Staging analytics DB; read-only."
|
|
|
|
|
|
|
|
|
|
await _aclose_all(connections)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
|
|
|
async def test_attach_is_fail_open_and_skips_a_failed_connection(
|
|
|
|
|
@@ -970,6 +1043,8 @@ async def test_attach_is_fail_open_and_skips_a_failed_connection(
|
|
|
|
|
assert registry.get("good") is not None
|
|
|
|
|
assert registry.get("bad") is None
|
|
|
|
|
|
|
|
|
|
await _aclose_all(connections)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
# --- provider on the registry ------------------------------------------------
|
|
|
|
|
|
|
|
|
|
@@ -1092,3 +1167,435 @@ async def test_errored_structured_output_is_wrapped_with_success_false() -> None
|
|
|
|
|
# Structured content serializes to a JSON string; it too is wrapped under
|
|
|
|
|
# ``content`` so the failure flag has a top-level dict to ride on.
|
|
|
|
|
assert out == {"success": False, "content": json.dumps({"error": "boom"})}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
# --- per-session isolation: containment, reconnect-retry, mark-dead ----------
|
|
|
|
|
# Each MCP connection's live session is owned by its own supervising task. A
|
|
|
|
|
# background failure in one session is contained to that task: the agent's call
|
|
|
|
|
# comes back as a value, the run keeps going, and the session reconnects once and
|
|
|
|
|
# retries the failed call once before it is marked unavailable.
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
class _DyingHttpServer(FakeMCPServer):
|
|
|
|
|
"""A connected server whose ``call_tool`` fails to model a session death.
|
|
|
|
|
|
|
|
|
|
``death`` is the exception raised on a call: a plain ``Exception`` models an
|
|
|
|
|
HTTP/transport error, and ``asyncio.CancelledError`` models the streamable-HTTP
|
|
|
|
|
transport's task group cancelling the supervising task from a background POST
|
|
|
|
|
error (for example a provider 403). ``alive`` flips to stop dying, so a
|
|
|
|
|
reconnected replacement can succeed.
|
|
|
|
|
"""
|
|
|
|
|
|
|
|
|
|
def __init__(
|
|
|
|
|
self,
|
|
|
|
|
name: str,
|
|
|
|
|
tools: list[MCPTool],
|
|
|
|
|
*,
|
|
|
|
|
death: BaseException,
|
|
|
|
|
alive: bool = False,
|
|
|
|
|
) -> None:
|
|
|
|
|
super().__init__(name, tools)
|
|
|
|
|
self._death = death
|
|
|
|
|
self.alive = alive
|
|
|
|
|
|
|
|
|
|
async def call_tool(
|
|
|
|
|
self,
|
|
|
|
|
tool_name: str,
|
|
|
|
|
arguments: dict[str, Any] | None,
|
|
|
|
|
meta: dict[str, Any] | None = None,
|
|
|
|
|
) -> CallToolResult:
|
|
|
|
|
if not self.alive:
|
|
|
|
|
raise self._death
|
|
|
|
|
return await super().call_tool(tool_name, arguments)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def _secret_config(name: str) -> McpConnectionConfig:
|
|
|
|
|
return McpConnectionConfig(
|
|
|
|
|
name=name,
|
|
|
|
|
url="https://mcp.example.com",
|
|
|
|
|
auth=BearerAuth(token="super-secret-bearer-token-42"),
|
|
|
|
|
allowed_tools=["read_file"],
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
async def _started_session(config: McpConnectionConfig) -> SupervisedMcpSession:
|
|
|
|
|
session = SupervisedMcpSession(config)
|
|
|
|
|
assert await session.start()
|
|
|
|
|
return session
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
|
|
|
async def test_call_mcp_reconnects_and_retries_after_a_session_death(
|
|
|
|
|
monkeypatch: pytest.MonkeyPatch,
|
|
|
|
|
) -> None:
|
|
|
|
|
# The first session dies on its call; the supervisor rebuilds the connection
|
|
|
|
|
# once (reusing the existing _build_server + connect), retries the one call
|
|
|
|
|
# once, and the retry lands on the healthy replacement.
|
|
|
|
|
first = _DyingHttpServer("fs", [_mcp_tool("read_file")], death=ConnectionError("403"))
|
|
|
|
|
second = FakeMCPServer("fs", [_mcp_tool("read_file")])
|
|
|
|
|
built = iter([first, second])
|
|
|
|
|
monkeypatch.setattr(mcp_client, "_build_server", lambda _config: next(built))
|
|
|
|
|
|
|
|
|
|
session = await _started_session(_secret_config("fs"))
|
|
|
|
|
registry = McpRegistry()
|
|
|
|
|
registry.add(name="fs", session=session, tool_count=1)
|
|
|
|
|
|
|
|
|
|
out = await call_mcp.on_invoke_tool(
|
|
|
|
|
_ctx(registry), json.dumps({"connection": "fs", "tool": "read_file"})
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
# The caller gets the tool output as a value, and the retried call ran on the
|
|
|
|
|
# reconnected server.
|
|
|
|
|
assert out == {"type": "text", "text": "routed:read_file"}
|
|
|
|
|
assert second.calls == [("read_file", {})]
|
|
|
|
|
assert session.is_dead is False
|
|
|
|
|
|
|
|
|
|
await session.aclose()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
|
|
|
async def test_call_mcp_marks_connection_dead_when_reconnect_keeps_failing(
|
|
|
|
|
monkeypatch: pytest.MonkeyPatch,
|
|
|
|
|
) -> None:
|
|
|
|
|
# The session dies and the reconnect attempt also fails: the connection is
|
|
|
|
|
# marked dead and the call returns the standard failed-tool output.
|
|
|
|
|
first = _DyingHttpServer("fs", [_mcp_tool("read_file")], death=ConnectionError("403"))
|
|
|
|
|
built = {"n": 0}
|
|
|
|
|
|
|
|
|
|
def _build(_config: McpConnectionConfig) -> MCPServer:
|
|
|
|
|
built["n"] += 1
|
|
|
|
|
if built["n"] == 1:
|
|
|
|
|
return first
|
|
|
|
|
raise ConnectionError("cannot reconnect")
|
|
|
|
|
|
|
|
|
|
monkeypatch.setattr(mcp_client, "_build_server", _build)
|
|
|
|
|
|
|
|
|
|
session = await _started_session(_secret_config("fs"))
|
|
|
|
|
registry = McpRegistry()
|
|
|
|
|
registry.add(name="fs", session=session, tool_count=1)
|
|
|
|
|
|
|
|
|
|
out = await call_mcp.on_invoke_tool(
|
|
|
|
|
_ctx(registry), json.dumps({"connection": "fs", "tool": "read_file"})
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
# A dead connection surfaces as an ordinary failed tool call, not an exception.
|
|
|
|
|
assert isinstance(out, dict)
|
|
|
|
|
assert out["success"] is False
|
|
|
|
|
assert "unavailable" in out["content"]
|
|
|
|
|
assert session.is_dead is True
|
|
|
|
|
|
|
|
|
|
# A later call short-circuits to the same failed output without a new attempt.
|
|
|
|
|
again = await call_mcp.on_invoke_tool(
|
|
|
|
|
_ctx(registry), json.dumps({"connection": "fs", "tool": "read_file"})
|
|
|
|
|
)
|
|
|
|
|
assert again["success"] is False
|
|
|
|
|
|
|
|
|
|
# describe_mcp reports the connection unavailable, and list_mcps still lists it.
|
|
|
|
|
described = await describe_mcp.on_invoke_tool(_ctx(registry), json.dumps({"connection": "fs"}))
|
|
|
|
|
assert "unavailable" in described
|
|
|
|
|
listed = await list_mcps.on_invoke_tool(_ctx(registry), "{}")
|
|
|
|
|
assert [c["id"] for c in listed["connections"]] == ["fs"]
|
|
|
|
|
|
|
|
|
|
await session.aclose()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
async def _pump_until(predicate: Callable[[], bool], *, limit: int = 100) -> None:
|
|
|
|
|
"""Yield to the event loop until ``predicate`` holds, so a background
|
|
|
|
|
supervising task can advance its reconnect without a real timer."""
|
|
|
|
|
for _ in range(limit):
|
|
|
|
|
if predicate():
|
|
|
|
|
return
|
|
|
|
|
await asyncio.sleep(0)
|
|
|
|
|
raise AssertionError("condition not reached")
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
|
|
|
async def test_idle_session_death_self_heals_on_reconnect(
|
|
|
|
|
monkeypatch: pytest.MonkeyPatch,
|
|
|
|
|
) -> None:
|
|
|
|
|
# A session that dies while idle (its supervising task cancelled between calls,
|
|
|
|
|
# modeling the transport scope dying with no call in flight) reconnects once on
|
|
|
|
|
# its own and keeps serving, rather than staying dead until a later call would
|
|
|
|
|
# have triggered a reconnect.
|
|
|
|
|
first = FakeMCPServer("fs", [_mcp_tool("read_file")])
|
|
|
|
|
second = FakeMCPServer("fs", [_mcp_tool("read_file")])
|
|
|
|
|
built = iter([first, second])
|
|
|
|
|
monkeypatch.setattr(mcp_client, "_build_server", lambda _config: next(built))
|
|
|
|
|
|
|
|
|
|
session = await _started_session(_secret_config("fs"))
|
|
|
|
|
registry = McpRegistry()
|
|
|
|
|
registry.add(name="fs", session=session, tool_count=1)
|
|
|
|
|
|
|
|
|
|
assert session._task is not None
|
|
|
|
|
session._task.cancel() # idle transport death: no call in flight
|
|
|
|
|
await _pump_until(lambda: session.server is second)
|
|
|
|
|
assert session.is_dead is False
|
|
|
|
|
|
|
|
|
|
# The reconnected session serves calls normally.
|
|
|
|
|
out = await call_mcp.on_invoke_tool(
|
|
|
|
|
_ctx(registry), json.dumps({"connection": "fs", "tool": "read_file"})
|
|
|
|
|
)
|
|
|
|
|
assert out == {"type": "text", "text": "routed:read_file"}
|
|
|
|
|
|
|
|
|
|
await session.aclose()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
|
|
|
async def test_flapping_idle_session_is_marked_dead_without_looping(
|
|
|
|
|
monkeypatch: pytest.MonkeyPatch,
|
|
|
|
|
) -> None:
|
|
|
|
|
# If a session reconnects after an idle death but dies again before serving any
|
|
|
|
|
# call, the supervisor stops reconnecting and marks the connection dead, so a
|
|
|
|
|
# server that instantly drops on connect cannot spin in a reconnect loop.
|
|
|
|
|
first = FakeMCPServer("fs", [_mcp_tool("read_file")])
|
|
|
|
|
second = FakeMCPServer("fs", [_mcp_tool("read_file")])
|
|
|
|
|
built = iter([first, second])
|
|
|
|
|
monkeypatch.setattr(mcp_client, "_build_server", lambda _config: next(built))
|
|
|
|
|
|
|
|
|
|
session = await _started_session(_secret_config("fs"))
|
|
|
|
|
registry = McpRegistry()
|
|
|
|
|
registry.add(name="fs", session=session, tool_count=1)
|
|
|
|
|
|
|
|
|
|
assert session._task is not None
|
|
|
|
|
# First idle death heals onto the second server (only two builds ever happen).
|
|
|
|
|
session._task.cancel()
|
|
|
|
|
await _pump_until(lambda: session.server is second)
|
|
|
|
|
assert session.is_dead is False
|
|
|
|
|
|
|
|
|
|
# Second idle death before any call is served: give up rather than reconnect.
|
|
|
|
|
session._task.cancel()
|
|
|
|
|
await _pump_until(lambda: session._task is not None and session._task.done())
|
|
|
|
|
assert session.is_dead is True
|
|
|
|
|
|
|
|
|
|
out = await call_mcp.on_invoke_tool(
|
|
|
|
|
_ctx(registry), json.dumps({"connection": "fs", "tool": "read_file"})
|
|
|
|
|
)
|
|
|
|
|
assert out["success"] is False
|
|
|
|
|
assert "unavailable" in out["content"]
|
|
|
|
|
|
|
|
|
|
await session.aclose()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
class _HangingCallServer(FakeMCPServer):
|
|
|
|
|
"""A connected server whose ``call_tool`` never returns, modeling a hung
|
|
|
|
|
in-flight call so teardown can be tested for boundedness."""
|
|
|
|
|
|
|
|
|
|
def __init__(self, name: str, tools: list[MCPTool]) -> None:
|
|
|
|
|
super().__init__(name, tools)
|
|
|
|
|
self.cleaned = False
|
|
|
|
|
|
|
|
|
|
async def call_tool(
|
|
|
|
|
self,
|
|
|
|
|
tool_name: str,
|
|
|
|
|
arguments: dict[str, Any] | None,
|
|
|
|
|
meta: dict[str, Any] | None = None,
|
|
|
|
|
) -> CallToolResult:
|
|
|
|
|
await asyncio.Event().wait()
|
|
|
|
|
raise AssertionError("unreachable")
|
|
|
|
|
|
|
|
|
|
async def cleanup(self) -> None:
|
|
|
|
|
self.cleaned = True
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
class _HangingConnectServer(FakeMCPServer):
|
|
|
|
|
"""A server whose ``connect`` never finishes, so ``start`` blocks on readiness
|
|
|
|
|
and can be cancelled mid-connect."""
|
|
|
|
|
|
|
|
|
|
def __init__(self, name: str, tools: list[MCPTool]) -> None:
|
|
|
|
|
super().__init__(name, tools)
|
|
|
|
|
self.cleaned = False
|
|
|
|
|
|
|
|
|
|
async def connect(self) -> None:
|
|
|
|
|
await asyncio.Event().wait()
|
|
|
|
|
|
|
|
|
|
async def cleanup(self) -> None:
|
|
|
|
|
self.cleaned = True
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
|
|
|
async def test_aclose_is_bounded_when_an_in_flight_call_hangs(
|
|
|
|
|
monkeypatch: pytest.MonkeyPatch,
|
|
|
|
|
) -> None:
|
|
|
|
|
# A hung call must not queue the shutdown sentinel behind itself forever:
|
|
|
|
|
# aclose falls back to cancelling the supervising task, and cleanup still runs.
|
|
|
|
|
monkeypatch.setattr(mcp_session_mod, "_SHUTDOWN_TIMEOUT", 0.2)
|
|
|
|
|
server = _HangingCallServer("fs", [_mcp_tool("read_file")])
|
|
|
|
|
monkeypatch.setattr(mcp_client, "_build_server", lambda _config: server)
|
|
|
|
|
|
|
|
|
|
session = await _started_session(_secret_config("fs"))
|
|
|
|
|
call = asyncio.create_task(session.dispatch("read_file", {}, label="fs_read_file"))
|
|
|
|
|
await asyncio.sleep(0.05) # let the serve loop pick up the request and hang
|
|
|
|
|
|
|
|
|
|
# Must return promptly rather than block on the hung call.
|
|
|
|
|
await asyncio.wait_for(session.aclose(), timeout=3.0)
|
|
|
|
|
assert session._task is not None and session._task.done()
|
|
|
|
|
assert server.cleaned is True
|
|
|
|
|
|
|
|
|
|
# The abandoned caller gets a value (dead), not a hang.
|
|
|
|
|
out = await asyncio.wait_for(call, timeout=3.0)
|
|
|
|
|
assert isinstance(out, dict) and out["success"] is False
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
|
|
|
async def test_aclose_cleans_up_when_connect_is_cancelled_mid_await(
|
|
|
|
|
monkeypatch: pytest.MonkeyPatch,
|
|
|
|
|
) -> None:
|
|
|
|
|
# If the scan is cancelled while start() awaits readiness, the readiness future
|
|
|
|
|
# is cancelled; aclose must not raise on it and must still cancel + clean up the
|
|
|
|
|
# partially connected supervisor.
|
|
|
|
|
server = _HangingConnectServer("fs", [_mcp_tool("read_file")])
|
|
|
|
|
monkeypatch.setattr(mcp_client, "_build_server", lambda _config: server)
|
|
|
|
|
|
|
|
|
|
session = SupervisedMcpSession(_secret_config("fs"))
|
|
|
|
|
start = asyncio.create_task(session.start())
|
|
|
|
|
await asyncio.sleep(0.05) # let the supervisor reach the hanging connect()
|
|
|
|
|
start.cancel()
|
|
|
|
|
with contextlib.suppress(asyncio.CancelledError):
|
|
|
|
|
await start
|
|
|
|
|
|
|
|
|
|
await asyncio.wait_for(session.aclose(), timeout=3.0)
|
|
|
|
|
assert session._task is not None and session._task.done()
|
|
|
|
|
assert server.cleaned is True
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
|
|
|
async def test_a_session_death_is_contained_and_other_connections_survive(
|
|
|
|
|
monkeypatch: pytest.MonkeyPatch,
|
|
|
|
|
) -> None:
|
|
|
|
|
# A background failure that surfaces as a cancellation (the transport scope
|
|
|
|
|
# dying) is contained to that one session: the caller gets a value, not a
|
|
|
|
|
# raised CancelledError, and a second healthy connection keeps working.
|
|
|
|
|
dying = _DyingHttpServer("dying", [_mcp_tool("read_file")], death=asyncio.CancelledError())
|
|
|
|
|
healthy = FakeMCPServer("healthy", [_mcp_tool("read_file")])
|
|
|
|
|
dying_builds = {"n": 0}
|
|
|
|
|
|
|
|
|
|
def _build(config: McpConnectionConfig) -> MCPServer:
|
|
|
|
|
if config.name == "healthy":
|
|
|
|
|
return healthy
|
|
|
|
|
# The dying connection connects once, then its rebuild raises, so it ends
|
|
|
|
|
# up marked dead rather than recovering.
|
|
|
|
|
dying_builds["n"] += 1
|
|
|
|
|
if dying_builds["n"] == 1:
|
|
|
|
|
return dying
|
|
|
|
|
raise ConnectionError("cannot reconnect")
|
|
|
|
|
|
|
|
|
|
monkeypatch.setattr(mcp_client, "_build_server", _build)
|
|
|
|
|
|
|
|
|
|
dying_session = await _started_session(_secret_config("dying"))
|
|
|
|
|
healthy_session = await _started_session(_secret_config("healthy"))
|
|
|
|
|
registry = McpRegistry()
|
|
|
|
|
registry.add(name="dying", session=dying_session, tool_count=1)
|
|
|
|
|
registry.add(name="healthy", session=healthy_session, tool_count=1)
|
|
|
|
|
|
|
|
|
|
dead_out = await call_mcp.on_invoke_tool(
|
|
|
|
|
_ctx(registry), json.dumps({"connection": "dying", "tool": "read_file"})
|
|
|
|
|
)
|
|
|
|
|
# Contained: a value came back rather than a CancelledError tearing down the run.
|
|
|
|
|
assert isinstance(dead_out, dict)
|
|
|
|
|
assert dead_out["success"] is False
|
|
|
|
|
|
|
|
|
|
good_out = await call_mcp.on_invoke_tool(
|
|
|
|
|
_ctx(registry), json.dumps({"connection": "healthy", "tool": "read_file"})
|
|
|
|
|
)
|
|
|
|
|
assert good_out == {"type": "text", "text": "routed:read_file"}
|
|
|
|
|
|
|
|
|
|
await dying_session.aclose()
|
|
|
|
|
await healthy_session.aclose()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
|
|
|
async def test_reconnect_reuses_the_stored_config_and_never_logs_the_token(
|
|
|
|
|
monkeypatch: pytest.MonkeyPatch,
|
|
|
|
|
caplog: pytest.LogCaptureFixture,
|
|
|
|
|
) -> None:
|
|
|
|
|
# The reconnect path rebuilds from the config held on the session, reusing the
|
|
|
|
|
# same bearer token, and that token never reaches a log line, a repr, or the
|
|
|
|
|
# inventory list_mcps emits.
|
|
|
|
|
seen_tokens: list[str | None] = []
|
|
|
|
|
|
|
|
|
|
def _build(config: McpConnectionConfig) -> MCPServer:
|
|
|
|
|
seen_tokens.append(config.auth.token if config.auth else None)
|
|
|
|
|
if len(seen_tokens) == 1:
|
|
|
|
|
return _DyingHttpServer("fs", [_mcp_tool("read_file")], death=ConnectionError("403"))
|
|
|
|
|
return FakeMCPServer("fs", [_mcp_tool("read_file")])
|
|
|
|
|
|
|
|
|
|
monkeypatch.setattr(mcp_client, "_build_server", _build)
|
|
|
|
|
|
|
|
|
|
config = _secret_config("fs")
|
|
|
|
|
token = config.auth.token if config.auth else ""
|
|
|
|
|
session = await _started_session(config)
|
|
|
|
|
registry = McpRegistry()
|
|
|
|
|
entry = registry.add(name="fs", session=session, tool_count=1)
|
|
|
|
|
|
|
|
|
|
with caplog.at_level("DEBUG", logger="strix.tools.mcp.session"):
|
|
|
|
|
out = await call_mcp.on_invoke_tool(
|
|
|
|
|
_ctx(registry), json.dumps({"connection": "fs", "tool": "read_file"})
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
# The retry succeeded, and both the initial connect and the reconnect used the
|
|
|
|
|
# same token from the stored config (never re-fetched).
|
|
|
|
|
assert out == {"type": "text", "text": "routed:read_file"}
|
|
|
|
|
assert seen_tokens == [token, token]
|
|
|
|
|
|
|
|
|
|
# The token appears in no log line, no repr of the session or entry, and not in
|
|
|
|
|
# the inventory the agent sees.
|
|
|
|
|
assert token not in caplog.text
|
|
|
|
|
assert token not in repr(session)
|
|
|
|
|
assert token not in repr(entry)
|
|
|
|
|
listed = await list_mcps.on_invoke_tool(_ctx(registry), "{}")
|
|
|
|
|
assert token not in json.dumps(listed)
|
|
|
|
|
# The config is still reachable in memory for the reconnect path.
|
|
|
|
|
assert entry.config is config
|
|
|
|
|
|
|
|
|
|
await session.aclose()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
# --- connection status signal ------------------------------------------------
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
class _RaisingMCPServer(FakeMCPServer):
|
|
|
|
|
"""A connected server whose every call raises, so the session dies."""
|
|
|
|
|
|
|
|
|
|
async def call_tool(
|
|
|
|
|
self,
|
|
|
|
|
tool_name: str,
|
|
|
|
|
arguments: dict[str, Any] | None,
|
|
|
|
|
meta: dict[str, Any] | None = None,
|
|
|
|
|
) -> CallToolResult:
|
|
|
|
|
raise RuntimeError("connection lost")
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
|
|
|
async def test_session_on_dead_fires_once_on_the_death_transition() -> None:
|
|
|
|
|
# An adopted session with no config cannot reconnect, so the first failed
|
|
|
|
|
# call marks it dead; the on-dead callback fires exactly once, on the edge.
|
|
|
|
|
server = _RaisingMCPServer("db", [_mcp_tool("read")])
|
|
|
|
|
session = SupervisedMcpSession.adopt(server, name="db")
|
|
|
|
|
fires: list[int] = []
|
|
|
|
|
session.set_on_dead(lambda: fires.append(1))
|
|
|
|
|
|
|
|
|
|
out = await session.dispatch("read", {}, label="db_read")
|
|
|
|
|
|
|
|
|
|
assert session.is_dead is True
|
|
|
|
|
assert isinstance(out, dict) and out.get("success") is False
|
|
|
|
|
assert fires == [1]
|
|
|
|
|
|
|
|
|
|
# A later call to the already-dead session must not fire the callback again.
|
|
|
|
|
await session.dispatch("read", {}, label="db_read")
|
|
|
|
|
assert fires == [1]
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def test_registry_statuses_report_the_live_dead_flag_and_provider() -> None:
|
|
|
|
|
registry = McpRegistry()
|
|
|
|
|
alive = SupervisedMcpSession.adopt(FakeMCPServer("a", []), name="a")
|
|
|
|
|
gone = SupervisedMcpSession.adopt(FakeMCPServer("b", []), name="b")
|
|
|
|
|
registry.add(name="a", session=alive, tool_count=2, provider="supabase")
|
|
|
|
|
registry.add(name="b", session=gone, tool_count=1, provider=None)
|
|
|
|
|
gone._mark_dead()
|
|
|
|
|
|
|
|
|
|
statuses = {status.name: status for status in registry.statuses()}
|
|
|
|
|
assert statuses["a"].dead is False
|
|
|
|
|
assert statuses["a"].tool_count == 2
|
|
|
|
|
assert statuses["a"].provider == "supabase"
|
|
|
|
|
assert statuses["b"].dead is True
|
|
|
|
|
assert statuses["b"].provider is None
|
|
|
|
|
|