Files
turnstone/tests/test_mcp_client.py
T
Patrick Buckley 0c2c534c86 fix(mcp): route pool transport lifecycles through per-entry owner tasks (#788)
* fix(mcp): route static transport lifecycles through per-server owner tasks

A crash-looping MCP server drove the mcp-loop thread to a sustained,
climbing 100%+ CPU spin. Root cause: anyio cancel scopes are
host-task-bound, and the static path entered the SDK's transport /
ClientSession task-group scopes from short-lived connect tasks (every
health tick is a new task since #768). Once such a scope was cancelled
after its host task had finished - by anyio's task_done when a
transport child died with the server, or by ClientSession.__aexit__
during a cross-task teardown - CancelScope._deliver_cancellation could
never make progress (task.cancel() on a done task is a no-op) and
re-armed itself via call_soon every loop iteration, forever: ~900k
callbacks/s per zombie scope, one more per flap cycle (verified against
anyio 4.14.1; no upstream fix exists as of that release).

Fix: each static server's transport + session cms are now entered,
parked, and exited by ONE long-lived owner task
(_static_transport_owner), so scopes always have a live host and always
exit in the task that entered them. Teardown follows a one-cancel close
protocol (signal the close event before the first await, graceful
grace, then at most ONE cancel - never a second, which would abandon a
scope exit mid-flight). Connect timeouts now cancel only the waiting
caller; connect failures are delivered through a readiness future;
unrequested owner death (server died under a live session) evicts the
session immediately via a done-callback instead of waiting for the next
liveness ping. A rate-limited, mcp-loop-scoped gc-walk backstop
(_maybe_disarm_orphaned_scopes) disarms any zombie minted by paths not
yet migrated (the oauth_user pool keeps the old cross-task-close shape;
follow-up).

Also fixed: BaseExceptionGroup (BaseException-derived, as raised by
anyio task groups wrapping a stray CancelledError, e.g. an
accept-then-RST server) escaped `except Exception` in _connect_all and
killed it before the health/sweep loops were created - silently
disabling all autonomous recovery. Handled there and in the
health/sweep/eviction loops and the reconnect/refresh callers.

Verified: a live SIGKILL-flap repro went from 130%+ CPU (climbing, one
armed scope per cycle) to 0.3% flat with zero armed scopes; the RST
repro now leaves both background loops alive (previously both silently
dead). New tests: owner-lifecycle + close-protocol units (incl. an
exactly-one-cancel pin), a _connect_all BaseExceptionGroup regression,
a discriminating disarm-sweep test, and a ~10s live SIGKILL-flap smoke
test (real FastMCP subprocess, skips on environment gaps) asserting
zero armed scopes, exactly one live owner, and a post-recovery tool
call. Full 8470-test suite green; ruff+mypy clean.

* fix(mcp): route pool transport lifecycles through per-entry owner tasks

Completes the owner-task migration started for the static path: the
oauth_user pool path had the same latent anyio cancel-scope exposure
(host-task-bound scopes entered by short-lived connect tasks; a scope
cancelled after its host finished re-delivers cancellation via
call_soon forever - the 100%-CPU zombie), previously covered only by
the disarm backstop.

Each (user, server) pool entry's transport + ClientSession cms are now
entered, parked, and exited by ONE long-lived owner task
(_pool_transport_owner). The caller keeps building client_kwargs (the
per-user bearer and, when an auth-capture carrier is active, the
httpx_client_factory response hook) so 401/WWW-Authenticate capture
semantics are unchanged. Teardown is the shared one-cancel close
protocol (_teardown_pool_entry: signal before first await, graceful
grace, at most ONE cancel), used by the connect stale-guard, idle/LRU
eviction, and shutdown (parallel signal-then-reap). Unrequested owner
death evicts the session but keeps the entry and its discovered
catalog, matching the existing evict-session-keep-entry semantics the
auth_401 retry relies on.

Discovery still runs in the connecting caller while the transport is
hosted by the owner, so a transport collapse mid-discovery (e.g. the
SDK tearing its task group down on an upstream 401) cancels the OWNER,
not the caller - a bare await on the response stream would hang until
the 30s phase timeout. _await_pool_discovery races each discovery
await against owner completion and converts owner death into a prompt
ConnectionError (the owner is never cancelled there; teardown owns its
lifecycle). Carrier-first failure classification preserves auth_401
semantics for captured 401s.

With no cross-task stack closes left, _safe_close_stack and
_safe_teardown_on_connect_failure are deleted (zero callers).

Tests: new tests/test_mcp_pool_owner.py pins the pool close protocol
(graceful event-before-await close, exactly-one-cancel escalation,
owner-death eviction retaining entry+catalog, caller-cancel-mid-connect
cm-exit guarantee, factory-present-iff-capture, and the
owner-death-during-discovery fast-fail). 1010 mcp tests and the full
8471-test suite green, including the historical cross-task-anyio
sentinel test_integration_pool_reuse_401_refresh_and_retry_succeeds;
ruff+mypy clean; zero destroyed-task warnings.

* fix(mcp): harden disarm-sweep loop guard and owner BaseException arm

Review follow-ups on the owner-task migration:

- _maybe_disarm_orphaned_scopes now enforces its mcp-loop requirement
  instead of trusting callers: it returns without walking (and without
  advancing the rate-limit clock) unless the currently running loop IS
  self._loop. A suppressed close can fire before start() or after
  shutdown(), where the walk would be wasted at best and a cross-thread
  reach at worst.

- The transport owner's BaseException arm now re-raises non-Exception,
  non-group escapees (KeyboardInterrupt, SystemExit) after delivering
  them to the readiness future - failure delivery is the arm's job;
  swallowing an interpreter-level exit was not.

* fix(mcp): extend owner-death discovery fast-fail to the static path

The static connect path had the same exposure the pool's discovery race
closed: discovery runs in the connecting caller while the transport is
hosted by the owner task, so a transport collapse mid-discovery cancels
the OWNER and the caller's bare await on the response stream hung until
the caller-side attempt timeout (~45s) instead of failing promptly.

_await_pool_discovery is renamed to _await_owner_discovery (it is now
path-neutral) and wired into _connect_one_locked's four discovery
awaits. The helper also converts a discovery future that completes
CANCELLED without the race's own reap (an SDK-internal cancellation
shape) into the same ConnectionError, instead of leaking a bare
CancelledError the caller would misread as its own cancellation.

The pool transport owner's BaseException arm gains the same refinement
the static owner received in review: interpreter-level exits
(KeyboardInterrupt, SystemExit) re-raise after delivery to the
readiness future instead of being swallowed.

Tests: static owner-death-during-discovery fast-fail (<1s vs the ~45s
hang), and a direct pin on the cancelled-discovery-future conversion.

* fix(mcp): replace owner BaseException arm with targeted catch + finally delivery

The owner's failure arm now catches only (BaseExceptionGroup, Exception);
waiter delivery for everything else moves to a finally that resolves the
readiness future with a clean transport-failure ConnectionError before
the task unwinds. Interpreter exits and BaseException-derived library
control-flow escapes propagate from the owner exactly once, uncaught -
and the waiter can never be left hanging on an unresolved future (the
initial _connect_all connect has no outer bound). For SystemExit /
KeyboardInterrupt asyncio additionally stops the loop right after, so
the delivery is load-bearing for the non-exit BaseException shapes and
free for the exits.

Pinned by a test driving a BaseException-derived escape through the
owner: the waiter resolves promptly with ConnectionError while the
escape propagates unswallowed.

* fix(mcp): mirror targeted-catch + finally delivery in the pool owner

Same shape the static owner received in review: the failure arm catches
only (BaseExceptionGroup, Exception), and waiter delivery for anything
else moves to a finally that resolves the readiness future with a clean
ConnectionError before the task unwinds - interpreter exits and
BaseException-derived library escapes propagate exactly once, uncaught,
and the waiter can never be left hanging.

* test(mcp): narrow the escape test's waiter catch to explicit types

* test(mcp): narrow discovery-race waiter catches to explicit types

* refactor(mcp): make reap/synchronization awaits explicit to analyzers

Full-absorb reaps (cancel-then-drain of a future whose outcome is
deliberately consumed) become `await asyncio.gather(x,
return_exceptions=True)` - one line, self-describing, and in the
owner-died discovery reap it is also a small semantic improvement: a
caller cancellation arriving during the reap now propagates instead of
being masked by the ConnectionError. Bare synchronization awaits and
selective suppress blocks in tests keep their raise-through semantics
via throwaway assignment. Applied uniformly across the owner-task
test files, including sites introduced by the static-path PR.
2026-07-06 18:46:28 -07:00

4230 lines
177 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""Tests for turnstone.core.mcp_client — MCP client manager and config loading."""
from __future__ import annotations
import asyncio
import concurrent.futures
import contextlib
import inspect
import json
import time
from contextlib import AsyncExitStack, suppress
from typing import Any
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from tests.conftest import _seed_static_state
from turnstone.core.mcp_client import (
MCPClientManager,
_db_servers_to_config,
_is_dead_transport,
_mcp_to_openai,
load_mcp_config,
)
from turnstone.core.tools import INTERACTIVE_TOOLS, TOOLS, merge_mcp_tools
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _dispatch_stub(mock_future: MagicMock) -> Any:
"""Stand-in for ``asyncio.run_coroutine_threadsafe`` in sync-bridge tests.
Closes the never-scheduled coroutine before handing back the canned
future — a mocked dispatch never awaits it, and an unawaited coroutine
GC-fires "coroutine ... was never awaited" inside whatever unrelated
test happens to be running when collection finally occurs (cross-test
bleed that per-test filterwarnings markers cannot catch).
"""
def _rct(coro: Any, _loop: Any) -> MagicMock:
# Only real coroutines need (or survive) closing — several tests
# dispatch a plain MagicMock return value through this seam.
if inspect.iscoroutine(coro):
coro.close()
return mock_future
return _rct
def _fake_mcp_tool(name: str = "search", description: str = "Search stuff") -> MagicMock:
"""Create a mock MCP tool object matching the SDK's Tool type."""
tool = MagicMock()
tool.name = name
tool.description = description
tool.inputSchema = {
"type": "object",
"properties": {"query": {"type": "string"}},
"required": ["query"],
}
return tool
def _fake_openai_tool(name: str = "mcp__test__search") -> dict[str, Any]:
"""Create a fake OpenAI-format tool dict."""
return {
"type": "function",
"function": {
"name": name,
"description": "[MCP: test] Search stuff",
"parameters": {
"type": "object",
"properties": {"query": {"type": "string"}},
"required": ["query"],
},
},
}
def _fake_mcp_resource(
uri: str = "file:///README.md",
name: str = "readme",
description: str = "Project readme",
mime_type: str = "text/plain",
) -> MagicMock:
"""Create a mock MCP Resource object matching the SDK's Resource type."""
res = MagicMock()
res.uri = uri
res.name = name
res.description = description
res.mimeType = mime_type
return res
def _fake_resource_dict(
uri: str = "file:///README.md",
name: str = "readme",
description: str = "Project readme",
mime_type: str = "text/plain",
server: str = "test",
) -> dict[str, Any]:
"""Create a fake resource dict as stored in per-server state."""
return {
"uri": uri,
"name": name,
"description": description,
"mimeType": mime_type,
"server": server,
}
def _fake_mcp_prompt(
name: str = "code_review",
description: str = "Generate a code review",
arguments: list[dict[str, Any]] | None = None,
) -> MagicMock:
"""Create a mock MCP Prompt object matching the SDK's Prompt type."""
prompt = MagicMock()
prompt.name = name
prompt.description = description
if arguments is None:
arg = MagicMock()
arg.name = "language"
arg.description = "Programming language"
arg.required = True
prompt.arguments = [arg]
else:
mock_args = []
for a in arguments:
arg = MagicMock()
arg.name = a["name"]
arg.description = a.get("description", "")
arg.required = a.get("required", False)
mock_args.append(arg)
prompt.arguments = mock_args
return prompt
def _fake_prompt_dict(
name: str = "mcp__test__code_review",
original_name: str = "code_review",
server: str = "test",
description: str = "Generate a code review",
) -> dict[str, Any]:
"""Create a fake prompt dict as stored in per-server state."""
return {
"name": name,
"original_name": original_name,
"server": server,
"description": description,
"arguments": [
{"name": "language", "description": "Programming language", "required": True}
],
}
@pytest.fixture
def running_loop_mgr():
"""Yield a (mgr, loop, thread) triple with a background loop already running.
Spawns an MCPClientManager with a default config of {"srv": stdio echo}
and a fresh asyncio loop driven by a daemon thread. Tests that need a
different config can mutate ``mgr._server_configs`` directly. The
fixture stops the loop and joins the thread on teardown so each test
leaves a clean slate.
"""
import threading as _threading
cfg = {"srv": {"type": "stdio", "command": "echo"}}
mgr = MCPClientManager(cfg)
loop = asyncio.new_event_loop()
thread = _threading.Thread(target=loop.run_forever, daemon=True)
thread.start()
mgr._loop = loop
try:
yield mgr, loop, thread
finally:
# Drain BEFORE stopping: a task left pending (or finished-but-
# unretrieved) on a stopped loop becomes cross-test global state —
# asyncio reports it at GC time, mid-suite, onto whatever stream
# pytest has attached THEN (the "I/O operation on closed file"
# spew), and a silently-abandoned loop thread keeps running
# manager code against torn-down mocks.
async def _cancel_pending() -> None:
tasks = [t for t in asyncio.all_tasks() if t is not asyncio.current_task()]
for t in tasks:
t.cancel()
await asyncio.gather(*tasks, return_exceptions=True)
with suppress(Exception):
asyncio.run_coroutine_threadsafe(_cancel_pending(), loop).result(timeout=5)
loop.call_soon_threadsafe(loop.stop)
thread.join(timeout=5)
assert not thread.is_alive(), "mcp test loop thread failed to stop within 5s"
loop.close()
# ---------------------------------------------------------------------------
# Schema conversion
# ---------------------------------------------------------------------------
class TestMcpToOpenai:
def test_basic_conversion(self):
tool = _fake_mcp_tool("search_repos", "Search GitHub repos")
result = _mcp_to_openai("github", tool)
assert result["type"] == "function"
func = result["function"]
assert func["name"] == "mcp__github__search_repos"
assert func["description"] == "Search GitHub repos"
assert func["parameters"]["type"] == "object"
assert "query" in func["parameters"]["properties"]
def test_name_prefixing(self):
tool = _fake_mcp_tool("list_files")
result = _mcp_to_openai("fs", tool)
assert result["function"]["name"] == "mcp__fs__list_files"
def test_missing_input_schema(self):
tool = MagicMock()
tool.name = "ping"
tool.description = "Ping the server"
tool.inputSchema = None
result = _mcp_to_openai("test", tool)
assert result["function"]["parameters"] == {"type": "object", "properties": {}}
def test_empty_description(self):
tool = MagicMock()
tool.name = "noop"
tool.description = ""
tool.inputSchema = {"type": "object", "properties": {}}
result = _mcp_to_openai("test", tool)
assert result["function"]["description"] == ""
# ---------------------------------------------------------------------------
# Config loading
# ---------------------------------------------------------------------------
class TestLoadMcpConfig:
def test_load_from_json_file(self, tmp_path):
config_file = tmp_path / "mcp.json"
config_file.write_text(
json.dumps(
{
"mcpServers": {
"github": {
"command": "npx",
"args": ["-y", "@modelcontextprotocol/server-github"],
"env": {"GITHUB_TOKEN": "test"},
}
}
}
)
)
result = load_mcp_config(str(config_file))
assert "github" in result
assert result["github"]["command"] == "npx"
assert result["github"]["env"]["GITHUB_TOKEN"] == "test"
def test_load_from_toml(self):
mock_config = {
"servers": {
"postgres": {
"type": "http",
"url": "https://mcp.example.com/mcp",
}
}
}
with patch("turnstone.core.mcp_client.load_config", return_value=mock_config):
result = load_mcp_config(None)
assert "postgres" in result
assert result["postgres"]["url"] == "https://mcp.example.com/mcp"
def test_empty_when_no_config(self):
with patch("turnstone.core.mcp_client.load_config", return_value={}):
result = load_mcp_config(None)
assert result == {}
def test_json_file_not_found(self, tmp_path):
with patch("turnstone.core.mcp_client.load_config", return_value={}):
result = load_mcp_config(str(tmp_path / "nonexistent.json"))
assert result == {}
def test_toml_config_path_redirect(self):
"""TOML [mcp] config_path redirects to JSON file."""
# load_config returns a section with config_path pointing to a nonexistent file
mock_config = {"config_path": "/tmp/nonexistent_mcp.json"}
with patch("turnstone.core.mcp_client.load_config", return_value=mock_config):
result = load_mcp_config(None)
assert result == {}
def test_invalid_json(self, tmp_path):
config_file = tmp_path / "bad.json"
config_file.write_text("not json")
with patch("turnstone.core.mcp_client.load_config", return_value={}):
result = load_mcp_config(str(config_file))
assert result == {}
class TestDBServersToConfig:
"""``_db_servers_to_config`` shapes DB rows for the MCP client."""
def test_static_streamable_http_row_passes_through(self) -> None:
rows = [
{
"name": "static-srv",
"transport": "streamable-http",
"url": "https://mcp.example.com",
"headers": '{"Authorization": "Bearer token"}',
"auth_type": "static",
}
]
result = _db_servers_to_config(rows)
assert "static-srv" in result
assert result["static-srv"]["url"] == "https://mcp.example.com"
assert result["static-srv"]["headers"] == {"Authorization": "Bearer token"}
def test_db_servers_to_config_skips_oauth_user_rows(self) -> None:
"""Rows with auth_type=oauth_user must be invisible to the static
auto-connect path.
Auto-connecting these with empty headers fails the AS check and
trips the circuit breaker on startup. Per-user OAuth servers
come online lazily once the user has consented.
"""
rows = [
{
"name": "static-srv",
"transport": "streamable-http",
"url": "https://static.example.com",
"headers": "{}",
"auth_type": "static",
},
{
"name": "oauth-srv",
"transport": "streamable-http",
"url": "https://oauth.example.com",
"headers": "{}",
"auth_type": "oauth_user",
},
{
"name": "stdio-srv",
"transport": "stdio",
"command": "echo",
"args": "[]",
"env": "{}",
"auth_type": "none",
},
]
result = _db_servers_to_config(rows)
assert set(result) == {"static-srv", "stdio-srv"}
assert "oauth-srv" not in result
# ---------------------------------------------------------------------------
# merge_mcp_tools
# ---------------------------------------------------------------------------
class TestMergeTools:
def test_merge_preserves_builtin(self):
mcp_tools = [_fake_openai_tool()]
merged = merge_mcp_tools(TOOLS, mcp_tools)
# First N should be built-in
for i, t in enumerate(TOOLS):
assert merged[i] is t
def test_merge_appends_mcp(self):
mcp_tools = [_fake_openai_tool("mcp__a__x"), _fake_openai_tool("mcp__b__y")]
merged = merge_mcp_tools(TOOLS, mcp_tools)
assert len(merged) == len(TOOLS) + 2
assert merged[-2]["function"]["name"] == "mcp__a__x"
assert merged[-1]["function"]["name"] == "mcp__b__y"
def test_merge_empty_mcp(self):
merged = merge_mcp_tools(TOOLS, [])
assert merged == TOOLS
def test_merge_does_not_mutate_input(self):
mcp_tools = [_fake_openai_tool()]
original_len = len(TOOLS)
merge_mcp_tools(TOOLS, mcp_tools)
assert len(TOOLS) == original_len
# ---------------------------------------------------------------------------
# MCPClientManager unit tests (no real MCP servers)
# ---------------------------------------------------------------------------
class TestMCPClientManager:
def test_init_state(self):
mgr = MCPClientManager({"test": {"command": "echo"}})
assert mgr.get_tools() == []
assert mgr.is_mcp_tool("anything") is False
assert mgr.server_count == 0
def test_get_tools_returns_copy(self):
mgr = MCPClientManager({})
mgr._tools = [_fake_openai_tool()]
tools = mgr.get_tools()
assert len(tools) == 1
tools.clear() # mutate the copy
assert len(mgr.get_tools()) == 1 # original unchanged
def test_is_mcp_tool(self):
mgr = MCPClientManager({})
mgr._tool_map["mcp__gh__search"] = ("gh", "search")
assert mgr.is_mcp_tool("mcp__gh__search") is True
assert mgr.is_mcp_tool("bash") is False
def test_server_count(self):
mgr = MCPClientManager({})
_seed_static_state(mgr, "a", session=MagicMock())
_seed_static_state(mgr, "b", session=MagicMock())
assert mgr.server_count == 2
def test_call_tool_sync_unknown_tool(self):
mgr = MCPClientManager({})
with pytest.raises(ValueError, match="Unknown MCP tool"):
mgr.call_tool_sync("mcp__no__such", {})
def test_call_tool_sync_disconnected_server(self):
mgr = MCPClientManager({})
mgr._tool_map["mcp__dead__ping"] = ("dead", "ping")
# No session registered for "dead", no config/loop → reconnect fails
with pytest.raises(RuntimeError, match="not connected"):
mgr.call_tool_sync("mcp__dead__ping", {})
def test_shutdown_on_unstarted_manager(self):
"""shutdown() should not raise when called on a manager that was never started."""
mgr = MCPClientManager({})
mgr.shutdown() # should be a no-op
# -- Phase 7: per-user catalog scoping ---------------------------------
def test_is_mcp_tool_user_id_default_none_unchanged(self):
"""Sanity: default ``user_id=None`` answers static-only.
The legacy single-arg call still works, and unknown names still
return False — Phase 7 adds an optional keyword without
rewriting the static-path semantics.
"""
mgr = MCPClientManager({})
mgr._tool_map["mcp__static__list"] = ("static", "list")
# Legacy single-arg call still works.
assert mgr.is_mcp_tool("mcp__static__list") is True
assert mgr.is_mcp_tool("mcp__static__list", user_id=None) is True
assert mgr.is_mcp_tool("nonexistent") is False
assert mgr.is_mcp_tool("nonexistent", user_id=None) is False
def test_is_mcp_tool_user_keyed_pool_tool(self):
"""A name visible only via ``_user_tool_map`` resolves only for
the matching ``user_id``.
Verifies the new branch: ``_tool_map`` miss + ``user_id`` hit.
"""
mgr = MCPClientManager({})
mgr._user_tool_map["user-1"] = {
"mcp__pool-srv__do": ("pool-srv", "do"),
}
# Visible to user-1.
assert mgr.is_mcp_tool("mcp__pool-srv__do", user_id="user-1") is True
# Invisible to None caller (admin / web-search backend resolution).
assert mgr.is_mcp_tool("mcp__pool-srv__do", user_id=None) is False
# Invisible to a different user.
assert mgr.is_mcp_tool("mcp__pool-srv__do", user_id="user-2") is False
def test_is_mcp_tool_static_wins_for_any_user(self):
"""Static-path tools are visible regardless of ``user_id`` —
the merged view is ``static user-pool``."""
mgr = MCPClientManager({})
mgr._tool_map["mcp__static__list"] = ("static", "list")
# Even for an unknown user, a static tool is still reachable —
# static-path is process-global.
assert mgr.is_mcp_tool("mcp__static__list", user_id="user-1") is True
assert mgr.is_mcp_tool("mcp__static__list", user_id="anybody") is True
def test_get_tools_user_id_none_returns_static_only(self):
"""``get_tools(user_id=None)`` returns the global static catalog.
Pool tools are NEVER included in the default-arg view — that's
the legacy contract every pre-Phase-7 caller relies on.
"""
mgr = MCPClientManager({})
mgr._tools = [_fake_openai_tool("mcp__static__list")]
# Seed a pool entry that should NOT appear in the default view.
from turnstone.core.mcp_client import PoolEntryState
entry = PoolEntryState(key=("user-1", "pool-srv"), open_lock=MagicMock())
entry.tools = [_fake_openai_tool("mcp__pool-srv__do")]
mgr._user_pool_entries[("user-1", "pool-srv")] = entry
tools = mgr.get_tools()
names = [t["function"]["name"] for t in tools]
assert names == ["mcp__static__list"]
def test_get_tools_user_id_merges_pool(self):
"""``get_tools(user_id='user-1')`` merges static + that user's pool tools.
Other users' pool entries MUST NOT leak into the result —
privacy / RBAC invariant.
"""
from turnstone.core.mcp_client import PoolEntryState
mgr = MCPClientManager({})
mgr._tools = [_fake_openai_tool("mcp__static__list")]
e1 = PoolEntryState(key=("user-1", "pool-srv"), open_lock=MagicMock())
e1.tools = [_fake_openai_tool("mcp__pool-srv__do")]
e2 = PoolEntryState(key=("user-2", "pool-srv"), open_lock=MagicMock())
e2.tools = [_fake_openai_tool("mcp__pool-srv__other")]
mgr._user_pool_entries[("user-1", "pool-srv")] = e1
mgr._user_pool_entries[("user-2", "pool-srv")] = e2
# Production invariant: ``_connect_one_pool`` /
# ``_refresh_pool_server_tools`` / ``_evict_session`` /
# ``_close_pool_entry_if_idle`` all call ``_rebuild_user_tool_map``
# immediately after mutating ``_user_pool_entries``. Tests that
# seed pool entries directly must mirror that invariant —
# ``get_tools(user_id=...)`` reads from the ``_user_tools``
# snapshot (built by ``_rebuild_user_tool_map``), never iterating
# ``_user_pool_entries`` directly.
mgr._rebuild_user_tool_map("user-1")
mgr._rebuild_user_tool_map("user-2")
u1_names = [t["function"]["name"] for t in mgr.get_tools(user_id="user-1")]
assert sorted(u1_names) == ["mcp__pool-srv__do", "mcp__static__list"]
u2_names = [t["function"]["name"] for t in mgr.get_tools(user_id="user-2")]
assert sorted(u2_names) == ["mcp__pool-srv__other", "mcp__static__list"]
# Default still global-only — unaffected by either user's entries.
default_names = [t["function"]["name"] for t in mgr.get_tools()]
assert default_names == ["mcp__static__list"]
def test_get_tools_user_id_returns_copies(self):
"""Mirror existing ``test_get_tools_returns_copy``: caller mutation
of the returned list MUST NOT affect the manager's catalog."""
from turnstone.core.mcp_client import PoolEntryState
mgr = MCPClientManager({})
mgr._tools = [_fake_openai_tool("mcp__static__a")]
entry = PoolEntryState(key=("user-1", "pool-srv"), open_lock=MagicMock())
entry.tools = [_fake_openai_tool("mcp__pool-srv__b")]
mgr._user_pool_entries[("user-1", "pool-srv")] = entry
# See note in ``test_get_tools_user_id_merges_pool``.
mgr._rebuild_user_tool_map("user-1")
tools = mgr.get_tools(user_id="user-1")
assert len(tools) == 2
tools.clear()
# Re-fetch — original catalog unchanged.
assert len(mgr.get_tools(user_id="user-1")) == 2
def test_get_tools_user_with_none_tools_skipped(self):
"""A pool entry that hasn't completed discovery (``entry.tools is None``)
contributes no tools — the merged view skips it cleanly."""
from turnstone.core.mcp_client import PoolEntryState
mgr = MCPClientManager({})
mgr._tools = [_fake_openai_tool("mcp__static__a")]
# Brand-new pool entry, discovery not yet run.
entry = PoolEntryState(key=("user-1", "pool-srv"), open_lock=MagicMock())
assert entry.tools is None
mgr._user_pool_entries[("user-1", "pool-srv")] = entry
# Rebuild observes ``entry.tools is None`` and skips this entry.
mgr._rebuild_user_tool_map("user-1")
names = [t["function"]["name"] for t in mgr.get_tools(user_id="user-1")]
assert names == ["mcp__static__a"]
def test_rebuild_user_tool_map_populates(self):
"""``_rebuild_user_tool_map`` materializes the per-user index from
pool entries owned by that user."""
from turnstone.core.mcp_client import PoolEntryState
mgr = MCPClientManager({})
entry = PoolEntryState(key=("user-1", "pool-srv"), open_lock=MagicMock())
entry.tools = [_fake_openai_tool("mcp__pool-srv__do")]
mgr._user_pool_entries[("user-1", "pool-srv")] = entry
mgr._rebuild_user_tool_map("user-1")
assert mgr._user_tool_map["user-1"] == {"mcp__pool-srv__do": ("pool-srv", "do")}
# Sibling _user_tools cache (bug-1 fix) MUST be populated alongside
# the map — otherwise get_tools(user_id="user-1") would silently
# return the static-only view despite is_mcp_tool returning True.
assert mgr._user_tools["user-1"] == [_fake_openai_tool("mcp__pool-srv__do")]
def test_rebuild_user_tool_map_drops_empty_user(self):
"""Rebuilding for a user with no pool entries removes the key
rather than retaining an empty-dict sentinel."""
from turnstone.core.mcp_client import PoolEntryState
mgr = MCPClientManager({})
entry = PoolEntryState(key=("user-1", "pool-srv"), open_lock=MagicMock())
entry.tools = [_fake_openai_tool("mcp__pool-srv__do")]
mgr._user_pool_entries[("user-1", "pool-srv")] = entry
mgr._rebuild_user_tool_map("user-1")
assert "user-1" in mgr._user_tool_map
assert "user-1" in mgr._user_tools
# Drop the entry, rebuild — user_id key should be removed from BOTH
# the map and the sibling tool list (bug-1 fix). A drop in only one
# would leave get_tools and is_mcp_tool out of sync.
mgr._user_pool_entries.clear()
mgr._rebuild_user_tool_map("user-1")
assert "user-1" not in mgr._user_tool_map
assert "user-1" not in mgr._user_tools
def test_rebuild_user_tool_map_isolates_users(self):
"""Rebuilding for ``user-1`` MUST NOT touch ``user-2``'s entry."""
from turnstone.core.mcp_client import PoolEntryState
mgr = MCPClientManager({})
e1 = PoolEntryState(key=("user-1", "pool-srv"), open_lock=MagicMock())
e1.tools = [_fake_openai_tool("mcp__pool-srv__one")]
e2 = PoolEntryState(key=("user-2", "pool-srv"), open_lock=MagicMock())
e2.tools = [_fake_openai_tool("mcp__pool-srv__two")]
mgr._user_pool_entries[("user-1", "pool-srv")] = e1
mgr._user_pool_entries[("user-2", "pool-srv")] = e2
mgr._rebuild_user_tool_map("user-1")
mgr._rebuild_user_tool_map("user-2")
assert mgr._user_tool_map["user-1"] == {"mcp__pool-srv__one": ("pool-srv", "one")}
assert mgr._user_tool_map["user-2"] == {"mcp__pool-srv__two": ("pool-srv", "two")}
# Clear user-1's entry only; rebuild user-1; user-2 must remain.
mgr._user_pool_entries.pop(("user-1", "pool-srv"))
mgr._rebuild_user_tool_map("user-1")
assert "user-1" not in mgr._user_tool_map
assert mgr._user_tool_map["user-2"] == {"mcp__pool-srv__two": ("pool-srv", "two")}
def test_rebuild_user_tool_map_does_not_touch_static(self):
"""Invariant 1: per-user rebuild must NOT mutate ``_tool_map``."""
from turnstone.core.mcp_client import PoolEntryState
mgr = MCPClientManager({})
mgr._tool_map["mcp__static__list"] = ("static", "list")
entry = PoolEntryState(key=("user-1", "pool-srv"), open_lock=MagicMock())
entry.tools = [_fake_openai_tool("mcp__pool-srv__do")]
mgr._user_pool_entries[("user-1", "pool-srv")] = entry
before = dict(mgr._tool_map)
mgr._rebuild_user_tool_map("user-1")
assert mgr._tool_map == before
# ---------------------------------------------------------------------------
# Session integration (mock MCP client)
# ---------------------------------------------------------------------------
class TestSessionIntegration:
@pytest.fixture()
def tmp_db(self, tmp_path):
from turnstone.core.storage import init_storage, reset_storage
reset_storage()
init_storage("sqlite", path=str(tmp_path / "test.db"), run_migrations=False)
yield
reset_storage()
def _make_session(self, mcp_client=None, **kwargs):
from turnstone.core.session import ChatSession
defaults: dict[str, Any] = dict(
client=MagicMock(),
model="test-model",
ui=MagicMock(),
instructions=None,
temperature=0.5,
max_tokens=4096,
tool_timeout=30,
mcp_client=mcp_client,
)
defaults.update(kwargs)
return ChatSession(**defaults)
def test_session_without_mcp(self, tmp_db):
session = self._make_session(mcp_client=None)
# Interactive session surface — coordinator tools excluded.
assert session._tools is INTERACTIVE_TOOLS
assert session._mcp_client is None
def test_session_with_mcp(self, tmp_db):
mock_mcp = MagicMock()
mock_mcp.get_tools.return_value = [_fake_openai_tool()]
session = self._make_session(mcp_client=mock_mcp)
assert len(session._tools) == len(INTERACTIVE_TOOLS) + 1
assert session._tools[-1]["function"]["name"] == "mcp__test__search"
def test_task_tools_include_mcp(self, tmp_db):
from turnstone.core.tools import TASK_AGENT_TOOLS
mock_mcp = MagicMock()
mock_mcp.get_tools.return_value = [_fake_openai_tool()]
session = self._make_session(mcp_client=mock_mcp)
assert len(session._task_tools) == len(TASK_AGENT_TOOLS) + 1
def test_prepare_mcp_tool(self, tmp_db):
mock_mcp = MagicMock()
mock_mcp.get_tools.return_value = [_fake_openai_tool()]
mock_mcp.is_mcp_tool.return_value = True
session = self._make_session(mcp_client=mock_mcp)
tc = {
"id": "call_123",
"function": {
"name": "mcp__test__search",
"arguments": '{"query": "hello"}',
},
}
prepared = session._prepare_tool(tc)
assert prepared["func_name"] == "mcp__test__search"
assert prepared["needs_approval"] is True
assert "mcp:test/search" in prepared["header"]
assert callable(prepared["execute"])
def test_unknown_tool_without_mcp(self, tmp_db):
session = self._make_session(mcp_client=None)
tc = {
"id": "call_456",
"function": {"name": "nonexistent", "arguments": "{}"},
}
prepared = session._prepare_tool(tc)
assert "error" in prepared
assert "Unknown tool" in prepared["error"]
# Error lists available tools so the model can self-correct
assert "bash" in prepared["error"]
# Surfaces warning to user
session.ui.on_error.assert_called_once()
assert "nonexistent" in session.ui.on_error.call_args[0][0]
def test_prepare_tool_strips_whitespace_from_name(self, tmp_db):
"""Local models may produce tool names with leading/trailing whitespace."""
session = self._make_session(mcp_client=None)
tc = {
"id": "call_strip",
"function": {"name": " bash\n", "arguments": '{"command": "echo hi"}'},
}
prepared = session._prepare_tool(tc)
assert prepared["func_name"] == "bash"
assert "error" not in prepared
def test_prepare_tool_malformed_json_surfaces_error(self, tmp_db):
"""Malformed JSON args should surface a warning to the user and
give the model a hint about expected format."""
session = self._make_session(mcp_client=None)
tc = {
"id": "call_bad",
"function": {"name": "bash", "arguments": "{command: echo hi}"},
}
prepared = session._prepare_tool(tc)
assert "error" in prepared
assert "JSON parse error" in prepared["error"]
assert "command" in prepared["error"] # hint about expected key
assert "Please retry" in prepared["error"]
# User-facing warning
session.ui.on_error.assert_called_once()
assert "Malformed tool call" in session.ui.on_error.call_args[0][0]
def test_ensure_tool_call_ids_dict(self, tmp_db):
"""_ensure_tool_call_ids fills empty IDs on streaming-style dict."""
from turnstone.core.session import ChatSession
tool_calls_acc = {
0: {"id": "", "function": {"name": "bash", "arguments": "{}"}},
1: {"id": "", "function": {"name": "read_file", "arguments": "{}"}},
}
ChatSession._ensure_tool_call_ids(tool_calls_acc)
ids = [tc["id"] for tc in tool_calls_acc.values()]
assert all(id_.startswith("call_") for id_ in ids)
assert len(set(ids)) == 2 # unique
def test_ensure_tool_call_ids_list(self, tmp_db):
"""_ensure_tool_call_ids fills empty IDs on list (agent path)."""
from turnstone.core.session import ChatSession
tool_calls = [
{"id": None, "function": {"name": "bash", "arguments": "{}"}},
{"id": "call_existing", "function": {"name": "bash", "arguments": "{}"}},
]
ChatSession._ensure_tool_call_ids(tool_calls)
assert tool_calls[0]["id"].startswith("call_")
assert tool_calls[1]["id"] == "call_existing" # preserved
def test_mcp_command_no_client(self, tmp_db):
session = self._make_session(mcp_client=None)
session.handle_command("/mcp")
session.ui.on_info.assert_called_once()
assert "No MCP servers" in session.ui.on_info.call_args[0][0]
def test_mcp_command_with_tools(self, tmp_db):
mock_mcp = MagicMock()
mock_mcp.get_tools.return_value = [_fake_openai_tool()]
session = self._make_session(mcp_client=mock_mcp)
session.handle_command("/mcp")
session.ui.on_info.assert_called_once()
output = session.ui.on_info.call_args[0][0]
assert "MCP tools (1)" in output
assert "mcp__test__search" in output
def test_exec_mcp_tool(self, tmp_db):
mock_mcp = MagicMock()
mock_mcp.get_tools.return_value = [_fake_openai_tool()]
mock_mcp.is_mcp_tool.return_value = True
mock_mcp.call_tool_sync.return_value = "result text"
session = self._make_session(mcp_client=mock_mcp)
item = {
"call_id": "call_789",
"mcp_func_name": "mcp__test__search",
"mcp_args": {"query": "hello"},
}
call_id, output = session._exec_mcp_tool(item)
assert call_id == "call_789"
assert output == "result text"
mock_mcp.call_tool_sync.assert_called_once_with(
"mcp__test__search",
{"query": "hello"},
user_id=None,
timeout=30,
is_interactive_for_consent=True,
)
def test_exec_mcp_tool_error(self, tmp_db):
mock_mcp = MagicMock()
mock_mcp.get_tools.return_value = [_fake_openai_tool()]
mock_mcp.is_mcp_tool.return_value = True
mock_mcp.call_tool_sync.side_effect = RuntimeError("server crashed")
session = self._make_session(mcp_client=mock_mcp)
item = {
"call_id": "call_err",
"mcp_func_name": "mcp__test__search",
"mcp_args": {"query": "hello"},
}
call_id, output = session._exec_mcp_tool(item)
assert call_id == "call_err"
assert "MCP tool error" in output
assert "server crashed" in output
# -- Phase 7: per-user catalog scoping ---------------------------------
def test_session_passes_user_id_to_get_tools(self, tmp_db):
"""ChatSession threads its ``user_id`` into ``get_tools`` so the
merged static + pool view is scoped to the session's user."""
mock_mcp = MagicMock()
mock_mcp.get_tools.return_value = [_fake_openai_tool()]
self._make_session(mcp_client=mock_mcp, user_id="user-7")
mock_mcp.get_tools.assert_called_with(user_id="user-7")
def test_session_get_tools_empty_user_id_passes_none(self, tmp_db):
"""Sentinel ``user_id=""`` (CLI / service / unknown) collapses to
``user_id=None`` so the static-only view is returned."""
mock_mcp = MagicMock()
mock_mcp.get_tools.return_value = [_fake_openai_tool()]
self._make_session(mcp_client=mock_mcp, user_id="")
mock_mcp.get_tools.assert_called_with(user_id=None)
def test_session_passes_user_id_to_add_listener(self, tmp_db):
"""ChatSession registers its tool-change listener under its own
``user_id`` so pool-only changes for OTHER users do not fire it."""
mock_mcp = MagicMock()
mock_mcp.get_tools.return_value = [_fake_openai_tool()]
self._make_session(mcp_client=mock_mcp, user_id="user-7")
# ``add_listener`` was called with ``user_id="user-7"``.
listener_calls = mock_mcp.add_listener.call_args_list
assert listener_calls, "ChatSession did not register a tool listener"
first_call = listener_calls[0]
assert first_call.kwargs.get("user_id") == "user-7"
def test_session_close_removes_listener_with_same_user_id(self, tmp_db):
"""R4 critical: register and remove MUST agree on ``user_id`` —
the listener identity is ``(user_id, callback)``."""
mock_mcp = MagicMock()
mock_mcp.get_tools.return_value = [_fake_openai_tool()]
session = self._make_session(mcp_client=mock_mcp, user_id="user-7")
session.close()
# ``remove_listener`` must be called with the same ``user_id``.
remove_calls = mock_mcp.remove_listener.call_args_list
assert remove_calls, "ChatSession.close did not unregister a tool listener"
first_remove = remove_calls[0]
assert first_remove.kwargs.get("user_id") == "user-7"
# And the callback identity must match what was registered.
registered_cb = mock_mcp.add_listener.call_args_list[0].args[0]
removed_cb = first_remove.args[0]
assert registered_cb is removed_cb
def test_session_unknown_tool_lists_user_scoped_catalog(self, tmp_db):
"""The "Unknown tool" error message lists tools the session can
actually invoke — drawn from the merged user-scoped catalog,
not the manager's private static-only ``_tool_map``."""
mock_mcp = MagicMock()
# Pretend the user's merged view contains a static + pool entry.
mock_mcp.get_tools.return_value = [
_fake_openai_tool("mcp__static__list"),
_fake_openai_tool("mcp__pool-srv__do"),
]
mock_mcp.is_mcp_tool.return_value = False
session = self._make_session(mcp_client=mock_mcp, user_id="user-7")
# Reset the call counter so we observe only the _prepare_tool call.
mock_mcp.get_tools.reset_mock()
tc = {
"id": "call_unknown",
"function": {"name": "no_such_tool", "arguments": "{}"},
}
prepared = session._prepare_tool(tc)
assert "error" in prepared
# The error mentions both static and pool tools — proves we're
# consulting the merged catalog rather than ``_tool_map``.
assert "mcp__static__list" in prepared["error"]
assert "mcp__pool-srv__do" in prepared["error"]
# And the catalog request was scoped to this session's user.
assert any(
call.kwargs.get("user_id") == "user-7" for call in mock_mcp.get_tools.call_args_list
)
# ---------------------------------------------------------------------------
# Server name validation
# ---------------------------------------------------------------------------
class TestServerNameValidation:
def test_double_underscore_in_name(self):
"""Server names with __ should be rejected during _connect_one."""
import asyncio
async def _run() -> None:
mgr = MCPClientManager({"my__bad": {"command": "echo"}})
async with AsyncExitStack() as stack:
mgr._exit_stack = stack
await mgr._connect_one("my__bad", {"command": "echo"})
# Should not have connected
assert "my__bad" not in mgr._static_servers
assert mgr.get_tools() == []
asyncio.run(_run())
# ---------------------------------------------------------------------------
# create_mcp_client guard
# ---------------------------------------------------------------------------
class TestCreateMcpClient:
def test_returns_none_when_no_config(self):
with patch("turnstone.core.mcp_client.load_mcp_config", return_value={}):
from turnstone.core.mcp_client import create_mcp_client
result = create_mcp_client()
assert result is None
# ---------------------------------------------------------------------------
# Tool refresh — _rebuild_tools, _refresh_server, listeners
# ---------------------------------------------------------------------------
class TestRebuildTools:
def test_rebuild_from_per_server(self):
mgr = MCPClientManager({})
_seed_static_state(mgr, "github", tools=[_fake_openai_tool("mcp__github__search")])
_seed_static_state(mgr, "slack", tools=[_fake_openai_tool("mcp__slack__send")])
mgr._rebuild_tools()
assert len(mgr._tools) == 2
names = {t["function"]["name"] for t in mgr._tools}
assert names == {"mcp__github__search", "mcp__slack__send"}
assert mgr._tool_map["mcp__github__search"] == ("github", "search")
assert mgr._tool_map["mcp__slack__send"] == ("slack", "send")
def test_rebuild_copy_on_write(self):
mgr = MCPClientManager({})
_seed_static_state(mgr, "a", tools=[_fake_openai_tool("mcp__a__x")])
mgr._rebuild_tools()
old_tools = mgr._tools
old_map = mgr._tool_map
_seed_static_state(mgr, "b", tools=[_fake_openai_tool("mcp__b__y")])
mgr._rebuild_tools()
assert mgr._tools is not old_tools
assert mgr._tool_map is not old_map
def test_rebuild_empty(self):
mgr = MCPClientManager({})
mgr._static_servers = {}
mgr._rebuild_tools()
assert mgr._tools == []
assert mgr._tool_map == {}
class TestRefreshServer:
@staticmethod
def _add_empty_resource_prompt_mocks(
mgr: MCPClientManager, server_name: str, mock_session: MagicMock
) -> None:
"""Add empty list_resources/list_prompts mocks so _refresh_server works."""
_seed_static_state(mgr, server_name, supports_resources=True, supports_prompts=True)
empty_res = MagicMock()
empty_res.resources = []
mock_session.list_resources = AsyncMock(return_value=empty_res)
empty_tmpl = MagicMock()
empty_tmpl.resourceTemplates = []
mock_session.list_resource_templates = AsyncMock(return_value=empty_tmpl)
empty_prompts = MagicMock()
empty_prompts.prompts = []
mock_session.list_prompts = AsyncMock(return_value=empty_prompts)
def test_refresh_detects_added_tools(self):
async def _run() -> None:
mgr = MCPClientManager({})
mock_session = MagicMock()
mock_result = MagicMock()
mock_result.tools = [
_fake_mcp_tool("search"),
_fake_mcp_tool("create"), # new tool
]
mock_session.list_tools = AsyncMock(return_value=mock_result)
self._add_empty_resource_prompt_mocks(mgr, "github", mock_session)
_seed_static_state(
mgr,
"github",
session=mock_session,
tools=[_fake_openai_tool("mcp__github__search")],
)
mgr._rebuild_tools()
added, removed = await mgr._refresh_server("github")
assert "mcp__github__create" in added
assert removed == []
assert len(mgr._tools) == 2
asyncio.run(_run())
def test_refresh_detects_removed_tools(self):
async def _run() -> None:
mgr = MCPClientManager({})
mock_session = MagicMock()
mock_result = MagicMock()
mock_result.tools = [] # all tools removed
mock_session.list_tools = AsyncMock(return_value=mock_result)
self._add_empty_resource_prompt_mocks(mgr, "github", mock_session)
_seed_static_state(
mgr,
"github",
session=mock_session,
tools=[_fake_openai_tool("mcp__github__search")],
)
mgr._rebuild_tools()
added, removed = await mgr._refresh_server("github")
assert added == []
assert "mcp__github__search" in removed
assert mgr._tools == []
asyncio.run(_run())
def test_refresh_no_changes(self):
async def _run() -> None:
mgr = MCPClientManager({})
mock_session = MagicMock()
mock_result = MagicMock()
mock_result.tools = [_fake_mcp_tool("search")]
mock_session.list_tools = AsyncMock(return_value=mock_result)
self._add_empty_resource_prompt_mocks(mgr, "github", mock_session)
_seed_static_state(
mgr,
"github",
session=mock_session,
tools=[_fake_openai_tool("mcp__github__search")],
)
mgr._rebuild_tools()
added, removed = await mgr._refresh_server("github")
assert added == []
assert removed == []
asyncio.run(_run())
def test_refresh_disconnected_raises(self):
async def _run() -> None:
mgr = MCPClientManager({})
with pytest.raises(RuntimeError, match="not connected"):
await mgr._refresh_server_tools("ghost")
asyncio.run(_run())
class TestLastRefreshTracking:
"""Phase 9 admin status pill — ``_last_refresh`` is written on every
refresh path so the admin UI reflects manual-refresh AND auto-
reconnect outcomes uniformly. This test class pins the contract.
"""
@staticmethod
def _seed_minimal(mgr: MCPClientManager, name: str = "srv") -> MagicMock:
mock_session = MagicMock()
mock_session.list_tools = AsyncMock(return_value=MagicMock(tools=[]))
mock_session.list_resources = AsyncMock(return_value=MagicMock(resources=[]))
mock_session.list_resource_templates = AsyncMock(
return_value=MagicMock(resourceTemplates=[])
)
mock_session.list_prompts = AsyncMock(return_value=MagicMock(prompts=[]))
_seed_static_state(
mgr,
name,
session=mock_session,
tools=[],
supports_resources=True,
supports_prompts=True,
)
return mock_session
def test_last_refresh_written_on_success(self) -> None:
async def _run() -> None:
mgr = MCPClientManager({})
self._seed_minimal(mgr)
assert "srv" not in mgr._last_refresh
await mgr._refresh_server("srv")
entry = mgr._last_refresh.get("srv")
assert entry is not None
ts, outcome = entry
assert outcome == "ok"
assert isinstance(ts, float) and ts > 0
asyncio.run(_run())
def test_last_refresh_written_on_tool_refresh_failure(self) -> None:
"""When ``_refresh_server_tools`` raises, the outcome reflects
the exception class and the exception still propagates."""
async def _run() -> None:
mgr = MCPClientManager({})
mock_session = self._seed_minimal(mgr)
mock_session.list_tools = AsyncMock(side_effect=RuntimeError("upstream down"))
with pytest.raises(RuntimeError, match="upstream down"):
await mgr._refresh_server("srv")
entry = mgr._last_refresh.get("srv")
assert entry is not None
_, outcome = entry
assert outcome == "error:RuntimeError"
asyncio.run(_run())
def test_last_refresh_records_first_exception_when_multiple_fail(
self,
) -> None:
"""``return_exceptions=True`` lets sibling tasks complete; the
outcome reflects the FIRST exception encountered."""
async def _run() -> None:
mgr = MCPClientManager({})
mock_session = self._seed_minimal(mgr)
# Tools succeeds; resources raises first (gather preserves
# argument order in its results list, so resources is the
# first failure regardless of which awaitable finished first
# in wall-clock terms).
mock_session.list_resources = AsyncMock(side_effect=ValueError("res boom"))
mock_session.list_prompts = AsyncMock(side_effect=KeyError("prompts boom"))
with pytest.raises((ValueError, KeyError)):
await mgr._refresh_server("srv")
entry = mgr._last_refresh.get("srv")
assert entry is not None
_, outcome = entry
# Either of the two failing tasks could be "first" in
# gather's results list ordering — the order is positional
# so resources (arg #2) comes before prompts (arg #3).
assert outcome == "error:ValueError"
asyncio.run(_run())
def test_refresh_all_overwrites_stale_ok_on_reconnect_failure(
self,
) -> None:
"""The chokepoint bug-1 fix: a prior successful refresh's ``'ok'``
entry MUST be overwritten when a subsequent reconnect fails —
otherwise the admin pill shows misleading "ok" while the server
is in fact broken."""
async def _run() -> None:
mgr = MCPClientManager({})
# Server is configured but has no live session — _refresh_all
# routes to the reconnect branch.
mgr._server_configs["srv"] = {"type": "stdio", "command": "x"}
# Pre-seed a stale "ok" from an earlier successful refresh.
mgr._last_refresh["srv"] = (1000.0, "ok")
async def _raise(*_a: object, **_kw: object) -> None:
raise ConnectionError("reconnect failed")
# The reconnect branch routes through _ensure_static_connected,
# which calls the LOCKED connect body.
mgr._connect_one_locked = _raise # type: ignore[assignment]
await mgr._refresh_all("srv")
entry = mgr._last_refresh.get("srv")
assert entry is not None
ts, outcome = entry
# Outcome reflects the new failure, not the stale ok.
assert outcome == "error:ConnectionError"
assert ts > 1000.0
asyncio.run(_run())
def test_get_server_status_surfaces_last_refresh_fields(self) -> None:
"""``get_server_status`` surfaces ``last_refresh_at`` and
``last_refresh_outcome`` for the admin pill — null when no
refresh has occurred yet, populated after one."""
mgr = MCPClientManager({})
mgr._server_configs["srv"] = {"type": "stdio", "command": "x"}
# No refresh yet — fields must be present and null so the JS
# renderer can branch on absence cleanly.
status = mgr.get_server_status("srv")
assert status["last_refresh_at"] is None
assert status["last_refresh_outcome"] is None
# Populate the tuple directly and re-read.
mgr._last_refresh["srv"] = (12345.5, "ok")
status = mgr.get_server_status("srv")
assert status["last_refresh_at"] == 12345.5
assert status["last_refresh_outcome"] == "ok"
class TestListeners:
def test_add_and_notify(self):
mgr = MCPClientManager({})
calls: list[int] = []
mgr.add_listener(lambda: calls.append(1))
_seed_static_state(mgr, "a", tools=[_fake_openai_tool("mcp__a__x")])
mgr._rebuild_tools()
assert len(calls) == 1
def test_remove_listener(self):
mgr = MCPClientManager({})
calls: list[int] = []
cb = lambda: calls.append(1) # noqa: E731
mgr.add_listener(cb)
mgr.remove_listener(cb)
mgr._rebuild_tools()
assert calls == []
def test_remove_nonexistent_listener(self):
mgr = MCPClientManager({})
mgr.remove_listener(lambda: None) # should not raise
def test_listener_error_does_not_propagate(self):
mgr = MCPClientManager({})
mgr.add_listener(lambda: 1 / 0) # will raise ZeroDivisionError
mgr._rebuild_tools() # should not raise
# -- Phase 7: user-keyed listener fan-out ------------------------------
def test_add_listener_records_user_id(self):
"""``add_listener`` stores ``(user_id, callback)`` tuples — the
listener identity carries the user_id."""
mgr = MCPClientManager({})
cb_admin = lambda: None # noqa: E731
cb_user = lambda: None # noqa: E731
mgr.add_listener(cb_admin) # default: user_id=None (admin)
mgr.add_listener(cb_user, user_id="user-1")
assert (None, cb_admin) in mgr._listeners
assert ("user-1", cb_user) in mgr._listeners
def test_remove_listener_requires_matching_user_id(self):
"""Removing with a different ``user_id`` must NOT remove the
original registration — listener identity is the pair."""
mgr = MCPClientManager({})
calls: list[int] = []
cb = lambda: calls.append(1) # noqa: E731
mgr.add_listener(cb, user_id="user-1")
# Try to remove with the wrong user_id — should be a no-op.
mgr.remove_listener(cb, user_id="user-2")
# The user-1 listener should still be live.
mgr._notify_user_tool_listeners("user-1")
assert calls == [1]
# Now remove with the right user_id.
mgr.remove_listener(cb, user_id="user-1")
mgr._notify_user_tool_listeners("user-1")
assert calls == [1] # not invoked again
def test_static_change_fires_all_listeners(self):
"""``_rebuild_tools`` (static-path change) fires ALL registered
listeners — admin + every user. RFC §3.3."""
mgr = MCPClientManager({})
admin_calls: list[int] = []
u1_calls: list[int] = []
u2_calls: list[int] = []
mgr.add_listener(lambda: admin_calls.append(1))
mgr.add_listener(lambda: u1_calls.append(1), user_id="user-1")
mgr.add_listener(lambda: u2_calls.append(1), user_id="user-2")
_seed_static_state(mgr, "a", tools=[_fake_openai_tool("mcp__a__x")])
mgr._rebuild_tools()
assert admin_calls == [1]
assert u1_calls == [1]
assert u2_calls == [1]
def test_user_tool_listeners_only_fire_for_matching_user(self):
"""``_notify_user_tool_listeners('user-1')`` fires admin (None)
and user-1 listeners; user-2's listener is silent."""
mgr = MCPClientManager({})
admin_calls: list[int] = []
u1_calls: list[int] = []
u2_calls: list[int] = []
mgr.add_listener(lambda: admin_calls.append(1))
mgr.add_listener(lambda: u1_calls.append(1), user_id="user-1")
mgr.add_listener(lambda: u2_calls.append(1), user_id="user-2")
mgr._notify_user_tool_listeners("user-1")
assert admin_calls == [1]
assert u1_calls == [1]
assert u2_calls == []
mgr._notify_user_tool_listeners("user-2")
assert admin_calls == [1, 1]
assert u1_calls == [1]
assert u2_calls == [1]
class TestServerNames:
def test_server_names_property(self):
mgr = MCPClientManager({"github": {}, "slack": {}})
assert sorted(mgr.server_names) == ["github", "slack"]
def test_server_names_empty(self):
mgr = MCPClientManager({})
assert mgr.server_names == []
# ---------------------------------------------------------------------------
# Session integration — tool refresh propagation
# ---------------------------------------------------------------------------
class TestSessionRefresh:
@pytest.fixture()
def tmp_db(self, tmp_path):
from turnstone.core.storage import init_storage, reset_storage
reset_storage()
init_storage("sqlite", path=str(tmp_path / "test.db"), run_migrations=False)
yield
reset_storage()
def _make_session(self, mcp_client=None, **kwargs):
from turnstone.core.session import ChatSession
defaults: dict[str, Any] = dict(
client=MagicMock(),
model="test-model",
ui=MagicMock(),
instructions=None,
temperature=0.5,
max_tokens=4096,
tool_timeout=30,
mcp_client=mcp_client,
)
defaults.update(kwargs)
return ChatSession(**defaults)
def test_listener_registered_on_init(self, tmp_db):
mock_mcp = MagicMock()
mock_mcp.get_tools.return_value = []
session = self._make_session(mcp_client=mock_mcp)
mock_mcp.add_listener.assert_called_once()
assert session._mcp_refresh_cb is not None
def test_no_listener_without_mcp(self, tmp_db):
session = self._make_session(mcp_client=None)
assert session._mcp_refresh_cb is None
def test_close_removes_listener(self, tmp_db):
mock_mcp = MagicMock()
mock_mcp.get_tools.return_value = []
session = self._make_session(mcp_client=mock_mcp)
session.close()
mock_mcp.remove_listener.assert_called_once()
assert session._mcp_refresh_cb is None
def test_close_idempotent(self, tmp_db):
mock_mcp = MagicMock()
mock_mcp.get_tools.return_value = []
session = self._make_session(mcp_client=mock_mcp)
session.close()
session.close() # should not raise
assert mock_mcp.remove_listener.call_count == 1
def test_on_mcp_tools_changed_rebuilds_tools(self, tmp_db):
mock_mcp = MagicMock()
mock_mcp.get_tools.return_value = [_fake_openai_tool("mcp__test__a")]
session = self._make_session(mcp_client=mock_mcp)
initial_count = len(session._tools)
# Simulate a tool refresh — MCP now has 2 tools
mock_mcp.get_tools.return_value = [
_fake_openai_tool("mcp__test__a"),
_fake_openai_tool("mcp__test__b"),
]
session._on_mcp_tools_changed()
assert len(session._tools) == initial_count + 1
def test_tool_search_preserved_across_refresh(self, tmp_db):
# Create enough MCP tools to trigger tool search
mcp_tools = [_fake_openai_tool(f"mcp__srv__tool{i}") for i in range(25)]
mock_mcp = MagicMock()
mock_mcp.get_tools.return_value = mcp_tools
session = self._make_session(
mcp_client=mock_mcp,
tool_search="auto",
tool_search_threshold=20,
)
assert session._tool_search is not None
# Expand a tool
session._tool_search.expand_visible(["mcp__srv__tool0"])
assert "mcp__srv__tool0" in session._tool_search.get_expanded_names()
# Refresh with same tools
session._on_mcp_tools_changed()
assert session._tool_search is not None
assert "mcp__srv__tool0" in session._tool_search.get_expanded_names()
def test_tool_search_prunes_removed_from_expanded(self, tmp_db):
mcp_tools = [_fake_openai_tool(f"mcp__srv__tool{i}") for i in range(25)]
mock_mcp = MagicMock()
mock_mcp.get_tools.return_value = mcp_tools
session = self._make_session(
mcp_client=mock_mcp,
tool_search="auto",
tool_search_threshold=20,
)
session._tool_search.expand_visible(["mcp__srv__tool0"])
# Refresh with tool0 removed
new_tools = [_fake_openai_tool(f"mcp__srv__tool{i}") for i in range(1, 25)]
mock_mcp.get_tools.return_value = new_tools
session._on_mcp_tools_changed()
# tool0 was removed, so it should no longer be in expanded
expanded = session._tool_search.get_expanded_names()
assert "mcp__srv__tool0" not in expanded
def test_mcp_refresh_command(self, tmp_db):
mock_mcp = MagicMock()
mock_mcp.get_tools.return_value = [_fake_openai_tool()]
mock_mcp.server_names = ["test"]
mock_mcp.refresh_sync.return_value = {"test": (["mcp__test__new"], [])}
session = self._make_session(mcp_client=mock_mcp)
session.handle_command("/mcp refresh")
mock_mcp.refresh_sync.assert_called_once_with(None)
session.ui.on_info.assert_called()
def test_mcp_refresh_specific_server(self, tmp_db):
mock_mcp = MagicMock()
mock_mcp.get_tools.return_value = [_fake_openai_tool()]
mock_mcp.server_names = ["github", "slack"]
mock_mcp.refresh_sync.return_value = {"github": ([], [])}
session = self._make_session(mcp_client=mock_mcp)
session.handle_command("/mcp refresh github")
mock_mcp.refresh_sync.assert_called_once_with("github")
def test_mcp_refresh_unknown_server(self, tmp_db):
mock_mcp = MagicMock()
mock_mcp.get_tools.return_value = [_fake_openai_tool()]
mock_mcp.server_names = ["github"]
session = self._make_session(mcp_client=mock_mcp)
session.handle_command("/mcp refresh nonexistent")
session.ui.on_error.assert_called_once()
assert "Unknown MCP server" in session.ui.on_error.call_args[0][0]
def test_mcp_refresh_error_handling(self, tmp_db):
mock_mcp = MagicMock()
mock_mcp.get_tools.return_value = [_fake_openai_tool()]
mock_mcp.server_names = ["test"]
mock_mcp.refresh_sync.side_effect = TimeoutError("timed out")
session = self._make_session(mcp_client=mock_mcp)
session.handle_command("/mcp refresh")
session.ui.on_error.assert_called_once()
assert "MCP refresh failed" in session.ui.on_error.call_args[0][0]
# ---------------------------------------------------------------------------
# MCP Resources
# ---------------------------------------------------------------------------
class TestMCPResources:
def test_resource_discovery(self):
"""Mock list_resources() returning 2 resources, verify get_resources()."""
mgr = MCPClientManager({})
_seed_static_state(
mgr,
"fs",
resources=[
_fake_resource_dict("file:///a.txt", "a", "File A", "text/plain", "fs"),
_fake_resource_dict("file:///b.txt", "b", "File B", "text/plain", "fs"),
],
)
mgr._rebuild_resources()
resources = mgr.get_resources()
assert len(resources) == 2
uris = {r["uri"] for r in resources}
assert uris == {"file:///a.txt", "file:///b.txt"}
assert all(r["server"] == "fs" for r in resources)
def test_rebuild_resources_copy_on_write(self):
"""Verify mutation safety — get_resources() returns independent copy."""
mgr = MCPClientManager({})
_seed_static_state(mgr, "a", resources=[_fake_resource_dict("file:///x", "x", "", "", "a")])
mgr._rebuild_resources()
old_resources = mgr._resources
old_map = mgr._resource_map
_seed_static_state(mgr, "b", resources=[_fake_resource_dict("file:///y", "y", "", "", "b")])
mgr._rebuild_resources()
assert mgr._resources is not old_resources
assert mgr._resource_map is not old_map
def test_get_resources_returns_copy(self):
mgr = MCPClientManager({})
_seed_static_state(mgr, "a", resources=[_fake_resource_dict("file:///x", "x", "", "", "a")])
mgr._rebuild_resources()
resources = mgr.get_resources()
assert len(resources) == 1
resources.clear()
assert len(mgr.get_resources()) == 1
def test_read_resource_sync(self):
"""Mock session.read_resource(), verify text extraction."""
mgr = MCPClientManager({})
mgr._resource_map = {"file:///readme": ("fs", "file:///readme")}
mock_session = MagicMock()
_seed_static_state(mgr, "fs", session=mock_session)
mgr._loop = asyncio.new_event_loop()
# Mock the read_resource result
text_content = MagicMock(spec=["text"])
text_content.text = "Hello, world!"
mock_result = MagicMock()
mock_result.contents = [text_content]
mock_session.read_resource = AsyncMock(return_value=mock_result)
thread = None
try:
thread = __import__("threading").Thread(target=mgr._loop.run_forever, daemon=True)
thread.start()
output = mgr.read_resource_sync("file:///readme", timeout=5)
assert output == "Hello, world!"
mock_session.read_resource.assert_awaited_once_with("file:///readme")
finally:
mgr._loop.call_soon_threadsafe(mgr._loop.stop)
if thread:
thread.join(timeout=5)
mgr._loop.close()
def test_read_resource_sync_blob(self):
"""Verify base64 blob extraction."""
mgr = MCPClientManager({})
mgr._resource_map = {"file:///img.png": ("fs", "file:///img.png")}
mock_session = MagicMock()
_seed_static_state(mgr, "fs", session=mock_session)
mgr._loop = asyncio.new_event_loop()
blob_content = MagicMock(spec=["blob"])
blob_content.blob = "aGVsbG8="
mock_result = MagicMock()
mock_result.contents = [blob_content]
mock_session.read_resource = AsyncMock(return_value=mock_result)
thread = None
try:
thread = __import__("threading").Thread(target=mgr._loop.run_forever, daemon=True)
thread.start()
output = mgr.read_resource_sync("file:///img.png", timeout=5)
assert output == "aGVsbG8="
finally:
mgr._loop.call_soon_threadsafe(mgr._loop.stop)
if thread:
thread.join(timeout=5)
mgr._loop.close()
def test_read_resource_sync_unknown_uri(self):
mgr = MCPClientManager({})
with pytest.raises(ValueError, match="Unknown MCP resource"):
mgr.read_resource_sync("file:///nonexistent")
def test_read_resource_sync_disconnected(self):
mgr = MCPClientManager({})
mgr._resource_map = {"file:///x": ("dead", "file:///x")}
with pytest.raises(RuntimeError, match="not connected"):
mgr.read_resource_sync("file:///x")
def test_read_resource_sync_timeout(self):
"""Verify timeout handling."""
mgr = MCPClientManager({})
mgr._resource_map = {"file:///x": ("fs", "file:///x")}
mock_session = MagicMock()
_seed_static_state(mgr, "fs", session=mock_session)
mgr._loop = asyncio.new_event_loop()
async def _slow_read(_uri: str) -> None:
await asyncio.sleep(10)
mock_session.read_resource = _slow_read
thread = None
try:
thread = __import__("threading").Thread(target=mgr._loop.run_forever, daemon=True)
thread.start()
with pytest.raises(TimeoutError):
mgr.read_resource_sync("file:///x", timeout=1)
finally:
mgr._loop.call_soon_threadsafe(mgr._loop.stop)
if thread:
thread.join(timeout=5)
mgr._loop.close()
def test_resource_listener_notification(self):
"""Verify callback fires on rebuild."""
mgr = MCPClientManager({})
calls: list[int] = []
mgr.add_resource_listener(lambda: calls.append(1))
_seed_static_state(mgr, "a", resources=[_fake_resource_dict()])
mgr._rebuild_resources()
assert len(calls) == 1
def test_resource_listener_remove(self):
mgr = MCPClientManager({})
calls: list[int] = []
cb = lambda: calls.append(1) # noqa: E731
mgr.add_resource_listener(cb)
mgr.remove_resource_listener(cb)
mgr._rebuild_resources()
assert calls == []
def test_resource_listener_error_does_not_propagate(self):
mgr = MCPClientManager({})
mgr.add_resource_listener(lambda: 1 / 0)
mgr._rebuild_resources() # should not raise
def test_resource_refresh_on_notification(self):
"""Mock notification, verify re-fetch of resources."""
async def _run() -> None:
mgr = MCPClientManager({})
mock_session = MagicMock()
_seed_static_state(
mgr,
"fs",
session=mock_session,
supports_resources=True,
resources=[_fake_resource_dict("file:///old", server="fs")],
)
mgr._rebuild_resources()
assert len(mgr.get_resources()) == 1
# Mock the re-fetch returning a new resource
new_res = _fake_mcp_resource("file:///new", "new")
mock_res_result = MagicMock()
mock_res_result.resources = [new_res]
mock_session.list_resources = AsyncMock(return_value=mock_res_result)
mock_tmpl_result = MagicMock()
mock_tmpl_result.resourceTemplates = []
mock_session.list_resource_templates = AsyncMock(return_value=mock_tmpl_result)
await mgr._refresh_server_resources("fs")
resources = mgr.get_resources()
assert len(resources) == 1
assert resources[0]["uri"] == "file:///new"
asyncio.run(_run())
def test_rebuild_resources_empty(self):
mgr = MCPClientManager({})
mgr._static_servers = {}
mgr._rebuild_resources()
assert mgr._resources == []
assert mgr._resource_map == {}
def test_rebuild_resources_multi_server(self):
mgr = MCPClientManager({})
_seed_static_state(mgr, "fs", resources=[_fake_resource_dict("file:///a", server="fs")])
_seed_static_state(
mgr,
"db",
resources=[_fake_resource_dict("db://table", name="table", server="db")],
)
mgr._rebuild_resources()
assert len(mgr._resources) == 2
assert mgr._resource_map["file:///a"] == ("fs", "file:///a")
assert mgr._resource_map["db://table"] == ("db", "db://table")
def test_template_prefix_matching(self):
"""Expanded URI matches template by prefix."""
mgr = MCPClientManager({})
_seed_static_state(
mgr,
"db",
resources=[
{
"uri": "db://tables/{table}/rows/{id}",
"name": "row",
"description": "A row",
"mimeType": "application/json",
"server": "db",
"template": True,
},
],
)
mgr._rebuild_resources()
# Template should not be in resource_map
assert "db://tables/{table}/rows/{id}" not in mgr._resource_map
# But prefix matching should find it
result = mgr._match_template("db://tables/users/rows/1")
assert result is not None
server, template_uri = result
assert server == "db"
assert template_uri == "db://tables/{table}/rows/{id}"
def test_template_longest_prefix_wins(self):
"""When two templates have overlapping prefixes, the longer one wins."""
mgr = MCPClientManager({})
# Use templates with genuinely different prefix lengths:
# "db://data/" (6 chars after scheme) vs "db://data/tables/" (13 chars after scheme)
_seed_static_state(
mgr,
"short",
resources=[
{
"uri": "db://data/{collection}",
"name": "collection",
"description": "",
"mimeType": "",
"server": "short",
"template": True,
},
],
)
_seed_static_state(
mgr,
"long",
resources=[
{
"uri": "db://data/tables/{table}",
"name": "table",
"description": "",
"mimeType": "",
"server": "long",
"template": True,
},
],
)
mgr._rebuild_resources()
# "db://data/tables/users" matches both prefixes ("db://data/" and
# "db://data/tables/") — the longer one should win
result = mgr._match_template("db://data/tables/users")
assert result is not None
server, template_uri = result
assert server == "long"
assert template_uri == "db://data/tables/{table}"
# URI that only matches the short prefix
result2 = mgr._match_template("db://data/views/active")
assert result2 is not None
assert result2[0] == "short"
def test_template_no_match_raises(self):
"""Completely unrelated URI still raises ValueError."""
mgr = MCPClientManager({})
_seed_static_state(
mgr,
"db",
resources=[
{
"uri": "db://tables/{table}",
"name": "table",
"description": "",
"mimeType": "",
"server": "db",
"template": True,
},
],
)
mgr._rebuild_resources()
assert mgr._match_template("file:///something") is None
with pytest.raises(ValueError, match="Unknown MCP resource"):
mgr.read_resource_sync("file:///something")
def test_read_resource_sync_with_template_uri(self):
"""End-to-end: template discovered, expanded URI dispatched to correct server."""
mgr = MCPClientManager({})
_seed_static_state(
mgr,
"db",
resources=[
{
"uri": "db://tables/{table}/rows/{id}",
"name": "row",
"description": "A row",
"mimeType": "application/json",
"server": "db",
"template": True,
},
],
)
mgr._rebuild_resources()
mock_session = MagicMock()
_seed_static_state(mgr, "db", session=mock_session)
mgr._loop = asyncio.new_event_loop()
text_content = MagicMock(spec=["text"])
text_content.text = '{"name": "Alice"}'
mock_result = MagicMock()
mock_result.contents = [text_content]
mock_session.read_resource = AsyncMock(return_value=mock_result)
thread = None
try:
thread = __import__("threading").Thread(target=mgr._loop.run_forever, daemon=True)
thread.start()
output = mgr.read_resource_sync("db://tables/users/rows/1", timeout=5)
assert output == '{"name": "Alice"}'
mock_session.read_resource.assert_awaited_once_with("db://tables/users/rows/1")
finally:
mgr._loop.call_soon_threadsafe(mgr._loop.stop)
if thread:
thread.join(timeout=5)
mgr._loop.close()
# ---------------------------------------------------------------------------
# MCP Prompts
# ---------------------------------------------------------------------------
class TestMCPPrompts:
def test_prompt_discovery(self):
"""Mock list_prompts(), verify get_prompts() with correct prefixed names."""
mgr = MCPClientManager({})
_seed_static_state(
mgr,
"tmpl",
prompts=[
_fake_prompt_dict("mcp__tmpl__code_review", "code_review", "tmpl"),
_fake_prompt_dict("mcp__tmpl__summarize", "summarize", "tmpl"),
],
)
mgr._rebuild_prompts()
prompts = mgr.get_prompts()
assert len(prompts) == 2
names = {p["name"] for p in prompts}
assert names == {"mcp__tmpl__code_review", "mcp__tmpl__summarize"}
# Verify map entries
assert mgr._prompt_map["mcp__tmpl__code_review"] == ("tmpl", "code_review")
assert mgr._prompt_map["mcp__tmpl__summarize"] == ("tmpl", "summarize")
def test_rebuild_prompts_copy_on_write(self):
"""Verify mutation safety."""
mgr = MCPClientManager({})
_seed_static_state(mgr, "a", prompts=[_fake_prompt_dict("mcp__a__p1", "p1", "a")])
mgr._rebuild_prompts()
old_prompts = mgr._prompts
old_map = mgr._prompt_map
_seed_static_state(mgr, "b", prompts=[_fake_prompt_dict("mcp__b__p2", "p2", "b")])
mgr._rebuild_prompts()
assert mgr._prompts is not old_prompts
assert mgr._prompt_map is not old_map
def test_get_prompts_returns_copy(self):
mgr = MCPClientManager({})
_seed_static_state(mgr, "a", prompts=[_fake_prompt_dict("mcp__a__p1", "p1", "a")])
mgr._rebuild_prompts()
prompts = mgr.get_prompts()
assert len(prompts) == 1
prompts.clear()
assert len(mgr.get_prompts()) == 1
def test_get_prompt_sync(self):
"""Mock session.get_prompt(), verify message conversion."""
mgr = MCPClientManager({})
mgr._prompt_map = {"mcp__tmpl__review": ("tmpl", "review")}
mock_session = MagicMock()
_seed_static_state(mgr, "tmpl", session=mock_session)
mgr._loop = asyncio.new_event_loop()
# Build mock PromptMessage
msg1 = MagicMock()
msg1.role = "user"
msg1.content = MagicMock()
msg1.content.text = "Review this code"
msg2 = MagicMock()
msg2.role = "assistant"
msg2.content = MagicMock()
msg2.content.text = "Looks good!"
mock_result = MagicMock()
mock_result.messages = [msg1, msg2]
mock_session.get_prompt = AsyncMock(return_value=mock_result)
thread = None
try:
thread = __import__("threading").Thread(target=mgr._loop.run_forever, daemon=True)
thread.start()
messages = mgr.get_prompt_sync(
"mcp__tmpl__review", arguments={"language": "python"}, timeout=5
)
assert len(messages) == 2
assert messages[0] == {"role": "user", "content": "Review this code"}
assert messages[1] == {"role": "assistant", "content": "Looks good!"}
mock_session.get_prompt.assert_awaited_once_with(
"review", arguments={"language": "python"}
)
finally:
mgr._loop.call_soon_threadsafe(mgr._loop.stop)
if thread:
thread.join(timeout=5)
mgr._loop.close()
def test_get_prompt_sync_unknown(self):
mgr = MCPClientManager({})
with pytest.raises(ValueError, match="Unknown MCP prompt"):
mgr.get_prompt_sync("mcp__no__such")
def test_get_prompt_sync_disconnected(self):
mgr = MCPClientManager({})
mgr._prompt_map = {"mcp__dead__p": ("dead", "p")}
with pytest.raises(RuntimeError, match="not connected"):
mgr.get_prompt_sync("mcp__dead__p")
def test_get_prompt_sync_timeout(self):
"""Verify timeout handling."""
mgr = MCPClientManager({})
mgr._prompt_map = {"mcp__tmpl__slow": ("tmpl", "slow")}
mock_session = MagicMock()
_seed_static_state(mgr, "tmpl", session=mock_session)
mgr._loop = asyncio.new_event_loop()
async def _slow_prompt(_name: str, *, arguments: dict[str, str] | None = None) -> None:
await asyncio.sleep(10)
mock_session.get_prompt = _slow_prompt
thread = None
try:
thread = __import__("threading").Thread(target=mgr._loop.run_forever, daemon=True)
thread.start()
with pytest.raises(TimeoutError):
mgr.get_prompt_sync("mcp__tmpl__slow", timeout=1)
finally:
mgr._loop.call_soon_threadsafe(mgr._loop.stop)
if thread:
thread.join(timeout=5)
mgr._loop.close()
def test_prompt_listener_notification(self):
"""Verify callback fires on rebuild."""
mgr = MCPClientManager({})
calls: list[int] = []
mgr.add_prompt_listener(lambda: calls.append(1))
_seed_static_state(mgr, "a", prompts=[_fake_prompt_dict()])
mgr._rebuild_prompts()
assert len(calls) == 1
def test_prompt_listener_remove(self):
mgr = MCPClientManager({})
calls: list[int] = []
cb = lambda: calls.append(1) # noqa: E731
mgr.add_prompt_listener(cb)
mgr.remove_prompt_listener(cb)
mgr._rebuild_prompts()
assert calls == []
def test_prompt_listener_error_does_not_propagate(self):
mgr = MCPClientManager({})
mgr.add_prompt_listener(lambda: 1 / 0)
mgr._rebuild_prompts() # should not raise
def test_is_mcp_prompt(self):
"""Verify name lookup."""
mgr = MCPClientManager({})
mgr._prompt_map["mcp__tmpl__review"] = ("tmpl", "review")
assert mgr.is_mcp_prompt("mcp__tmpl__review") is True
assert mgr.is_mcp_prompt("nonexistent") is False
def test_prompt_refresh_on_notification(self):
"""Mock notification, verify re-fetch of prompts."""
async def _run() -> None:
mgr = MCPClientManager({})
mock_session = MagicMock()
_seed_static_state(
mgr,
"tmpl",
session=mock_session,
supports_prompts=True,
prompts=[_fake_prompt_dict("mcp__tmpl__old", "old", "tmpl")],
)
mgr._rebuild_prompts()
assert len(mgr.get_prompts()) == 1
# Mock re-fetch returning a new prompt
new_prompt = _fake_mcp_prompt("new_prompt", "A new prompt")
mock_prompt_result = MagicMock()
mock_prompt_result.prompts = [new_prompt]
mock_session.list_prompts = AsyncMock(return_value=mock_prompt_result)
await mgr._refresh_server_prompts("tmpl")
prompts = mgr.get_prompts()
assert len(prompts) == 1
assert prompts[0]["name"] == "mcp__tmpl__new_prompt"
assert prompts[0]["original_name"] == "new_prompt"
asyncio.run(_run())
def test_rebuild_prompts_empty(self):
mgr = MCPClientManager({})
mgr._static_servers = {}
mgr._rebuild_prompts()
assert mgr._prompts == []
assert mgr._prompt_map == {}
def test_rebuild_prompts_multi_server(self):
mgr = MCPClientManager({})
_seed_static_state(mgr, "a", prompts=[_fake_prompt_dict("mcp__a__p1", "p1", "a")])
_seed_static_state(mgr, "b", prompts=[_fake_prompt_dict("mcp__b__p2", "p2", "b")])
mgr._rebuild_prompts()
assert len(mgr._prompts) == 2
assert mgr._prompt_map["mcp__a__p1"] == ("a", "p1")
assert mgr._prompt_map["mcp__b__p2"] == ("b", "p2")
# ---------------------------------------------------------------------------
# Shutdown cleans up new state
# ---------------------------------------------------------------------------
class TestShutdownCleanup:
def test_shutdown_clears_resources_and_prompts(self):
mgr = MCPClientManager({})
_seed_static_state(mgr, "a", resources=[_fake_resource_dict()])
mgr._rebuild_resources()
_seed_static_state(mgr, "a", prompts=[_fake_prompt_dict()])
mgr._rebuild_prompts()
assert mgr.get_resources() != []
assert mgr.get_prompts() != []
mgr.shutdown()
assert mgr.get_resources() == []
assert mgr.get_prompts() == []
assert mgr._resource_map == {}
assert mgr._prompt_map == {}
def test_shutdown_closes_owned_loop_and_clears_refs(self):
"""When the manager owns the loop thread, shutdown must close the loop
(selector resources leak otherwise) and drop both refs; a second
shutdown is then a clean no-op."""
import threading as _threading
mgr = MCPClientManager({})
loop = asyncio.new_event_loop()
thread = _threading.Thread(target=loop.run_forever, daemon=True)
thread.start()
mgr._loop = loop
mgr._thread = thread
mgr.shutdown()
assert loop.is_closed()
assert mgr._loop is None
assert mgr._thread is None
mgr.shutdown() # idempotent
def test_shutdown_leaves_unowned_loop_open(self):
"""Tests (and any embedder) that wire ``_loop`` directly without a
thread own the loop's lifecycle — shutdown must not close it."""
mgr = MCPClientManager({})
loop = asyncio.new_event_loop()
mgr._loop = loop
try:
mgr.shutdown()
assert not loop.is_closed()
finally:
loop.close()
# ---------------------------------------------------------------------------
# TCP probe and unreachable server handling
# ---------------------------------------------------------------------------
class TestTCPProbe:
"""MCPClientManager._tcp_probe should fail fast on unreachable servers."""
def test_tcp_probe_unreachable_raises_connection_error(self):
"""Unreachable host raises ConnectionError, not TimeoutError."""
mgr = MCPClientManager({})
async def _run():
with pytest.raises(ConnectionError, match="unreachable"):
await mgr._tcp_probe("test-server", "http://127.0.0.1:1")
asyncio.run(_run())
def test_tcp_probe_parses_url_correctly(self):
"""Port and host are extracted from the URL."""
mgr = MCPClientManager({})
async def _run():
# Non-routable port — should fail with ConnectionError
with pytest.raises(ConnectionError):
await mgr._tcp_probe("srv", "https://127.0.0.1:1/mcp")
asyncio.run(_run())
def test_tcp_probe_default_port_http(self):
"""Default port 80 used for http:// URLs without explicit port."""
mgr = MCPClientManager({})
async def _run():
# Will fail (nothing on port 80), but should not crash on parsing
with pytest.raises(ConnectionError):
await mgr._tcp_probe("srv", "http://127.0.0.1")
asyncio.run(_run())
def test_tcp_probe_dns_failure(self):
"""Unresolvable hostname raises ConnectionError."""
mgr = MCPClientManager({})
async def _run():
with pytest.raises(ConnectionError):
await mgr._tcp_probe("srv", "http://this.host.does.not.exist.invalid:8080/mcp")
asyncio.run(_run())
class TestConnectOneUnreachable:
"""_connect_one should handle unreachable HTTP servers gracefully."""
def test_unreachable_http_server_raises_connection_error(self):
"""Unreachable HTTP MCP server raises ConnectionError without spinning."""
mgr = MCPClientManager({})
mgr._loop = asyncio.new_event_loop()
async def _run():
with pytest.raises(ConnectionError, match="unreachable"):
await mgr._connect_one(
"bad-server",
{
"type": "http",
"url": "http://127.0.0.1:1/mcp",
},
)
mgr._loop.run_until_complete(_run())
mgr._loop.close()
# Server should NOT have a live session (connection failed)
bad_state = mgr._static_servers.get("bad-server")
assert bad_state is None or bad_state.session is None
def test_connect_all_continues_after_unreachable_server(self):
"""_connect_all logs error and continues to next server."""
mgr = MCPClientManager(
{
"bad": {"type": "http", "url": "http://127.0.0.1:1/mcp"},
}
)
loop = asyncio.new_event_loop()
loop.run_until_complete(mgr._connect_all())
# _connect_all spawns the long-lived token-freshness sweep task; cancel
# it before closing the loop so it isn't destroyed while pending.
sweep = mgr._user_token_sweep_task
if sweep is not None:
sweep.cancel()
with contextlib.suppress(BaseException):
loop.run_until_complete(sweep)
# _connect_all spawns the long-lived static health task; cancel it before
# closing the loop so it isn't destroyed while pending.
health = mgr._static_health_task
if health is not None:
health.cancel()
with suppress(BaseException):
loop.run_until_complete(health)
loop.close()
bad_state = mgr._static_servers.get("bad")
assert bad_state is None or bad_state.session is None
assert "bad" in mgr._last_error
# ---------------------------------------------------------------------------
# Fix 1: Cancel orphaned futures on timeout
# ---------------------------------------------------------------------------
class TestFutureCancellation:
"""Verify future.cancel() is called when sync bridge methods time out."""
def _make_manager_with_session(self) -> MCPClientManager:
mgr = MCPClientManager({"test": {"type": "stdio", "command": "echo"}})
mock_session = MagicMock()
# Prevent auto-spec from creating async coroutines that trigger warnings
mock_session.call_tool = MagicMock(return_value="sentinel")
mock_session.read_resource = MagicMock(return_value="sentinel")
mock_session.get_prompt = MagicMock(return_value="sentinel")
_seed_static_state(mgr, "test", session=mock_session)
mgr._loop = MagicMock()
mgr._tool_map["mcp__test__search"] = ("test", "search")
mgr._resource_map["file:///a.txt"] = ("test", "file:///a.txt")
mgr._prompt_map["mcp__test__review"] = ("test", "review")
return mgr
def test_call_tool_sync_cancels_future_on_timeout(self):
mgr = self._make_manager_with_session()
mock_future = MagicMock()
mock_future.result.side_effect = concurrent.futures.TimeoutError()
with (
patch("asyncio.run_coroutine_threadsafe", new=_dispatch_stub(mock_future)),
pytest.raises(TimeoutError, match="timed out"),
):
mgr.call_tool_sync("mcp__test__search", {"query": "x"}, timeout=1)
mock_future.cancel.assert_called_once()
def test_read_resource_sync_cancels_future_on_timeout(self):
mgr = self._make_manager_with_session()
mock_future = MagicMock()
mock_future.result.side_effect = concurrent.futures.TimeoutError()
with (
patch("asyncio.run_coroutine_threadsafe", new=_dispatch_stub(mock_future)),
pytest.raises(TimeoutError, match="timed out"),
):
mgr.read_resource_sync("file:///a.txt", timeout=1)
mock_future.cancel.assert_called_once()
def test_get_prompt_sync_cancels_future_on_timeout(self):
mgr = self._make_manager_with_session()
mock_future = MagicMock()
mock_future.result.side_effect = concurrent.futures.TimeoutError()
with (
patch("asyncio.run_coroutine_threadsafe", new=_dispatch_stub(mock_future)),
pytest.raises(TimeoutError, match="timed out"),
):
mgr.get_prompt_sync("mcp__test__review", timeout=1)
mock_future.cancel.assert_called_once()
def test_refresh_sync_cancels_future_on_timeout(self):
mgr = MCPClientManager({})
mgr._loop = MagicMock()
mock_future = MagicMock()
mock_future.result.side_effect = concurrent.futures.TimeoutError()
with (
patch.object(mgr, "_refresh_all", return_value=MagicMock()),
patch("asyncio.run_coroutine_threadsafe", new=_dispatch_stub(mock_future)),
pytest.raises(TimeoutError, match="timed out"),
):
mgr.refresh_sync(timeout=1)
mock_future.cancel.assert_called_once()
# ---------------------------------------------------------------------------
# Fix 2: Per-server circuit breaker
# ---------------------------------------------------------------------------
class TestCircuitBreaker:
"""Verify per-server circuit breaker behavior."""
def test_circuit_stays_closed_below_threshold(self):
mgr = MCPClientManager({})
mgr._cb_record_failure("srv")
mgr._cb_record_failure("srv")
is_open, _ = mgr._cb_check("srv")
assert not is_open
def test_circuit_opens_at_threshold(self):
mgr = MCPClientManager({})
for _ in range(3):
mgr._cb_record_failure("srv")
is_open, cooldown_expired = mgr._cb_check("srv")
assert is_open
assert not cooldown_expired # just opened, cooldown not expired
def test_circuit_half_open_after_cooldown(self):
mgr = MCPClientManager({})
for _ in range(3):
mgr._cb_record_failure("srv")
# Simulate cooldown expiry
mgr._circuit_open_until["srv"] = time.monotonic() - 1
is_open, cooldown_expired = mgr._cb_check("srv")
assert is_open
assert cooldown_expired
def test_circuit_resets_on_success(self):
mgr = MCPClientManager({})
for _ in range(3):
mgr._cb_record_failure("srv")
assert "srv" in mgr._circuit_open_until
mgr._cb_record_success("srv")
is_open, _ = mgr._cb_check("srv")
assert not is_open
assert mgr._consecutive_failures.get("srv") is None
def test_success_decays_trip_count(self):
"""Success decays trip_count by 1 so flapping servers escalate backoff."""
mgr = MCPClientManager({})
mgr._circuit_trip_count["srv"] = 3
mgr._cb_record_success("srv")
assert mgr._circuit_trip_count["srv"] == 2
mgr._cb_record_success("srv")
assert mgr._circuit_trip_count["srv"] == 1
mgr._cb_record_success("srv")
assert "srv" not in mgr._circuit_trip_count
def test_cooldown_is_exponential(self):
mgr = MCPClientManager({})
# First trip (trip_count starts at 0)
for _ in range(3):
mgr._cb_record_failure("srv")
deadline1 = mgr._circuit_open_until["srv"]
base1 = deadline1 - time.monotonic()
# Reset circuit but keep trip_count at 1 (set by first trip)
mgr._cb_record_success("srv")
# trip_count decayed from 1 to 0 — manually set to 1 for test
mgr._circuit_trip_count["srv"] = 1
for _ in range(3):
mgr._cb_record_failure("srv")
deadline2 = mgr._circuit_open_until["srv"]
base2 = deadline2 - time.monotonic()
# Second trip should have longer cooldown (roughly 2x, within jitter)
assert base2 > base1 * 1.5
def test_cooldown_capped_at_max(self):
mgr = MCPClientManager({})
mgr._circuit_trip_count["srv"] = 100 # very high trip count
for _ in range(3):
mgr._cb_record_failure("srv")
deadline = mgr._circuit_open_until["srv"]
cooldown = deadline - time.monotonic()
# Should not exceed max (300s) + 10% jitter = 330s
assert cooldown <= mgr._CB_MAX_COOLDOWN * 1.11
def test_cb_gate_rejects_when_open(self):
mgr = MCPClientManager({})
for _ in range(3):
mgr._cb_record_failure("srv")
with pytest.raises(RuntimeError, match="circuit open"):
mgr._cb_gate("srv")
def test_cb_gate_allows_after_cooldown(self):
mgr = MCPClientManager({})
for _ in range(3):
mgr._cb_record_failure("srv")
mgr._circuit_open_until["srv"] = time.monotonic() - 1
# Should not raise
mgr._cb_gate("srv")
# Deadline should be removed (half-open probe allowed)
assert "srv" not in mgr._circuit_open_until
def test_cb_clear_removes_all_state(self):
mgr = MCPClientManager({})
for _ in range(3):
mgr._cb_record_failure("srv")
mgr._cb_clear("srv")
assert "srv" not in mgr._consecutive_failures
assert "srv" not in mgr._circuit_open_until
assert "srv" not in mgr._circuit_trip_count
def test_call_tool_sync_records_failure_on_timeout(self):
mgr = MCPClientManager({"test": {"type": "stdio", "command": "echo"}})
mock_session = MagicMock()
mock_session.call_tool = MagicMock(return_value="sentinel")
_seed_static_state(mgr, "test", session=mock_session)
mgr._loop = MagicMock()
mgr._tool_map["mcp__test__ping"] = ("test", "ping")
mock_future = MagicMock()
mock_future.result.side_effect = concurrent.futures.TimeoutError()
with (
patch("asyncio.run_coroutine_threadsafe", new=_dispatch_stub(mock_future)),
pytest.raises(TimeoutError),
):
mgr.call_tool_sync("mcp__test__ping", {}, timeout=1)
assert mgr._consecutive_failures.get("test", 0) == 1
def test_call_tool_sync_records_success(self):
mgr = MCPClientManager({"test": {"type": "stdio", "command": "echo"}})
mock_session = MagicMock()
mock_session.call_tool = MagicMock(return_value="sentinel")
_seed_static_state(mgr, "test", session=mock_session)
mgr._loop = MagicMock()
mgr._tool_map["mcp__test__ping"] = ("test", "ping")
# Pre-set a failure
mgr._consecutive_failures["test"] = 2
mock_result = MagicMock()
mock_result.content = []
mock_result.isError = False
mock_future = MagicMock()
mock_future.result.return_value = mock_result
with patch("asyncio.run_coroutine_threadsafe", new=_dispatch_stub(mock_future)):
mgr.call_tool_sync("mcp__test__ping", {}, timeout=5)
assert mgr._consecutive_failures.get("test") is None
def test_connection_error_evicts_session(self):
mgr = MCPClientManager({"test": {"type": "stdio", "command": "echo"}})
mock_session = MagicMock()
mock_session.call_tool = MagicMock(return_value="sentinel")
# Seed session + owner so the test can verify the owner survives.
old_owner = MagicMock()
old_streams = (MagicMock(), MagicMock())
_seed_static_state(
mgr, "test", session=mock_session, owner_task=old_owner, streams=old_streams
)
mgr._loop = MagicMock()
mgr._tool_map["mcp__test__ping"] = ("test", "ping")
mock_future = MagicMock()
mock_future.result.side_effect = BrokenPipeError("dead")
with (
patch("asyncio.run_coroutine_threadsafe", new=_dispatch_stub(mock_future)),
pytest.raises(BrokenPipeError),
):
mgr.call_tool_sync("mcp__test__ping", {}, timeout=5)
# Session evicted, but the owner/streams remain for the stale guard in
# _connect_one_locked to close on the next reconnect attempt.
state = mgr._static_servers["test"]
assert state.session is None
assert state.owner_task is old_owner
assert state.streams is old_streams
def test_independent_circuits_per_server(self):
mgr = MCPClientManager({})
for _ in range(3):
mgr._cb_record_failure("a")
is_open_a, _ = mgr._cb_check("a")
is_open_b, _ = mgr._cb_check("b")
assert is_open_a
assert not is_open_b
def test_mcp_error_does_not_trip_circuit(self):
"""Protocol errors (McpError) should not count as transport failures."""
from mcp import McpError
from mcp.types import ErrorData
mgr = MCPClientManager({"test": {"type": "stdio", "command": "echo"}})
mock_session = MagicMock()
mock_session.call_tool = MagicMock(return_value="sentinel")
_seed_static_state(mgr, "test", session=mock_session)
mgr._loop = MagicMock()
mgr._tool_map["mcp__test__ping"] = ("test", "ping")
mock_future = MagicMock()
mock_future.result.side_effect = McpError(ErrorData(code=-32601, message="tool not found"))
with (
patch("asyncio.run_coroutine_threadsafe", new=_dispatch_stub(mock_future)),
pytest.raises(McpError),
):
mgr.call_tool_sync("mcp__test__ping", {}, timeout=5)
# Circuit should NOT have recorded a failure
assert mgr._consecutive_failures.get("test", 0) == 0
def test_closed_resource_error_evicts_session_and_trips_circuit(self):
"""Regression: the MCP SDK's streamable-http transport raises
``anyio.ClosedResourceError`` (NOT BrokenPipeError) when its write
stream is dead. That must evict the session AND trip the breaker —
otherwise the corpse session is re-used on every call forever and
only a full process restart recovers it."""
import anyio
mgr = MCPClientManager({"test": {"type": "stdio", "command": "echo"}})
mock_session = MagicMock()
mock_session.call_tool = MagicMock(return_value="sentinel")
_seed_static_state(mgr, "test", session=mock_session)
mgr._loop = MagicMock()
mgr._tool_map["mcp__test__ping"] = ("test", "ping")
mock_future = MagicMock()
mock_future.result.side_effect = anyio.ClosedResourceError()
with (
patch("asyncio.run_coroutine_threadsafe", new=_dispatch_stub(mock_future)),
pytest.raises(anyio.ClosedResourceError),
):
mgr.call_tool_sync("mcp__test__ping", {}, timeout=5)
assert mgr._static_servers["test"].session is None
assert mgr._consecutive_failures.get("test", 0) == 1
def test_session_terminated_mcperror_evicts_and_trips_circuit(self):
"""Regression: when the MCP SERVER restarts and loses its session map, our
held mcp-session-id is stale; the server returns HTTP 404 and the SDK
surfaces McpError(code=32600, 'Session terminated'). That is NOT a healthy
protocol rejection — the session must be evicted so the next dispatch
reconnects with a fresh initialize; reusing it 404s forever (restart-hang)."""
from mcp import McpError
from mcp.types import ErrorData
mgr = MCPClientManager({"test": {"type": "stdio", "command": "echo"}})
mock_session = MagicMock()
mock_session.call_tool = MagicMock(return_value="sentinel")
_seed_static_state(mgr, "test", session=mock_session)
mgr._loop = MagicMock()
mgr._tool_map["mcp__test__ping"] = ("test", "ping")
mock_future = MagicMock()
# Exactly what the streamable-http SDK injects on a 404 stale session.
mock_future.result.side_effect = McpError(
ErrorData(code=32600, message="Session terminated")
)
with (
patch("asyncio.run_coroutine_threadsafe", new=_dispatch_stub(mock_future)),
pytest.raises(McpError),
):
mgr.call_tool_sync("mcp__test__ping", {}, timeout=5)
assert mgr._static_servers["test"].session is None
assert mgr._consecutive_failures.get("test", 0) == 1
def test_httpx_connect_error_evicts_session(self):
"""A dead underlying httpx connection (server down mid-call) is transport
death, not a protocol rejection — evict so the next call reconnects."""
import httpx
mgr = MCPClientManager({"test": {"type": "stdio", "command": "echo"}})
mock_session = MagicMock()
mock_session.call_tool = MagicMock(return_value="sentinel")
_seed_static_state(mgr, "test", session=mock_session)
mgr._loop = MagicMock()
mgr._tool_map["mcp__test__ping"] = ("test", "ping")
mock_future = MagicMock()
mock_future.result.side_effect = httpx.ConnectError("connection refused")
with (
patch("asyncio.run_coroutine_threadsafe", new=_dispatch_stub(mock_future)),
pytest.raises(httpx.ConnectError),
):
mgr.call_tool_sync("mcp__test__ping", {}, timeout=5)
assert mgr._static_servers["test"].session is None
assert mgr._consecutive_failures.get("test", 0) == 1
def test_connection_closed_mcperror_evicts_and_trips_circuit(self):
"""Regression: when the SDK's ``post_writer`` swallows the transport
error, a dead connection surfaces as ``McpError(CONNECTION_CLOSED)``.
Unlike a genuine protocol rejection, this MUST evict + trip the
breaker so the next dispatch reconnects instead of looping."""
from mcp import McpError
from mcp.types import CONNECTION_CLOSED, ErrorData
mgr = MCPClientManager({"test": {"type": "stdio", "command": "echo"}})
mock_session = MagicMock()
mock_session.call_tool = MagicMock(return_value="sentinel")
_seed_static_state(mgr, "test", session=mock_session)
mgr._loop = MagicMock()
mgr._tool_map["mcp__test__ping"] = ("test", "ping")
mock_future = MagicMock()
mock_future.result.side_effect = McpError(
ErrorData(code=CONNECTION_CLOSED, message="connection closed")
)
with (
patch("asyncio.run_coroutine_threadsafe", new=_dispatch_stub(mock_future)),
pytest.raises(McpError),
):
mgr.call_tool_sync("mcp__test__ping", {}, timeout=5)
assert mgr._static_servers["test"].session is None
assert mgr._consecutive_failures.get("test", 0) == 1
def test_refresh_all_evicts_dead_session_so_next_tick_reconnects(self):
"""Regression: a periodic refresh that hits a dead-but-non-None
session must null the session so the reconnect branch (gated on
``session is None``) fires on the NEXT tick. Without this the
refresh re-probes the corpse forever — the bug that required a
full restart."""
import anyio
async def _run() -> None:
mgr = MCPClientManager({})
mgr._server_configs["test"] = {"type": "stdio", "command": "x"}
dead = anyio.ClosedResourceError()
mock_session = MagicMock()
mock_session.list_tools = AsyncMock(side_effect=dead)
mock_session.list_resources = AsyncMock(side_effect=dead)
mock_session.list_resource_templates = AsyncMock(side_effect=dead)
mock_session.list_prompts = AsyncMock(side_effect=dead)
_seed_static_state(mgr, "test", session=mock_session)
await mgr._refresh_all("test")
# Dead session evicted → next refresh tick / dispatch reconnects.
assert mgr._static_servers["test"].session is None
ts, outcome = mgr._last_refresh["test"]
assert outcome == "error:ClosedResourceError"
asyncio.run(_run())
def test_read_resource_sync_dead_transport_evicts_and_trips_circuit(self):
"""Regression (follow-up): read_resource_sync kept the old
BrokenPipe/ConnectionReset/EOF-only guard, so a dead streamable-http
transport surfacing as McpError(CONNECTION_CLOSED) reused the corpse
session forever — the exact restart-hang call_tool_sync already fixes.
It must now evict the session AND trip the breaker."""
from mcp import McpError
from mcp.types import CONNECTION_CLOSED, ErrorData
mgr = MCPClientManager({"test": {"type": "stdio", "command": "echo"}})
mock_session = MagicMock()
_seed_static_state(mgr, "test", session=mock_session)
mgr._loop = MagicMock()
mgr._resource_map = {"file:///x": ("test", "file:///x")}
mock_future = MagicMock()
mock_future.result.side_effect = McpError(
ErrorData(code=CONNECTION_CLOSED, message="connection closed")
)
with (
patch("asyncio.run_coroutine_threadsafe", new=_dispatch_stub(mock_future)),
pytest.raises(McpError),
):
mgr.read_resource_sync("file:///x", timeout=5)
assert mgr._static_servers["test"].session is None
assert mgr._consecutive_failures.get("test", 0) == 1
def test_read_resource_sync_protocol_mcperror_does_not_evict(self):
"""A healthy protocol rejection (resource not found) must NOT evict the
session or trip the breaker on the resource path."""
from mcp import McpError
from mcp.types import ErrorData
mgr = MCPClientManager({"test": {"type": "stdio", "command": "echo"}})
mock_session = MagicMock()
_seed_static_state(mgr, "test", session=mock_session)
mgr._loop = MagicMock()
mgr._resource_map = {"file:///x": ("test", "file:///x")}
mock_future = MagicMock()
mock_future.result.side_effect = McpError(
ErrorData(code=-32602, message="resource not found")
)
with (
patch("asyncio.run_coroutine_threadsafe", new=_dispatch_stub(mock_future)),
pytest.raises(McpError),
):
mgr.read_resource_sync("file:///x", timeout=5)
assert mgr._static_servers["test"].session is mock_session
assert mgr._consecutive_failures.get("test", 0) == 0
def test_get_prompt_sync_dead_transport_evicts_and_trips_circuit(self):
"""Regression (follow-up): get_prompt_sync had the same corpse-reuse
bug as read_resource_sync. A dead transport (anyio.ClosedResourceError)
must evict the session AND trip the breaker."""
import anyio
mgr = MCPClientManager({"test": {"type": "stdio", "command": "echo"}})
mock_session = MagicMock()
_seed_static_state(mgr, "test", session=mock_session)
mgr._loop = MagicMock()
mgr._prompt_map = {"mcp__test__p": ("test", "p")}
mock_future = MagicMock()
mock_future.result.side_effect = anyio.ClosedResourceError()
with (
patch("asyncio.run_coroutine_threadsafe", new=_dispatch_stub(mock_future)),
pytest.raises(anyio.ClosedResourceError),
):
mgr.get_prompt_sync("mcp__test__p", timeout=5)
assert mgr._static_servers["test"].session is None
assert mgr._consecutive_failures.get("test", 0) == 1
def test_get_prompt_sync_protocol_mcperror_does_not_evict(self):
"""A healthy protocol rejection must NOT evict on the prompt path."""
from mcp import McpError
from mcp.types import ErrorData
mgr = MCPClientManager({"test": {"type": "stdio", "command": "echo"}})
mock_session = MagicMock()
_seed_static_state(mgr, "test", session=mock_session)
mgr._loop = MagicMock()
mgr._prompt_map = {"mcp__test__p": ("test", "p")}
mock_future = MagicMock()
mock_future.result.side_effect = McpError(
ErrorData(code=-32602, message="prompt not found")
)
with (
patch("asyncio.run_coroutine_threadsafe", new=_dispatch_stub(mock_future)),
pytest.raises(McpError),
):
mgr.get_prompt_sync("mcp__test__p", timeout=5)
assert mgr._static_servers["test"].session is mock_session
assert mgr._consecutive_failures.get("test", 0) == 0
class TestIsDeadTransport:
"""Direct unit tests for ``_is_dead_transport`` — the single shared gate
that decides 'tear down and rebuild the session' vs 'healthy protocol
rejection' across every session-use site."""
def test_connection_closed_is_dead(self):
from mcp import McpError
from mcp.types import CONNECTION_CLOSED, ErrorData
assert _is_dead_transport(
McpError(ErrorData(code=CONNECTION_CLOSED, message="connection closed"))
)
def test_sdk_session_terminated_is_dead(self):
"""The streamable-http SDK synthesizes EXACTLY code=32600 /
'Session terminated' when a held mcp-session-id 404s after a server
restart — keyed off the code so it survives a message reword."""
from mcp import McpError
from mcp.types import ErrorData
assert _is_dead_transport(McpError(ErrorData(code=32600, message="Session terminated")))
def test_app_session_not_found_is_not_dead(self):
"""#2 regression: a HEALTHY session-owning MCP server (game/shell)
rejecting a stale id with 'session not found' is a protocol error, NOT
transport death. The old bare-substring match wrongly evicted the live
session and tripped the shared breaker for every user."""
from mcp import McpError
from mcp.types import ErrorData
assert not _is_dead_transport(
McpError(ErrorData(code=-32603, message="Backend session not found"))
)
def test_app_session_terminated_message_is_not_dead(self):
"""#8 regression: the message is application-controlled and is NOT matched
— only the SDK's synthesized code 32600 is. A healthy session-owning
server that returns a protocol error whose message is EXACTLY 'Session
terminated' (or a superstring) with a normal code stays breaker-safe."""
from mcp import McpError
from mcp.types import ErrorData
# Exact SDK message but an app protocol code (not 32600) — must NOT be dead.
assert not _is_dead_transport(
McpError(ErrorData(code=-32603, message="Session terminated"))
)
# Superstring likewise.
assert not _is_dead_transport(
McpError(ErrorData(code=-32603, message="Player session terminated by host"))
)
def test_plain_protocol_mcperror_is_not_dead(self):
from mcp import McpError
from mcp.types import ErrorData
assert not _is_dead_transport(McpError(ErrorData(code=-32601, message="method not found")))
def test_httpx_read_timeout_is_dead(self):
"""#7: an idle read timeout on a long-lived streamable-http stream is
the dominant idle-death mode — and is NOT a builtin TimeoutError, so it
must be caught here or it falls through to a healthy 'other'."""
import httpx
assert not issubclass(httpx.ReadTimeout, TimeoutError) # premise guard
assert _is_dead_transport(httpx.ReadTimeout("read timed out"))
def test_httpx_pool_timeout_is_not_dead(self):
"""PoolTimeout is connection-pool saturation, NOT a dead connection:
evicting the session can't relieve pool pressure and would trip the
shared breaker for all users under transient load. The Connect/Read/Write
timeouts (a dead/hung connection) stay dead."""
import httpx
assert not _is_dead_transport(httpx.PoolTimeout("pool exhausted"))
assert _is_dead_transport(httpx.WriteTimeout("write timed out"))
def test_httpx_read_error_is_dead(self):
"""#8: a connection that dies mid-read surfaces as httpx.ReadError (a
NetworkError sibling of the already-handled ConnectError)."""
import httpx
assert _is_dead_transport(httpx.ReadError("peer reset"))
def test_httpx_write_error_is_dead(self):
import httpx
assert _is_dead_transport(httpx.WriteError("broken pipe"))
def test_httpx_local_protocol_error_is_not_dead(self):
"""LocalProtocolError is OUR bug (a malformed request we built), not a
dead peer — it must NOT be mistaken for transport death."""
import httpx
assert not _is_dead_transport(httpx.LocalProtocolError("bad header"))
def test_anyio_closed_resource_is_dead(self):
import anyio
assert _is_dead_transport(anyio.ClosedResourceError())
# ---------------------------------------------------------------------------
# Fix 3: Safe transport stream pre-close
# ---------------------------------------------------------------------------
class TestSafeTransportStreams:
"""Verify stream references are stored and pre-closed."""
def test_pre_close_streams_closes_both(self):
mgr = MCPClientManager({})
stream_a = MagicMock()
stream_b = MagicMock()
_seed_static_state(mgr, "srv", streams=(stream_a, stream_b))
async def _run():
await mgr._pre_close_streams("srv")
asyncio.run(_run())
stream_a.aclose.assert_called_once()
stream_b.aclose.assert_called_once()
# Streams cleared, but the state entry itself can remain.
assert mgr._static_servers["srv"].streams is None
def test_pre_close_streams_ignores_missing(self):
mgr = MCPClientManager({})
async def _run():
await mgr._pre_close_streams("nonexistent")
asyncio.run(_run()) # should not raise
def test_pre_close_streams_suppresses_errors(self):
mgr = MCPClientManager({})
stream_a = MagicMock()
stream_a.aclose.side_effect = RuntimeError("boom")
stream_b = MagicMock()
_seed_static_state(mgr, "srv", streams=(stream_a, stream_b))
async def _run():
await mgr._pre_close_streams("srv")
asyncio.run(_run()) # should not raise despite stream_a error
stream_b.aclose.assert_called_once()
def test_shutdown_clears_stream_refs(self):
mgr = MCPClientManager({})
_seed_static_state(mgr, "srv", streams=(MagicMock(), MagicMock()))
mgr.shutdown()
assert len(mgr._static_servers) == 0
# ---------------------------------------------------------------------------
# Fix 4: Notification debounce
# ---------------------------------------------------------------------------
class TestNotificationDebounce:
"""Verify notification-triggered refreshes are debounced."""
def test_debounce_within_window(self):
mgr = MCPClientManager({})
mgr._last_notification_refresh["srv"] = time.monotonic()
# We can't easily call _on_notification (it's a closure), so test
# the debounce logic directly via the timestamp check
now = time.monotonic()
last = mgr._last_notification_refresh.get("srv", 0.0)
assert now - last < mgr._NOTIFICATION_DEBOUNCE
def test_debounce_passes_after_window(self):
mgr = MCPClientManager({})
# Set timestamp well in the past
mgr._last_notification_refresh["srv"] = time.monotonic() - 10
now = time.monotonic()
last = mgr._last_notification_refresh.get("srv", 0.0)
assert now - last >= mgr._NOTIFICATION_DEBOUNCE
def test_debounce_is_per_server(self):
mgr = MCPClientManager({})
mgr._last_notification_refresh["srv_a"] = time.monotonic()
# srv_b has no timestamp — should pass debounce
now = time.monotonic()
last_b = mgr._last_notification_refresh.get("srv_b", 0.0)
assert now - last_b >= mgr._NOTIFICATION_DEBOUNCE
# ---------------------------------------------------------------------------
# reconnect_sync — operator-driven full reconnect
# ---------------------------------------------------------------------------
class TestReconnectSync:
"""Verify reconnect_sync tears down old session, clears CB, calls _connect_one."""
def test_reconnect_unknown_server_returns_error(self):
mgr = MCPClientManager({})
result = mgr.reconnect_sync("missing")
assert result == {
"connected": False,
"tools": 0,
"resources": 0,
"prompts": 0,
"error": "unknown server",
}
def test_reconnect_clears_circuit_breaker(self, running_loop_mgr):
mgr, _loop, _thread = running_loop_mgr
async def _fake_connect_one(name: str, _cfg: dict[str, Any]) -> None:
_seed_static_state(mgr, name, session=MagicMock())
# Pre-trip the breaker
for _ in range(3):
mgr._cb_record_failure("srv")
assert "srv" in mgr._circuit_open_until
with (
# reconnect_sync holds the per-name lock and calls the LOCKED body.
patch.object(mgr, "_connect_one_locked", side_effect=_fake_connect_one),
patch.object(mgr, "_pre_close_streams", new=AsyncMock()),
):
result = mgr.reconnect_sync("srv")
assert result["connected"] is True
assert result["error"] == ""
assert "srv" not in mgr._circuit_open_until
assert "srv" not in mgr._consecutive_failures
def test_reconnect_closes_old_session_then_calls_connect_one(self, running_loop_mgr):
"""The FORCE-rebuild teardown lives in ``_connect_one_locked``'s
stale-guard (the one canonical ``_teardown_static_session`` sequence) —
``reconnect_sync`` no longer carries its own copy. Drive the REAL locked
body via a no-command stdio cfg: the stale-guard runs, then the connect
early-returns, so the ordering is observable without a live server."""
mgr, loop, _thread = running_loop_mgr
mgr._server_configs["srv"] = {"type": "stdio"} # no command → early return
order: list[str] = []
async def _make_owner() -> tuple[asyncio.Event, asyncio.Task[None]]:
ev = asyncio.Event()
async def _parked_owner() -> None:
await ev.wait()
order.append("owner_exit")
task = asyncio.create_task(_parked_owner())
await asyncio.sleep(0)
return ev, task
ev, old_owner = _run_hl(loop, _make_owner())
async def _pre_close(name: str) -> None:
order.append("pre_close")
# Session must already be nulled when streams close (canonical order).
assert mgr._static_servers["srv"].session is None
# Seed the old session/owner/streams that the stale-guard should close.
_seed_static_state(
mgr,
"srv",
session=MagicMock(),
owner_task=old_owner,
close_requested=ev,
streams=(MagicMock(), MagicMock()),
)
with patch.object(mgr, "_pre_close_streams", side_effect=_pre_close):
result = mgr.reconnect_sync("srv")
assert order == ["pre_close", "owner_exit"] # teardown ran, in order
assert result["connected"] is False # no command — nothing to rebuild
state = mgr._static_servers["srv"]
assert state.session is None
assert state.owner_task is None # old owner cleared from state
assert old_owner.done() and not old_owner.cancelled()
def test_reconnect_failure_returns_error_dict(self, running_loop_mgr):
mgr, _loop, _thread = running_loop_mgr
async def _connect_one_locked(name: str, _cfg: dict[str, Any]) -> None:
raise RuntimeError("handshake failed")
with (
patch.object(mgr, "_connect_one_locked", side_effect=_connect_one_locked),
patch.object(mgr, "_pre_close_streams", new=AsyncMock()),
):
result = mgr.reconnect_sync("srv")
assert result["connected"] is False
assert "handshake failed" in result["error"]
def test_reconnect_failure_clears_stale_catalog(self, running_loop_mgr):
# bug-2: when _connect_one fails mid-reconnect, the per-server
# catalog must be dropped so the merged tool/resource/prompt maps
# don't keep advertising entries with no live session.
mgr, _loop, _thread = running_loop_mgr
# Seed catalog state from a previous successful connect.
_seed_static_state(
mgr,
"srv",
tools=[_fake_openai_tool("mcp__srv__t")],
resources=[_fake_resource_dict(server="srv")],
prompts=[_fake_prompt_dict(server="srv")],
)
mgr._rebuild_tools()
mgr._rebuild_resources()
mgr._rebuild_prompts()
async def _connect_one_locked(name: str, _cfg: dict[str, Any]) -> None:
raise RuntimeError("handshake failed")
with (
patch.object(mgr, "_connect_one_locked", side_effect=_connect_one_locked),
patch.object(mgr, "_pre_close_streams", new=AsyncMock()),
):
result = mgr.reconnect_sync("srv")
assert result["connected"] is False
# Per-server catalog should be cleared and merged maps drained.
srv_state = mgr._static_servers.get("srv")
assert srv_state is not None
assert srv_state.tools == []
assert srv_state.resources == []
assert srv_state.prompts == []
assert "mcp__srv__t" not in mgr._tool_map
def test_reconnect_preserves_static_state_identity(self, running_loop_mgr):
# q-3: PR #296 invariant 5 — _static_servers[name] must be the SAME
# object across a connect → transient-failure → reconnect cycle.
# Guards against future refactors that pop-and-repopulate the entry,
# which would invalidate any references held by concurrent readers.
mgr, _loop, _thread = running_loop_mgr
# First connect: seed an initial entry as if _connect_one succeeded.
async def _first_connect(name: str, _cfg: dict[str, Any]) -> None:
_seed_static_state(mgr, name, session=MagicMock())
with (
patch.object(mgr, "_connect_one_locked", side_effect=_first_connect),
patch.object(mgr, "_pre_close_streams", new=AsyncMock()),
):
mgr.reconnect_sync("srv")
state_before = mgr._static_servers["srv"]
id_before = id(state_before)
# Simulate a transient transport failure: evict the session (as
# call_tool_sync would on BrokenPipeError) but keep the entry.
state_before.session = None
# Reconnect.
async def _reconnect(name: str, _cfg: dict[str, Any]) -> None:
_seed_static_state(mgr, name, session=MagicMock())
with (
patch.object(mgr, "_connect_one_locked", side_effect=_reconnect),
patch.object(mgr, "_pre_close_streams", new=AsyncMock()),
):
result = mgr.reconnect_sync("srv")
assert result["connected"] is True
state_after = mgr._static_servers["srv"]
assert id(state_after) == id_before
assert state_after is state_before
# ---------------------------------------------------------------------------
# _cb_auto_reconnect — refresh-on-reconnect
# ---------------------------------------------------------------------------
class TestCBAutoReconnectRefresh:
"""Verify _cb_auto_reconnect schedules catalog refresh after a successful reconnect.
The refresh runs as a fire-and-forget background task on the loop so it
doesn't block the caller (perf-1). Tests wait briefly for the scheduled
task to run and observe its effect.
"""
def test_auto_reconnect_schedules_refresh_server_on_success(self, running_loop_mgr):
import threading as _threading
mgr, _loop, _thread = running_loop_mgr
new_session = MagicMock()
refresh_event = _threading.Event()
async def _connect_one(name: str, _cfg: dict[str, Any]) -> None:
_seed_static_state(mgr, name, session=new_session)
async def _refresh(name: str) -> tuple[list[str], list[str]]:
refresh_event.set()
return [], []
with (
# _cb_auto_reconnect coordinates via the per-name lock and calls the
# LOCKED connect body directly.
patch.object(mgr, "_connect_one_locked", side_effect=_connect_one),
patch.object(mgr, "_refresh_server", side_effect=_refresh),
):
session = mgr._cb_auto_reconnect("srv")
# Wait for the scheduled refresh task to actually run on the loop,
# then for the tracked task to DRAIN — exiting the patch context
# while the task is still in flight would hand the un-patched
# method to its tail.
assert refresh_event.wait(timeout=5), "refresh task was not scheduled"
deadline = time.time() + 5
while mgr._background_tasks and time.time() < deadline:
time.sleep(0.02)
assert not mgr._background_tasks, "background refresh task never drained"
assert session is new_session
def test_auto_reconnect_retrieves_and_logs_refresh_failure(self, running_loop_mgr):
"""A refresh failure must be RETRIEVED and logged by the task's
done-callback — not abandoned for asyncio to report as "Task exception
was never retrieved" at GC time (which lands on whatever stream pytest
has attached by then: the closed-file CI spew)."""
import threading as _threading
mgr, _loop, _thread = running_loop_mgr
new_session = MagicMock()
refresh_started = _threading.Event()
async def _connect_one(name: str, _cfg: dict[str, Any]) -> None:
_seed_static_state(mgr, name, session=new_session)
async def _refresh_failing(name: str) -> tuple[list[str], list[str]]:
refresh_started.set()
raise RuntimeError("catalog fetch broke")
with (
patch.object(mgr, "_connect_one_locked", side_effect=_connect_one),
patch.object(mgr, "_refresh_server", side_effect=_refresh_failing),
patch("turnstone.core.mcp_client.log") as mock_log,
):
# Must not raise — refresh failures are non-fatal to the caller.
session = mgr._cb_auto_reconnect("srv")
assert refresh_started.wait(timeout=5), "refresh task was not scheduled"
# Poll for the WARNING while the patch is still active — gating on
# set-emptiness alone would race the un-patch (review-caught: the
# warning could land on the restored real logger).
deadline = time.time() + 5
warn_calls = []
while not warn_calls and time.time() < deadline:
warn_calls = [
c for c in mock_log.warning.call_args_list if "MCP background" in str(c.args[0])
]
time.sleep(0.02)
# The tracked task must also fully drain (emptiness now implies
# "done AND reported" — discard is the callback's LAST step).
deadline = time.time() + 5
while mgr._background_tasks and time.time() < deadline:
time.sleep(0.02)
assert not mgr._background_tasks, "background refresh task never drained"
assert session is new_session
assert warn_calls, (
"the refresh failure must be logged by the done-callback, not left "
"for GC-time reporting"
)
exc = warn_calls[0].kwargs.get("exc_info")
assert isinstance(exc, RuntimeError)
assert "catalog fetch broke" in str(exc)
def _run_hl(loop: asyncio.AbstractEventLoop, coro: Any) -> Any:
"""Submit *coro* to the running fixture loop and wait (5s)."""
return asyncio.run_coroutine_threadsafe(coro, loop).result(timeout=5)
class TestStaticHealthLoop:
"""Autonomous static-server reconnect + liveness.
The SDK gives up reconnecting after 2 attempts on one transport and never on
the others (verified: mcp 1.28.1), and every other Turnstone reconnect path
is dispatch- or operator-driven — so a static server that dies while idle
stays dead. This loop is the missing autonomous trigger: it reconnects
disconnected servers on a capped, jittered, forever backoff and pings
connected ones, evicting a dead-but-idle session (which nothing else would
notice) so it is reconnected.
"""
# -- backoff policy ------------------------------------------------------
def test_reconnect_delay_capped_and_jittered(self) -> None:
mgr = MCPClientManager({})
# attempt 0: ceiling = base (1s); every draw within [0, base].
assert all(
0.0 <= mgr._static_reconnect_delay(0) <= mgr._STATIC_RECONNECT_BASE_S
for _ in range(200)
)
# large / unbounded attempt: never exceeds the cap (no overflow, forever).
assert all(
0.0 <= mgr._static_reconnect_delay(a) <= mgr._STATIC_RECONNECT_MAX_S
for a in (10, 100, 10_000)
)
# jitter actually varies (not a fixed value).
assert len({round(mgr._static_reconnect_delay(8), 6) for _ in range(50)}) > 1
# -- reconnect (layer 1) -------------------------------------------------
def test_reconnect_one_success_closes_open_circuit_but_keeps_failure_count(
self, running_loop_mgr
) -> None:
mgr, loop, _ = running_loop_mgr
mgr._server_configs["down"] = {"type": "stdio", "command": "echo"}
_seed_static_state(mgr, "down", session=None)
mgr._static_reconnect_attempt["down"] = 4 # pretend we'd been failing
# Trip the breaker OPEN (3 failures) so there is a deadline to clear.
for _ in range(3):
mgr._cb_record_failure("down")
assert "down" in mgr._circuit_open_until
async def _fake_connect(name: str, _cfg: dict[str, Any]) -> None:
_seed_static_state(mgr, name, session=MagicMock())
with (
# The driver routes through _ensure_static_connected, which calls
# the LOCKED connect body under the per-name lock.
patch.object(mgr, "_connect_one_locked", side_effect=_fake_connect),
patch.object(mgr, "_refresh_server", new=AsyncMock()),
):
_run_hl(loop, mgr._static_reconnect_one("down"))
assert mgr._static_servers["down"].session is not None
assert "down" not in mgr._static_reconnect_attempt # health backoff reset
# A transport reconnect closes the OPEN CIRCUIT (dispatch can flow) ...
assert "down" not in mgr._circuit_open_until
# ... but leaves the failure COUNT for a real dispatch to confirm/reset,
# so a connect-ok / calls-fail server still escalates to a trip.
assert mgr._consecutive_failures.get("down", 0) >= 1
def test_reconnect_one_retries_forever_with_growing_backoff(self, running_loop_mgr) -> None:
mgr, loop, _ = running_loop_mgr
mgr._server_configs["down"] = {"type": "stdio", "command": "echo"}
_seed_static_state(mgr, "down", session=None)
async def _boom(name: str, _cfg: dict[str, Any]) -> None:
raise RuntimeError("still down")
with patch.object(mgr, "_connect_one_locked", side_effect=_boom):
for expected in range(1, 7):
mgr._static_reconnect_next.pop("down", None) # force it due
due = _run_hl(loop, mgr._static_reconnect_one("down"))
assert mgr._static_reconnect_attempt["down"] == expected # no cap
assert due > time.monotonic() - 1 # next attempt scheduled
def test_reconnect_one_queued_behind_connect_reuses_its_session(self, running_loop_mgr) -> None:
"""While another driver holds the per-name connect lock, a health
reconnect QUEUES on it (no ``.locked()`` skip anymore) and then REUSES
the session the holder installed — never a second ``_connect_one_locked``
pile-on tearing down the fresh session."""
mgr, loop, _ = running_loop_mgr
mgr._server_configs["down"] = {"type": "stdio", "command": "echo"}
_seed_static_state(mgr, "down", session=None)
sess = MagicMock()
async def _scenario() -> float:
lock = mgr._static_connect_lock_for("down")
await lock.acquire() # a dispatch reconnect in progress
recon = asyncio.ensure_future(mgr._static_reconnect_one("down"))
await asyncio.sleep(0.05)
assert not recon.done() # queued on the lock, not skipped/failed
mgr._static_servers["down"].session = sess # the holder's connect lands
lock.release()
return await asyncio.wait_for(recon, timeout=5)
with (
patch.object(mgr, "_connect_one_locked", new=AsyncMock()) as cm,
patch.object(mgr, "_refresh_server", new=AsyncMock()),
):
due = _run_hl(loop, _scenario())
cm.assert_not_awaited() # reused the holder's session; no second connect
assert mgr._static_servers["down"].session is sess # never torn down
assert due > time.monotonic() - 1 # success cadence scheduled
def test_connect_one_serializes_concurrent_reconnects(self, running_loop_mgr) -> None:
"""The per-name lock prevents two concurrent ``_connect_one`` for one
server from interleaving teardown/rebuild on the shared state."""
mgr, loop, _ = running_loop_mgr
active = 0
max_active = 0
async def _inner(name: str, _cfg: dict[str, Any]) -> None:
nonlocal active, max_active
active += 1
max_active = max(max_active, active)
await asyncio.sleep(0.02)
active -= 1
async def _two() -> None:
with patch.object(mgr, "_connect_one_locked", side_effect=_inner):
await asyncio.gather(mgr._connect_one("x", {}), mgr._connect_one("x", {}))
_run_hl(loop, _two())
assert max_active == 1 # serialized, never overlapping
# -- liveness ping (layer 2 — the idle-dead detection) -------------------
def test_ping_one_healthy_keeps_session(self, running_loop_mgr) -> None:
mgr, loop, _ = running_loop_mgr
sess = MagicMock()
sess.send_ping = AsyncMock()
_seed_static_state(mgr, "up", session=sess)
due = _run_hl(loop, mgr._static_ping_one("up", time.monotonic()))
sess.send_ping.assert_awaited_once()
assert mgr._static_servers["up"].session is sess # kept
assert due > time.monotonic() # next ping scheduled
def test_ping_one_dead_session_evicts_for_reconnect(self, running_loop_mgr) -> None:
"""The Turnstone case: a connected-but-dead session that nothing else
would notice is detected by the ping and evicted so the next tick
reconnects it."""
mgr, loop, _ = running_loop_mgr
sess = MagicMock()
sess.send_ping = AsyncMock(side_effect=ConnectionResetError("peer gone"))
_seed_static_state(mgr, "up", session=sess)
now = time.monotonic()
_run_hl(loop, mgr._static_ping_one("up", now))
assert mgr._static_servers["up"].session is None # evicted
# "Reconnect asap": the (fresh-clock) deadline is already due.
assert now <= mgr._static_reconnect_next["up"] <= time.monotonic()
assert "up" in mgr._consecutive_failures # breaker recorded a failure
def test_ping_one_timeout_does_not_evict(self, running_loop_mgr) -> None:
"""A ping TIMEOUT means 'slow', not 'dead': the session is kept and the
ping rescheduled. A strict ping timeout must not churn a heavy-but-
working server every cycle (only a ``_is_dead_transport`` failure evicts).
"""
mgr, loop, _ = running_loop_mgr
mgr._STATIC_HEALTH_PING_TIMEOUT_S = 0.05 # instance shadow for a fast test
async def _hang() -> None:
await asyncio.sleep(10)
sess = MagicMock()
sess.send_ping = _hang
_seed_static_state(mgr, "up", session=sess)
now = time.monotonic()
due = _run_hl(loop, mgr._static_ping_one("up", now))
assert mgr._static_servers["up"].session is sess # kept — timeout != dead
assert "up" not in mgr._consecutive_failures # breaker NOT tripped
assert due > now # next ping rescheduled, not "reconnect asap"
def test_ping_one_mcp_error_does_not_evict(self, running_loop_mgr) -> None:
"""A protocol ``McpError`` from a healthy connection (server gates/omits
``ping``) must NOT evict or trip the breaker — only a dead transport does
(mirrors ``_record_and_evict_on_dead_transport``)."""
from mcp import McpError
from mcp.types import ErrorData
mgr, loop, _ = running_loop_mgr
sess = MagicMock()
sess.send_ping = AsyncMock(
side_effect=McpError(ErrorData(code=-32601, message="method not found"))
)
_seed_static_state(mgr, "up", session=sess)
now = time.monotonic()
due = _run_hl(loop, mgr._static_ping_one("up", now))
assert mgr._static_servers["up"].session is sess # kept — protocol != dead
assert "up" not in mgr._consecutive_failures # breaker untouched
assert due > now # rescheduled, not "reconnect asap"
def test_ping_one_pool_timeout_does_not_evict(self, running_loop_mgr) -> None:
"""``httpx.PoolTimeout`` is pool saturation, not a dead session — evicting
wouldn't relieve it and would trip the shared breaker under load; the ping
rescheduled instead (``_is_dead_transport`` deliberately excludes it)."""
import httpx
mgr, loop, _ = running_loop_mgr
sess = MagicMock()
sess.send_ping = AsyncMock(side_effect=httpx.PoolTimeout("pool saturated"))
_seed_static_state(mgr, "up", session=sess)
now = time.monotonic()
due = _run_hl(loop, mgr._static_ping_one("up", now))
assert mgr._static_servers["up"].session is sess # kept — pool != dead
assert "up" not in mgr._consecutive_failures
assert due > now
def test_ping_one_skips_busy_server(self, running_loop_mgr) -> None:
"""A server with an in-flight dispatch (``in_flight`` > 0) is demonstrably
alive; the ping is skipped and it is NEVER evicted — the interlock that
keeps a long-running ``call_tool`` from being torn down under itself."""
mgr, loop, _ = running_loop_mgr
sess = MagicMock()
sess.send_ping = AsyncMock(side_effect=AssertionError("busy server pinged"))
state = _seed_static_state(mgr, "up", session=sess)
state.in_flight = 1 # a call_tool is in flight
now = time.monotonic()
due = _run_hl(loop, mgr._static_ping_one("up", now))
sess.send_ping.assert_not_awaited() # skipped entirely
assert mgr._static_servers["up"].session is sess # not evicted
assert "up" not in mgr._consecutive_failures
assert due > now # rescheduled
def test_ping_one_does_not_clobber_concurrent_reconnect(self, running_loop_mgr) -> None:
"""A reconnect that installs a FRESH session during the ping's await
window must not be undone: the failure handler only nulls the session it
actually pinged (session-identity check)."""
mgr, loop, _ = running_loop_mgr
fresh = MagicMock() # S2, installed mid-ping by a concurrent reconnect
async def _die_after_swap() -> None:
# Model the race: a concurrent reconnect swaps in a fresh session,
# THEN this stale (S1) ping fails with a dead transport.
mgr._static_servers["up"].session = fresh
raise ConnectionResetError("S1 transport died")
stale = MagicMock() # S1
stale.send_ping = _die_after_swap
_seed_static_state(mgr, "up", session=stale)
_run_hl(loop, mgr._static_ping_one("up", time.monotonic()))
# S1's failure handler must NOT have nulled the freshly-installed S2.
assert mgr._static_servers["up"].session is fresh
# -- in-flight interlock (dispatch side) --------------------------------
def test_session_op_increments_in_flight_around_op(self, running_loop_mgr) -> None:
mgr, loop, _ = running_loop_mgr
state = _seed_static_state(mgr, "srv", session=MagicMock())
seen: list[int] = []
async def _op() -> str:
seen.append(state.in_flight) # observed WHILE in flight
return "ok"
assert _run_hl(loop, mgr._static_session_op("srv", _op())) == "ok"
assert seen == [1] # bumped for the duration of the op
assert state.in_flight == 0 # decremented in finally
def test_session_op_decrements_in_flight_on_error(self, running_loop_mgr) -> None:
mgr, loop, _ = running_loop_mgr
state = _seed_static_state(mgr, "srv", session=MagicMock())
async def _boom() -> None:
raise RuntimeError("dead transport")
with pytest.raises(RuntimeError, match="dead transport"):
_run_hl(loop, mgr._static_session_op("srv", _boom()))
assert state.in_flight == 0 # finally ran even on failure
def test_ping_skips_while_session_op_in_flight(self, running_loop_mgr) -> None:
"""End-to-end interlock: while a ``_static_session_op`` is mid-flight the
concurrent liveness ping skips and leaves the session intact."""
mgr, loop, _ = running_loop_mgr
sess = MagicMock()
sess.send_ping = AsyncMock(side_effect=AssertionError("busy server pinged"))
_seed_static_state(mgr, "up", session=sess)
async def _scenario() -> None:
release = asyncio.Event()
async def _op() -> str:
await release.wait()
return "done"
op_task = asyncio.ensure_future(mgr._static_session_op("up", _op()))
await asyncio.sleep(0) # let the op start and bump in_flight
assert mgr._static_servers["up"].in_flight == 1
due = await mgr._static_ping_one("up", time.monotonic())
assert mgr._static_servers["up"].session is sess # not evicted
assert due > time.monotonic() # rescheduled
release.set()
assert await op_task == "done"
assert mgr._static_servers["up"].in_flight == 0 # decremented
_run_hl(loop, _scenario())
sess.send_ping.assert_not_awaited()
# -- loop robustness (bounded / concurrent / fresh clock) ---------------
def test_reconnect_one_bounded_when_connect_wedges(self, running_loop_mgr) -> None:
"""A connect that handshakes then stalls its (unbounded) discovery must
not wedge the single loop coroutine or hold the per-name lock forever: the
caller-side timeout bounds it, it counts as a failed attempt (backoff),
and the ``async with`` lock is released on the cancellation."""
mgr, loop, _ = running_loop_mgr
mgr._STATIC_RECONNECT_ATTEMPT_TIMEOUT_S = 0.05 # shrink for a fast test
mgr._server_configs["wedge"] = {"type": "stdio", "command": "echo"}
_seed_static_state(mgr, "wedge", session=None)
async def _wedge_locked(name: str, _cfg: dict[str, Any]) -> None:
await asyncio.sleep(3600) # handshake ok, then hangs forever
# Patch the LOCKED body so the real ``_connect_one`` still takes the lock.
with patch.object(mgr, "_connect_one_locked", side_effect=_wedge_locked):
due = _run_hl(loop, mgr._static_reconnect_one("wedge"))
assert mgr._static_reconnect_attempt["wedge"] == 1 # failed attempt
assert not mgr._static_connect_lock_for("wedge").locked() # lock released
assert due > time.monotonic() - 1 # backoff scheduled
def test_health_tick_processes_servers_concurrently(self, running_loop_mgr) -> None:
"""One server stalling its reconnect must not block another server's ping
in the same tick — per-server work runs concurrently."""
mgr, loop, _ = running_loop_mgr
mgr._server_configs = {
"slow": {"type": "stdio", "command": "echo"},
"fast": {"type": "stdio", "command": "echo"},
}
_seed_static_state(mgr, "slow", session=None) # disconnected → reconnect
fast_sess = MagicMock()
_seed_static_state(mgr, "fast", session=fast_sess) # connected → ping
async def _scenario() -> None:
slow_in_connect = asyncio.Event()
fast_pinged = asyncio.Event()
async def _slow_connect(name: str, _cfg: dict[str, Any]) -> None:
slow_in_connect.set()
await asyncio.sleep(3600) # would block the WHOLE tick if serial
async def _fast_ping() -> None:
fast_pinged.set()
fast_sess.send_ping = _fast_ping
with patch.object(mgr, "_connect_one_locked", side_effect=_slow_connect):
tick = asyncio.ensure_future(mgr._static_health_tick())
# Sequential-with-slow-first would never reach 'fast'; concurrency
# means 'fast' is pinged while 'slow' is still stuck connecting.
await asyncio.wait_for(slow_in_connect.wait(), timeout=2)
await asyncio.wait_for(fast_pinged.wait(), timeout=2)
tick.cancel()
with suppress(BaseException):
await tick
_run_hl(loop, _scenario())
def test_health_tick_sleep_uses_fresh_clock(self, running_loop_mgr) -> None:
"""The returned sleep is computed against a FRESH clock, so a tick whose
per-server work burned real time doesn't over-sleep by that elapsed time."""
mgr, loop, _ = running_loop_mgr
mgr._server_configs = {"srv": {"type": "stdio", "command": "echo"}}
_seed_static_state(mgr, "srv", session=None)
async def _slow_reconnect(name: str) -> float:
await asyncio.sleep(0.3) # burn real time during the pass
return time.monotonic() + 1.0 # due in ~1s from now
with patch.object(mgr, "_static_reconnect_one", side_effect=_slow_reconnect):
sleep_s = _run_hl(loop, mgr._static_health_tick())
# The reconnect returns "due ~1s from post-sleep clock", and the tick
# subtracts its own fresh after-clock — the result should be near 1.0.
assert 0.7 <= sleep_s <= 1.3
def test_tick_skips_double_underscore_names(self, running_loop_mgr) -> None:
"""A ``__``-containing server can never connect (reserved delimiter); the
tick skips it entirely rather than retrying forever + spamming log.error."""
mgr, loop, _ = running_loop_mgr
mgr._server_configs = {"bad__name": {"type": "stdio", "command": "echo"}}
with (
patch.object(mgr, "_static_reconnect_one", new=AsyncMock()) as recon,
patch.object(mgr, "_static_ping_one", new=AsyncMock()) as ping,
):
sleep_s = _run_hl(loop, mgr._static_health_tick())
recon.assert_not_awaited() # never even attempted
ping.assert_not_awaited()
assert sleep_s == mgr._static_health_check_s # nothing due → full cadence
# -- cross-path coordination (operator / dispatch) ----------------------
def test_reconnect_sync_holds_connect_lock(self, running_loop_mgr) -> None:
"""Operator ``reconnect_sync`` must hold the per-name connect lock across
teardown+rebuild so an autonomous health reconnect can't interleave."""
mgr, loop, _ = running_loop_mgr
held: list[bool] = []
async def _connect_locked(name: str, _cfg: dict[str, Any]) -> None:
held.append(mgr._static_connect_lock_for(name).locked())
_seed_static_state(mgr, name, session=MagicMock())
with (
patch.object(mgr, "_connect_one_locked", side_effect=_connect_locked),
patch.object(mgr, "_pre_close_streams", new=AsyncMock()),
):
result = mgr.reconnect_sync("srv")
assert result["connected"] is True
assert held == [True] # the lock was held while rebuilding
def test_reconnect_sync_timeout_cleans_catalog(self, running_loop_mgr) -> None:
"""The inner ``asyncio.timeout`` in ``reconnect_sync``'s ``_reconnect``
fires before the caller-side ``future.result(timeout=...)``, triggers the
catalog cleanup (empty tools/resources/prompts so the merged maps don't
advertise entries with no live session), nulls the stale session, and
returns an error dict."""
mgr, loop, _ = running_loop_mgr
mgr._server_configs["srv"] = {"type": "stdio", "command": "echo"}
mgr._STATIC_RECONNECT_ATTEMPT_TIMEOUT_S = 0.1 # fire quickly
_seed_static_state(mgr, "srv", session=MagicMock())
state = mgr._static_servers["srv"]
state.tools = [{"name": "ghost_tool"}]
state.resources = [{"uri": "ghost://resource"}]
state.prompts = [{"name": "ghost_prompt"}]
async def _stall(name: str, _cfg: dict[str, Any]) -> None:
await asyncio.sleep(3600) # never completes — timeout fires first
with (
patch.object(mgr, "_connect_one_locked", side_effect=_stall),
patch.object(mgr, "_pre_close_streams", new=AsyncMock()),
):
result = mgr.reconnect_sync("srv")
assert result["connected"] is False
assert "timed out" in result["error"].lower()
assert result["tools"] == 0
assert result["resources"] == 0
assert result["prompts"] == 0
# Stale session was nulled after the timeout
assert mgr._static_servers["srv"].session is None
# Lock was released after the timeout + cleanup
assert not mgr._static_connect_lock_for("srv").locked()
def test_cb_auto_reconnect_reuses_existing_session(self, running_loop_mgr) -> None:
"""If a session is already live (a health reconnect established it), a
dispatch's ``_cb_auto_reconnect`` REUSES it — no second reconnect, no
spurious breaker failure."""
mgr, loop, _ = running_loop_mgr
existing = MagicMock()
_seed_static_state(mgr, "srv", session=existing) # already up
with (
patch.object(mgr, "_connect_one_locked", new=AsyncMock()) as connect,
patch.object(mgr, "_refresh_server", new=AsyncMock()),
):
session = mgr._cb_auto_reconnect("srv")
assert session is existing # reused
connect.assert_not_awaited() # did NOT reconnect
assert "srv" not in mgr._consecutive_failures # no spurious failure
def test_cb_auto_reconnect_reuses_session_established_under_lock(
self, running_loop_mgr
) -> None:
"""A health reconnect that finishes while the dispatch waits for the lock
is reused: ``_cb_auto_reconnect`` re-checks the session AFTER acquiring
the lock and does not reconnect again or record a spurious failure."""
import threading as _threading
mgr, loop, _ = running_loop_mgr
mgr._server_configs["srv"] = {"type": "stdio", "command": "echo"}
_seed_static_state(mgr, "srv", session=None) # down at the first check
fresh = MagicMock()
async def _acquire_hold() -> asyncio.Lock:
lock = mgr._static_connect_lock_for("srv")
await lock.acquire() # stand in for an in-progress health reconnect
return lock
lock = asyncio.run_coroutine_threadsafe(_acquire_hold(), loop).result(timeout=5)
result: dict[str, Any] = {}
error: dict[str, Exception] = {}
def _dispatch() -> None:
try:
result["session"] = mgr._cb_auto_reconnect("srv")
except Exception as exc:
error["exc"] = exc
with (
patch.object(mgr, "_connect_one_locked", new=AsyncMock()) as connect,
patch.object(mgr, "_refresh_server", new=AsyncMock()),
):
worker = _threading.Thread(target=_dispatch)
worker.start()
try:
time.sleep(0.2) # let the dispatch reach the lock wait
async def _finish() -> None:
mgr._static_servers["srv"].session = fresh # reconnect finished
lock.release()
asyncio.run_coroutine_threadsafe(_finish(), loop).result(timeout=5)
worker.join(timeout=5)
finally:
# Never leak the worker (conftest thread guard): if it is still
# blocked, release the lock on the loop and let it drain.
if worker.is_alive():
async def _emergency() -> None:
if lock.locked():
lock.release()
asyncio.run_coroutine_threadsafe(_emergency(), loop).result(timeout=5)
worker.join(timeout=5)
assert not worker.is_alive()
assert "exc" not in error, error.get("exc")
assert result["session"] is fresh # reused the under-lock session
connect.assert_not_awaited() # did NOT reconnect again
assert "srv" not in mgr._consecutive_failures # no spurious failure
# -- tick scope + lifecycle ---------------------------------------------
def test_tick_skips_oauth_user_servers(self, running_loop_mgr) -> None:
mgr, loop, _ = running_loop_mgr
mgr._server_configs["pool"] = {"type": "http", "url": "https://x/mcp"}
mgr._oauth_user_server_names = {"pool"}
now = time.monotonic()
with (
patch.object(mgr, "_static_ping_one", new=AsyncMock(return_value=now)) as ping,
patch.object(mgr, "_static_reconnect_one", new=AsyncMock(return_value=now)) as recon,
):
_run_hl(loop, mgr._static_health_tick())
for call in ping.await_args_list + recon.await_args_list:
assert call.args[0] != "pool" # oauth_user pools are managed separately
def test_connect_all_starts_health_task_and_disable_gates_it(self, running_loop_mgr) -> None:
mgr, loop, _ = running_loop_mgr
mgr._server_configs = {}
_run_hl(loop, mgr._connect_all())
assert mgr._static_health_task is not None # started
mgr2 = MCPClientManager({})
mgr2._loop = loop
mgr2._server_configs = {}
mgr2._static_health_check_s = 0 # disabled
_run_hl(loop, mgr2._connect_all())
assert mgr2._static_health_task is None # not started
def test_health_loop_cancel_returns_cleanly(self, running_loop_mgr) -> None:
mgr, loop, _ = running_loop_mgr
mgr._server_configs = {} # nothing to do → parks in the sleep
task = _run_hl(loop, _spawn_task(mgr._static_health_loop()))
async def _cancel() -> None:
task.cancel()
with suppress(BaseException):
await task
_run_hl(loop, _cancel())
assert task.cancelled() or task.done()
# -- unified reconnect coordination (round 3) ----------------------------
def test_reconnect_one_defers_while_sibling_in_flight(self, running_loop_mgr) -> None:
"""The health loop DEFERS (short recheck; no backoff bump, no breaker)
while a sibling dispatch still runs on the old evicted stack — tearing
down now would abort that call mid-flight (review finding [0])."""
mgr, loop, _ = running_loop_mgr
mgr._server_configs["down"] = {"type": "stdio", "command": "echo"}
state = _seed_static_state(mgr, "down", session=None)
state.in_flight = 1 # a call_tool still draining on the evicted stack
before = time.monotonic()
with patch.object(mgr, "_connect_one_locked", new=AsyncMock()) as cm:
due = _run_hl(loop, mgr._static_reconnect_one("down"))
cm.assert_not_awaited() # no teardown-under-dispatch
assert before < due <= time.monotonic() + 1.5 # ~1s recheck, not backoff
assert "down" not in mgr._static_reconnect_attempt # not a failed attempt
assert "down" not in mgr._consecutive_failures # breaker untouched
def test_reconnect_one_failure_deadline_uses_fresh_clock(self, running_loop_mgr) -> None:
"""A failed attempt's next-due is written from a FRESH clock — scheduling
from a stale tick-start ``now`` (a slow sibling op ran first in the same
gather) lands the deadline in the past and collapses the backoff into an
every-tick retry storm (review finding [3])."""
mgr, loop, _ = running_loop_mgr
mgr._server_configs["down"] = {"type": "stdio", "command": "echo"}
_seed_static_state(mgr, "down", session=None)
async def _boom(name: str, _cfg: dict[str, Any]) -> None:
raise RuntimeError("still down")
with patch.object(mgr, "_connect_one_locked", side_effect=_boom):
before = time.monotonic()
due = _run_hl(loop, mgr._static_reconnect_one("down"))
assert due >= before # future-dated, not tick-start-relative
assert mgr._static_reconnect_next["down"] == due
def test_ping_one_deadline_uses_fresh_clock(self, running_loop_mgr) -> None:
"""The next-ping deadline is written from a FRESH clock — a stale
tick-start ``now`` would schedule the next ping in the past and re-ping
the server on every tick (review finding [3])."""
mgr, loop, _ = running_loop_mgr
sess = MagicMock()
sess.send_ping = AsyncMock()
_seed_static_state(mgr, "up", session=sess)
stale_now = time.monotonic() - 45.0
before = time.monotonic()
due = _run_hl(loop, mgr._static_ping_one("up", stale_now))
assert due >= before + mgr._static_health_check_s - 1.0
assert mgr._static_next_ping["up"] == due
def test_health_tick_isolates_per_server_cancelled_error(self, running_loop_mgr) -> None:
"""A CancelledError in the gather RESULTS is per-server fallout, not
shutdown — a genuine shutdown cancels the ``await gather`` itself and
never lands in the results list. Re-raising it killed the whole loop."""
mgr, loop, _ = running_loop_mgr
mgr._server_configs = {"srv": {"type": "stdio", "command": "echo"}}
_seed_static_state(mgr, "srv", session=None)
async def _stray(name: str) -> float:
raise asyncio.CancelledError
with patch.object(mgr, "_static_reconnect_one", side_effect=_stray):
sleep_s = _run_hl(loop, mgr._static_health_tick())
# Returned a sleep instead of re-raising; the server was rescheduled
# on the normal cadence.
assert 0.5 <= sleep_s <= mgr._static_health_check_s + 1.0
def test_health_loop_absorbs_stray_cancel_and_stops_on_real_cancel(
self, running_loop_mgr
) -> None:
"""A stray CancelledError escaping the tick (no pending cancel request
on the task) must not kill the loop; a genuine ``task.cancel()`` still
stops it promptly."""
import threading as _threading
mgr, loop, _ = running_loop_mgr
mgr._static_health_check_s = 0.05 # instance shadow for a fast test
survived = _threading.Event()
calls = 0
async def _tick() -> float:
nonlocal calls
calls += 1
if calls == 1:
raise asyncio.CancelledError # stray — the task was NOT cancelled
survived.set()
return 3600.0 # park until the genuine cancel below
with patch.object(mgr, "_static_health_tick", side_effect=_tick):
task = _run_hl(loop, _spawn_task(mgr._static_health_loop()))
assert survived.wait(timeout=5), "loop died on a stray per-server cancel"
assert not task.done()
async def _cancel() -> None:
task.cancel()
with suppress(BaseException):
await task
_run_hl(loop, _cancel())
assert task.done()
def test_dispatch_reconnect_lock_contention_records_no_breaker_failure(
self, running_loop_mgr
) -> None:
"""Review finding [1]: a dispatch reconnect that times out while merely
QUEUED on the per-name lock (held by a longer-bounded health-loop
attempt) must not advance the breaker — the server was never proven
unreachable. The breaker is owned by _ensure_static_connected."""
mgr, loop, _ = running_loop_mgr
# Instance-shadow the sync-boundary wait (now the caller-timeout constant,
# not _CONNECT_TIMEOUT) so the held-lock contention path resolves fast.
mgr._STATIC_RECONNECT_CALLER_TIMEOUT_S = 1.0
_seed_static_state(mgr, "srv", session=None)
async def _hold() -> asyncio.Lock:
lock = mgr._static_connect_lock_for("srv")
await lock.acquire() # stand in for a health-loop reconnect
return lock
lock = asyncio.run_coroutine_threadsafe(_hold(), loop).result(timeout=5)
try:
with (
patch.object(mgr, "_connect_one_locked", new=AsyncMock()) as cm,
pytest.raises(RuntimeError, match="reconnect timed out"),
):
mgr._cb_auto_reconnect("srv")
finally:
async def _release() -> None:
if lock.locked():
lock.release()
asyncio.run_coroutine_threadsafe(_release(), loop).result(timeout=5)
cm.assert_not_awaited()
assert "srv" not in mgr._consecutive_failures # lock wait != failure
def test_dispatch_reconnect_real_failure_records_breaker_once(self, running_loop_mgr) -> None:
"""A REAL connect failure through the dispatch path advances the breaker
exactly ONCE — recorded inside _ensure_static_connected; the sync
boundary must not double-record the same outcome."""
mgr, loop, _ = running_loop_mgr
_seed_static_state(mgr, "srv", session=None)
async def _boom(name: str, _cfg: dict[str, Any]) -> None:
raise ConnectionError("refused")
with (
patch.object(mgr, "_connect_one_locked", side_effect=_boom),
pytest.raises(RuntimeError, match="reconnect failed"),
):
mgr._cb_auto_reconnect("srv")
assert mgr._consecutive_failures.get("srv") == 1
def test_dispatch_reconnect_does_not_resurrect_removed_server(self, running_loop_mgr) -> None:
"""Review finding [2]: a dispatch racing remove_server_sync must not
rebuild the server from its pre-lock cfg snapshot once the config is
gone — _ensure_static_connected re-checks under the lock."""
import threading as _threading
mgr, loop, _ = running_loop_mgr
_seed_static_state(mgr, "srv", session=None)
async def _hold() -> asyncio.Lock:
lock = mgr._static_connect_lock_for("srv")
await lock.acquire()
return lock
lock = asyncio.run_coroutine_threadsafe(_hold(), loop).result(timeout=5)
error: dict[str, Exception] = {}
def _dispatch() -> None:
try:
mgr._cb_auto_reconnect("srv")
except Exception as exc:
error["exc"] = exc
with patch.object(mgr, "_connect_one_locked", new=AsyncMock()) as cm:
worker = _threading.Thread(target=_dispatch)
worker.start()
try:
time.sleep(0.2) # let the dispatch queue on the lock
async def _remove_and_release() -> None:
mgr._server_configs.pop("srv", None) # the remove wins the race
lock.release()
asyncio.run_coroutine_threadsafe(_remove_and_release(), loop).result(timeout=5)
worker.join(timeout=5)
finally:
# Never leak the worker (conftest thread guard).
if worker.is_alive():
async def _emergency() -> None:
if lock.locked():
lock.release()
asyncio.run_coroutine_threadsafe(_emergency(), loop).result(timeout=5)
worker.join(timeout=5)
assert not worker.is_alive()
cm.assert_not_awaited() # no rebuild from the stale cfg
assert isinstance(error.get("exc"), RuntimeError) # failed cleanly
assert "unavailable" in str(error["exc"])
assert "srv" not in mgr._consecutive_failures # not a breaker failure
class TestEnsureStaticConnected:
"""The ONE lazy-connect primitive every autonomous driver routes through
(health loop, dispatch _cb_auto_reconnect, _refresh_all). Operator
reconnect_sync deliberately stays a force rebuild outside it."""
def test_concurrent_drivers_collapse_to_single_connect(self, running_loop_mgr) -> None:
"""The reconnect STORM fix: N concurrent callers for one server queue on
the per-name lock and collapse to a SINGLE _connect_one_locked; everyone
else reuses the installed session (never tears it down to rebuild)."""
mgr, loop, _ = running_loop_mgr
_seed_static_state(mgr, "srv", session=None)
sess = MagicMock()
connects = 0
async def _connect(name: str, _cfg: dict[str, Any]) -> None:
nonlocal connects
connects += 1
await asyncio.sleep(0.05) # hold the lock so siblings queue
_seed_static_state(mgr, name, session=sess)
async def _storm() -> list[Any]:
cfg = mgr._server_configs["srv"]
return await asyncio.gather(
*(mgr._ensure_static_connected("srv", cfg) for _ in range(5))
)
with patch.object(mgr, "_connect_one_locked", side_effect=_connect):
sessions = _run_hl(loop, _storm())
assert connects == 1 # queued callers reused, not rebuilt
assert all(s is sess for s in sessions)
def test_reuses_live_session_without_teardown(self, running_loop_mgr) -> None:
mgr, loop, _ = running_loop_mgr
sess = MagicMock()
_seed_static_state(mgr, "srv", session=sess)
with patch.object(mgr, "_connect_one_locked", new=AsyncMock()) as cm:
out = _run_hl(loop, mgr._ensure_static_connected("srv", mgr._server_configs["srv"]))
assert out is sess
cm.assert_not_awaited()
def test_returns_none_for_removed_server(self, running_loop_mgr) -> None:
"""Config re-check under the lock: a concurrently-removed server is not
resurrected from the caller's pre-lock cfg snapshot."""
mgr, loop, _ = running_loop_mgr
cfg = {"type": "stdio", "command": "echo"} # caller's stale snapshot
assert "gone" not in mgr._server_configs
with patch.object(mgr, "_connect_one_locked", new=AsyncMock()) as cm:
out = _run_hl(loop, mgr._ensure_static_connected("gone", cfg))
assert out is None
cm.assert_not_awaited()
assert "gone" not in mgr._static_servers # nothing rebuilt
def test_defers_when_sibling_call_in_flight(self, running_loop_mgr) -> None:
"""session None + in_flight > 0 → defer (None) without teardown; once
the sibling call drains, the next call reconnects."""
mgr, loop, _ = running_loop_mgr
state = _seed_static_state(mgr, "srv", session=None)
state.in_flight = 1
sess = MagicMock()
async def _connect(name: str, _cfg: dict[str, Any]) -> None:
_seed_static_state(mgr, name, session=sess)
cfg = mgr._server_configs["srv"]
with patch.object(mgr, "_connect_one_locked", side_effect=_connect) as cm:
out = _run_hl(loop, mgr._ensure_static_connected("srv", cfg))
assert out is None
assert cm.await_count == 0 # no teardown-under-dispatch
state.in_flight = 0 # the sibling call finished
out2 = _run_hl(loop, mgr._ensure_static_connected("srv", cfg))
assert out2 is sess
assert cm.await_count == 1
def test_dispatch_does_not_defer_when_busy(self, running_loop_mgr) -> None:
"""Round-4 [3]: a DISPATCH (defer_if_busy=False) reconnects even with an
in-flight sibling — it needs the session now — rather than hard-failing a
reachable server. The autonomous default still defers (test above)."""
mgr, loop, _ = running_loop_mgr
state = _seed_static_state(mgr, "srv", session=None)
state.in_flight = 1
sess = MagicMock()
async def _connect(name: str, _cfg: dict[str, Any]) -> None:
_seed_static_state(mgr, name, session=sess)
with patch.object(mgr, "_connect_one_locked", side_effect=_connect) as cm:
out = _run_hl(
loop,
mgr._ensure_static_connected(
"srv", mgr._server_configs["srv"], defer_if_busy=False
),
)
assert out is sess # reconnected despite in_flight > 0
assert cm.await_count == 1
def test_cancelled_attempt_cleans_up_partial_session_and_no_breaker_record(
self, running_loop_mgr
) -> None:
"""A bare CancelledError delivered mid-connect (a caller
sync-boundary giving up, or shutdown) must NOT leave a half-discovered
session installed but must NOT record a breaker failure — CancelledError
proves nothing about the server."""
mgr, loop, _ = running_loop_mgr
state = _seed_static_state(mgr, "srv", session=None)
async def _handshake_then_cancel(name: str, _cfg: dict[str, Any]) -> None:
state.session = MagicMock() # handshake OK, session installed
raise asyncio.CancelledError() # discovery cancelled by a caller giving up
with (
patch.object(mgr, "_connect_one_locked", side_effect=_handshake_then_cancel),
# run_coroutine_threadsafe surfaces a re-raised asyncio.CancelledError
# as concurrent.futures.CancelledError at the .result() boundary.
pytest.raises((asyncio.CancelledError, concurrent.futures.CancelledError)),
):
_run_hl(loop, mgr._ensure_static_connected("srv", mgr._server_configs["srv"]))
assert mgr._static_servers["srv"].session is None # partial session dropped
assert "srv" not in mgr._consecutive_failures # CancelledError is NOT a server failure
def test_timeout_hierarchy_caller_exceeds_attempt(self) -> None:
"""The systemic round-4 fix: the caller wait must exceed the inner attempt
bound so the inner asyncio.timeout fires first (clean TimeoutError), never
a bare cancel escaping — and the attempt must cover a full handshake."""
assert (
MCPClientManager._STATIC_RECONNECT_CALLER_TIMEOUT_S
> MCPClientManager._STATIC_RECONNECT_ATTEMPT_TIMEOUT_S
> MCPClientManager._CONNECT_TIMEOUT
)
def test_failure_records_one_breaker_failure_and_raises(self, running_loop_mgr) -> None:
mgr, loop, _ = running_loop_mgr
_seed_static_state(mgr, "srv", session=None)
async def _boom(name: str, _cfg: dict[str, Any]) -> None:
raise ConnectionError("refused")
with (
patch.object(mgr, "_connect_one_locked", side_effect=_boom),
pytest.raises(ConnectionError),
):
_run_hl(loop, mgr._ensure_static_connected("srv", mgr._server_configs["srv"]))
assert mgr._consecutive_failures.get("srv") == 1
def test_success_clears_only_open_circuit_deadline(self, running_loop_mgr) -> None:
"""Finding-13 semantics, now owned by the primitive: success re-opens
dispatch (deadline cleared) but keeps _consecutive_failures so a
connect-ok / calls-fail server still escalates to a trip."""
mgr, loop, _ = running_loop_mgr
_seed_static_state(mgr, "srv", session=None)
for _ in range(3):
mgr._cb_record_failure("srv")
assert "srv" in mgr._circuit_open_until
sess = MagicMock()
async def _connect(name: str, _cfg: dict[str, Any]) -> None:
_seed_static_state(mgr, name, session=sess)
with patch.object(mgr, "_connect_one_locked", side_effect=_connect):
out = _run_hl(loop, mgr._ensure_static_connected("srv", mgr._server_configs["srv"]))
assert out is sess
assert "srv" not in mgr._circuit_open_until # dispatch flows again
assert mgr._consecutive_failures.get("srv", 0) >= 3 # count kept
def test_refresh_all_reconnect_keeps_breaker_failure_count(self) -> None:
"""_refresh_all's reconnect branch routes through the primitive: only
the open-circuit deadline clears — the old full _cb_record_success
reset let a connect-ok / calls-fail server oscillate 0->1->0 below the
breaker threshold forever."""
async def _run() -> None:
mgr = MCPClientManager({})
mgr._server_configs["srv"] = {"type": "stdio", "command": "x"}
mgr._consecutive_failures["srv"] = 2
sess = MagicMock()
async def _connect(name: str, _cfg: dict[str, Any]) -> None:
_seed_static_state(mgr, name, session=sess)
with patch.object(mgr, "_connect_one_locked", side_effect=_connect):
results = await mgr._refresh_all("srv")
assert results["srv"] == ([], []) # reconnected; no tools seeded
assert mgr._last_refresh["srv"][1] == "ok"
assert mgr._consecutive_failures.get("srv") == 2 # NOT reset
assert mgr._static_servers["srv"].session is sess
asyncio.run(_run())
def test_remove_server_sync_cancels_on_timeout_and_reports_failure(
self, running_loop_mgr
) -> None:
"""Round-4 [2]: if teardown can't finish within the caller timeout (a slow
reconnect holding the per-name lock), remove_server_sync CANCELS the
pending _remove — so it can't later pop a re-added entry — and returns
False rather than a false 'removed'."""
mgr, loop, _ = running_loop_mgr
_seed_static_state(mgr, "srv", session=MagicMock())
async def _hang(_name: str) -> None:
await asyncio.sleep(10) # teardown wedged (stands in for lock contention)
with patch.object(mgr, "_teardown_static_session", side_effect=_hang):
result = mgr.remove_server_sync("srv", timeout=0.2)
assert result is False # not a false success
assert "srv" not in mgr._server_configs # config still popped (won't reconnect)
def test_schedule_next_ping_sets_and_returns_fresh_deadline(self) -> None:
"""Round-4 [5]: the extracted ping-scheduling helper stores and returns
the same fresh deadline (was a copy-pasted triplet in 3 branches)."""
mgr = MCPClientManager({})
mgr._static_health_check_s = 30.0
before = time.monotonic()
due = mgr._schedule_next_ping("srv")
assert mgr._static_next_ping["srv"] == due
assert before + 30.0 <= due <= time.monotonic() + 30.0
class TestTeardownStaticSession:
"""The one canonical teardown sequence (shared by _connect_one_locked's
stale-guard and remove_server_sync)."""
def test_teardown_order_and_state_cleared(self, running_loop_mgr) -> None:
"""Close protocol: session nulled, close event set BEFORE the first
await (a teardown cancelled mid-flight must still have delivered the
owner's marching orders), streams pre-closed, then the parked owner
exits GRACEFULLY — no cancel."""
mgr, loop, _ = running_loop_mgr
order: list[str] = []
async def _make_owner() -> tuple[asyncio.Event, asyncio.Task[None]]:
ev = asyncio.Event()
async def _parked_owner() -> None:
await ev.wait()
order.append("owner_exit")
task = asyncio.create_task(_parked_owner())
await asyncio.sleep(0) # let the owner park
return ev, task
ev, owner = _run_hl(loop, _make_owner())
async def _pre_close(name: str) -> None:
order.append("pre_close")
# Session nulled FIRST so concurrent dispatch reads see
# "disconnected", not a corpse.
assert mgr._static_servers["srv"].session is None
# The close signal precedes the first await of the teardown.
assert ev.is_set()
_seed_static_state(mgr, "srv", session=MagicMock(), owner_task=owner, close_requested=ev)
with patch.object(mgr, "_pre_close_streams", side_effect=_pre_close):
_run_hl(loop, mgr._teardown_static_session("srv"))
assert order == ["pre_close", "owner_exit"]
state = mgr._static_servers["srv"]
assert state.session is None
assert state.owner_task is None
assert state.close_requested is None
assert owner.done() and not owner.cancelled() # graceful, no escalation
def test_teardown_escalates_to_single_cancel(self, running_loop_mgr) -> None:
"""An owner that ignores the close event gets EXACTLY one cancel — a
second cancel is the zombie-minting mistake the protocol forbids, so
the count is pinned, not just the final cancelled state."""
mgr, loop, _ = running_loop_mgr
mgr._OWNER_CLOSE_GRACE_S = 0.05 # keep the graceful window short
cancel_calls: list[Any] = []
async def _make_owner() -> asyncio.Task[None]:
async def _stubborn_owner() -> None:
await asyncio.sleep(3600) # never watches the event
task = asyncio.create_task(_stubborn_owner())
await asyncio.sleep(0)
real_cancel = task.cancel
def _counting_cancel(*args: Any, **kwargs: Any) -> bool:
cancel_calls.append(args)
return real_cancel(*args, **kwargs)
task.cancel = _counting_cancel # type: ignore[method-assign]
return task
owner = _run_hl(loop, _make_owner())
_seed_static_state(
mgr, "srv", session=MagicMock(), owner_task=owner, close_requested=asyncio.Event()
)
_run_hl(loop, mgr._teardown_static_session("srv"))
assert owner.cancelled()
assert len(cancel_calls) == 1 # one cancel, never a second
assert mgr._static_servers["srv"].owner_task is None
def test_teardown_missing_server_is_noop(self, running_loop_mgr) -> None:
mgr, loop, _ = running_loop_mgr
_run_hl(loop, mgr._teardown_static_session("nope")) # must not raise
async def _spawn_task(coro: Any) -> asyncio.Task[Any]:
return asyncio.ensure_future(coro)