refactor(mcp): consolidate per-server state into StaticServerState dataclass

Phase 0 of the OAuth-MCP RFC: prepare MCPClientManager for the per-(user,
server) session pool that lands in Phase 5, without changing static-path
behavior.

Two changes:

1. Hardening helpers _pre_close_streams and _tcp_probe rename their first
   parameter from `name` to `key`.  Type stays `str` for now; widening to
   `str | tuple[str, str]` happens in Phase 5 when callers actually pass
   tuples.  _safe_close_stack takes the stack directly and is unchanged.

2. The eleven parallel name-keyed dicts (_sessions, _per_server_stacks,
   _per_server_tools, _per_server_resources, _per_server_prompts,
   _supports_list_changed, _supports_resources, _supports_resource_list_changed,
   _supports_prompts, _supports_prompt_list_changed, _server_streams) are
   consolidated into _static_servers: dict[str, StaticServerState].  Server-
   level state (circuit breaker, notification debounce, last-error,
   db-managed, merged catalog maps, listener lists) stays on the manager,
   unchanged.

PoolEntryState is defined for Phase 5 use but no code instantiates it.  The
typed map declarations (dict[str, StaticServerState] vs dict[tuple[str, str],
PoolEntryState]) make accidental cross-keying lookups easier to catch.

PR #296 hardening preserved exactly:
- pre-close-streams atomic take-and-clear before stack teardown
- stale-session-and-stack guard at _connect_one top: both state.session and
  state.stack checked, cleared independently, entry preserved (not popped)
- transport-error session-eviction in dispatch sets state.session=None only,
  leaving stack/streams for the next connect-time guard sweep
- _safe_close_stack CancelledError suppression unchanged
- TCP probe before streamablehttp_client unchanged
- future.cancel() after TimeoutError in all sync bridges unchanged
- notification debounce stays manager-level (not migrated into the dataclass)

Refresh helpers (_refresh_server_tools/_resources/_prompts) snapshot
state.session into a local immediately after the None guard so concurrent
transport-error eviction during await cannot null the session reference
mid-call.

Tests: shared _seed_static_state helper in tests/conftest.py replaces eleven
direct dict mutations; new test_reconnect_preserves_static_state_identity
guards the entry-preservation invariant.  Pass count rises 5266 → 5267.

(cherry picked from commit be0950bb98)
This commit is contained in:
Patrick Buckley
2026-05-04 17:26:32 -07:00
parent bace928477
commit c823156af5
5 changed files with 561 additions and 355 deletions
+19
View File
@@ -7,9 +7,28 @@ from unittest.mock import MagicMock
import pytest
if TYPE_CHECKING:
from turnstone.core.mcp_client import MCPClientManager, StaticServerState
from turnstone.core.oidc import OIDCConfig
def _seed_static_state(mgr: MCPClientManager, name: str, **overrides: Any) -> StaticServerState:
"""Get-or-create a ``StaticServerState`` on ``mgr`` and apply ``overrides``.
Shared across MCP test files so the helper stays in one place. Imported
where needed; ``StaticServerState`` is constructed lazily so non-MCP
tests don't pay the import cost.
"""
from turnstone.core.mcp_client import StaticServerState
state = mgr._static_servers.get(name)
if state is None:
state = StaticServerState(name=name)
mgr._static_servers[name] = state
for k, v in overrides.items():
setattr(state, k, v)
return state
def make_oidc_test_config(**overrides: Any) -> OIDCConfig:
"""Build a test ``OIDCConfig`` with sensible defaults.
+195 -115
View File
@@ -12,6 +12,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from tests.conftest import _seed_static_state
from turnstone.core.mcp_client import (
MCPClientManager,
_mcp_to_openai,
@@ -319,8 +320,8 @@ class TestMCPClientManager:
def test_server_count(self):
mgr = MCPClientManager({})
mgr._sessions["a"] = MagicMock()
mgr._sessions["b"] = MagicMock()
_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):
@@ -557,7 +558,7 @@ class TestServerNameValidation:
mgr._exit_stack = stack
await mgr._connect_one("my__bad", {"command": "echo"})
# Should not have connected
assert "my__bad" not in mgr._sessions
assert "my__bad" not in mgr._static_servers
assert mgr.get_tools() == []
asyncio.run(_run())
@@ -585,10 +586,8 @@ class TestCreateMcpClient:
class TestRebuildTools:
def test_rebuild_from_per_server(self):
mgr = MCPClientManager({})
mgr._per_server_tools = {
"github": [_fake_openai_tool("mcp__github__search")],
"slack": [_fake_openai_tool("mcp__slack__send")],
}
_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}
@@ -598,18 +597,18 @@ class TestRebuildTools:
def test_rebuild_copy_on_write(self):
mgr = MCPClientManager({})
mgr._per_server_tools = {"a": [_fake_openai_tool("mcp__a__x")]}
_seed_static_state(mgr, "a", tools=[_fake_openai_tool("mcp__a__x")])
mgr._rebuild_tools()
old_tools = mgr._tools
old_map = mgr._tool_map
mgr._per_server_tools["b"] = [_fake_openai_tool("mcp__b__y")]
_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._per_server_tools = {}
mgr._static_servers = {}
mgr._rebuild_tools()
assert mgr._tools == []
assert mgr._tool_map == {}
@@ -621,8 +620,7 @@ class TestRefreshServer:
mgr: MCPClientManager, server_name: str, mock_session: MagicMock
) -> None:
"""Add empty list_resources/list_prompts mocks so _refresh_server works."""
mgr._supports_resources[server_name] = True
mgr._supports_prompts[server_name] = True
_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)
@@ -644,8 +642,12 @@ class TestRefreshServer:
]
mock_session.list_tools = AsyncMock(return_value=mock_result)
self._add_empty_resource_prompt_mocks(mgr, "github", mock_session)
mgr._sessions["github"] = mock_session
mgr._per_server_tools["github"] = [_fake_openai_tool("mcp__github__search")]
_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")
@@ -663,8 +665,12 @@ class TestRefreshServer:
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)
mgr._sessions["github"] = mock_session
mgr._per_server_tools["github"] = [_fake_openai_tool("mcp__github__search")]
_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")
@@ -682,8 +688,12 @@ class TestRefreshServer:
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)
mgr._sessions["github"] = mock_session
mgr._per_server_tools["github"] = [_fake_openai_tool("mcp__github__search")]
_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")
@@ -706,7 +716,7 @@ class TestListeners:
mgr = MCPClientManager({})
calls: list[int] = []
mgr.add_listener(lambda: calls.append(1))
mgr._per_server_tools = {"a": [_fake_openai_tool("mcp__a__x")]}
_seed_static_state(mgr, "a", tools=[_fake_openai_tool("mcp__a__x")])
mgr._rebuild_tools()
assert len(calls) == 1
@@ -903,12 +913,14 @@ class TestMCPResources:
def test_resource_discovery(self):
"""Mock list_resources() returning 2 resources, verify get_resources()."""
mgr = MCPClientManager({})
mgr._per_server_resources = {
"fs": [
_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
@@ -919,22 +931,18 @@ class TestMCPResources:
def test_rebuild_resources_copy_on_write(self):
"""Verify mutation safety — get_resources() returns independent copy."""
mgr = MCPClientManager({})
mgr._per_server_resources = {
"a": [_fake_resource_dict("file:///x", "x", "", "", "a")],
}
_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
mgr._per_server_resources["b"] = [_fake_resource_dict("file:///y", "y", "", "", "b")]
_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({})
mgr._per_server_resources = {
"a": [_fake_resource_dict("file:///x", "x", "", "", "a")],
}
_seed_static_state(mgr, "a", resources=[_fake_resource_dict("file:///x", "x", "", "", "a")])
mgr._rebuild_resources()
resources = mgr.get_resources()
assert len(resources) == 1
@@ -946,7 +954,7 @@ class TestMCPResources:
mgr = MCPClientManager({})
mgr._resource_map = {"file:///readme": ("fs", "file:///readme")}
mock_session = MagicMock()
mgr._sessions["fs"] = mock_session
_seed_static_state(mgr, "fs", session=mock_session)
mgr._loop = asyncio.new_event_loop()
# Mock the read_resource result
@@ -974,7 +982,7 @@ class TestMCPResources:
mgr = MCPClientManager({})
mgr._resource_map = {"file:///img.png": ("fs", "file:///img.png")}
mock_session = MagicMock()
mgr._sessions["fs"] = mock_session
_seed_static_state(mgr, "fs", session=mock_session)
mgr._loop = asyncio.new_event_loop()
blob_content = MagicMock(spec=["blob"])
@@ -1011,7 +1019,7 @@ class TestMCPResources:
mgr = MCPClientManager({})
mgr._resource_map = {"file:///x": ("fs", "file:///x")}
mock_session = MagicMock()
mgr._sessions["fs"] = mock_session
_seed_static_state(mgr, "fs", session=mock_session)
mgr._loop = asyncio.new_event_loop()
async def _slow_read(_uri: str) -> None:
@@ -1036,7 +1044,7 @@ class TestMCPResources:
mgr = MCPClientManager({})
calls: list[int] = []
mgr.add_resource_listener(lambda: calls.append(1))
mgr._per_server_resources = {"a": [_fake_resource_dict()]}
_seed_static_state(mgr, "a", resources=[_fake_resource_dict()])
mgr._rebuild_resources()
assert len(calls) == 1
@@ -1060,13 +1068,13 @@ class TestMCPResources:
async def _run() -> None:
mgr = MCPClientManager({})
mock_session = MagicMock()
mgr._sessions["fs"] = mock_session
mgr._supports_resources["fs"] = True
# Initial state
mgr._per_server_resources["fs"] = [
_fake_resource_dict("file:///old", server="fs"),
]
_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
@@ -1088,17 +1096,19 @@ class TestMCPResources:
def test_rebuild_resources_empty(self):
mgr = MCPClientManager({})
mgr._per_server_resources = {}
mgr._static_servers = {}
mgr._rebuild_resources()
assert mgr._resources == []
assert mgr._resource_map == {}
def test_rebuild_resources_multi_server(self):
mgr = MCPClientManager({})
mgr._per_server_resources = {
"fs": [_fake_resource_dict("file:///a", server="fs")],
"db": [_fake_resource_dict("db://table", name="table", server="db")],
}
_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")
@@ -1107,8 +1117,10 @@ class TestMCPResources:
def test_template_prefix_matching(self):
"""Expanded URI matches template by prefix."""
mgr = MCPClientManager({})
mgr._per_server_resources = {
"db": [
_seed_static_state(
mgr,
"db",
resources=[
{
"uri": "db://tables/{table}/rows/{id}",
"name": "row",
@@ -1118,7 +1130,7 @@ class TestMCPResources:
"template": True,
},
],
}
)
mgr._rebuild_resources()
# Template should not be in resource_map
assert "db://tables/{table}/rows/{id}" not in mgr._resource_map
@@ -1134,8 +1146,10 @@ class TestMCPResources:
mgr = MCPClientManager({})
# Use templates with genuinely different prefix lengths:
# "db://data/" (6 chars after scheme) vs "db://data/tables/" (13 chars after scheme)
mgr._per_server_resources = {
"short": [
_seed_static_state(
mgr,
"short",
resources=[
{
"uri": "db://data/{collection}",
"name": "collection",
@@ -1145,7 +1159,11 @@ class TestMCPResources:
"template": True,
},
],
"long": [
)
_seed_static_state(
mgr,
"long",
resources=[
{
"uri": "db://data/tables/{table}",
"name": "table",
@@ -1155,7 +1173,7 @@ class TestMCPResources:
"template": True,
},
],
}
)
mgr._rebuild_resources()
# "db://data/tables/users" matches both prefixes ("db://data/" and
# "db://data/tables/") — the longer one should win
@@ -1172,8 +1190,10 @@ class TestMCPResources:
def test_template_no_match_raises(self):
"""Completely unrelated URI still raises ValueError."""
mgr = MCPClientManager({})
mgr._per_server_resources = {
"db": [
_seed_static_state(
mgr,
"db",
resources=[
{
"uri": "db://tables/{table}",
"name": "table",
@@ -1183,7 +1203,7 @@ class TestMCPResources:
"template": True,
},
],
}
)
mgr._rebuild_resources()
assert mgr._match_template("file:///something") is None
with pytest.raises(ValueError, match="Unknown MCP resource"):
@@ -1192,8 +1212,10 @@ class TestMCPResources:
def test_read_resource_sync_with_template_uri(self):
"""End-to-end: template discovered, expanded URI dispatched to correct server."""
mgr = MCPClientManager({})
mgr._per_server_resources = {
"db": [
_seed_static_state(
mgr,
"db",
resources=[
{
"uri": "db://tables/{table}/rows/{id}",
"name": "row",
@@ -1203,11 +1225,11 @@ class TestMCPResources:
"template": True,
},
],
}
)
mgr._rebuild_resources()
mock_session = MagicMock()
mgr._sessions["db"] = mock_session
_seed_static_state(mgr, "db", session=mock_session)
mgr._loop = asyncio.new_event_loop()
text_content = MagicMock(spec=["text"])
@@ -1239,12 +1261,14 @@ class TestMCPPrompts:
def test_prompt_discovery(self):
"""Mock list_prompts(), verify get_prompts() with correct prefixed names."""
mgr = MCPClientManager({})
mgr._per_server_prompts = {
"tmpl": [
_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
@@ -1257,22 +1281,18 @@ class TestMCPPrompts:
def test_rebuild_prompts_copy_on_write(self):
"""Verify mutation safety."""
mgr = MCPClientManager({})
mgr._per_server_prompts = {
"a": [_fake_prompt_dict("mcp__a__p1", "p1", "a")],
}
_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
mgr._per_server_prompts["b"] = [_fake_prompt_dict("mcp__b__p2", "p2", "b")]
_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({})
mgr._per_server_prompts = {
"a": [_fake_prompt_dict("mcp__a__p1", "p1", "a")],
}
_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
@@ -1284,7 +1304,7 @@ class TestMCPPrompts:
mgr = MCPClientManager({})
mgr._prompt_map = {"mcp__tmpl__review": ("tmpl", "review")}
mock_session = MagicMock()
mgr._sessions["tmpl"] = mock_session
_seed_static_state(mgr, "tmpl", session=mock_session)
mgr._loop = asyncio.new_event_loop()
# Build mock PromptMessage
@@ -1335,7 +1355,7 @@ class TestMCPPrompts:
mgr = MCPClientManager({})
mgr._prompt_map = {"mcp__tmpl__slow": ("tmpl", "slow")}
mock_session = MagicMock()
mgr._sessions["tmpl"] = mock_session
_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:
@@ -1360,7 +1380,7 @@ class TestMCPPrompts:
mgr = MCPClientManager({})
calls: list[int] = []
mgr.add_prompt_listener(lambda: calls.append(1))
mgr._per_server_prompts = {"a": [_fake_prompt_dict()]}
_seed_static_state(mgr, "a", prompts=[_fake_prompt_dict()])
mgr._rebuild_prompts()
assert len(calls) == 1
@@ -1391,13 +1411,13 @@ class TestMCPPrompts:
async def _run() -> None:
mgr = MCPClientManager({})
mock_session = MagicMock()
mgr._sessions["tmpl"] = mock_session
mgr._supports_prompts["tmpl"] = True
# Initial state
mgr._per_server_prompts["tmpl"] = [
_fake_prompt_dict("mcp__tmpl__old", "old", "tmpl"),
]
_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
@@ -1417,17 +1437,15 @@ class TestMCPPrompts:
def test_rebuild_prompts_empty(self):
mgr = MCPClientManager({})
mgr._per_server_prompts = {}
mgr._static_servers = {}
mgr._rebuild_prompts()
assert mgr._prompts == []
assert mgr._prompt_map == {}
def test_rebuild_prompts_multi_server(self):
mgr = MCPClientManager({})
mgr._per_server_prompts = {
"a": [_fake_prompt_dict("mcp__a__p1", "p1", "a")],
"b": [_fake_prompt_dict("mcp__b__p2", "p2", "b")],
}
_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")
@@ -1442,9 +1460,9 @@ class TestMCPPrompts:
class TestShutdownCleanup:
def test_shutdown_clears_resources_and_prompts(self):
mgr = MCPClientManager({})
mgr._per_server_resources = {"a": [_fake_resource_dict()]}
_seed_static_state(mgr, "a", resources=[_fake_resource_dict()])
mgr._rebuild_resources()
mgr._per_server_prompts = {"a": [_fake_prompt_dict()]}
_seed_static_state(mgr, "a", prompts=[_fake_prompt_dict()])
mgr._rebuild_prompts()
assert mgr.get_resources() != []
assert mgr.get_prompts() != []
@@ -1528,8 +1546,9 @@ class TestConnectOneUnreachable:
mgr._loop.run_until_complete(_run())
mgr._loop.close()
# Server should NOT be in sessions (connection failed)
assert "bad-server" not in mgr._sessions
# 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."""
@@ -1543,7 +1562,8 @@ class TestConnectOneUnreachable:
loop.run_until_complete(mgr._connect_all())
loop.close()
assert "bad" not in mgr._sessions
bad_state = mgr._static_servers.get("bad")
assert bad_state is None or bad_state.session is None
assert "bad" in mgr._last_error
@@ -1598,7 +1618,7 @@ class TestFutureCancellation:
mock_session.call_tool = MagicMock(return_value="sentinel")
mock_session.read_resource = MagicMock(return_value="sentinel")
mock_session.get_prompt = MagicMock(return_value="sentinel")
mgr._sessions["test"] = mock_session
_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")
@@ -1766,7 +1786,7 @@ class TestCircuitBreaker:
mgr = MCPClientManager({"test": {"type": "stdio", "command": "echo"}})
mock_session = MagicMock()
mock_session.call_tool = MagicMock(return_value="sentinel")
mgr._sessions["test"] = mock_session
_seed_static_state(mgr, "test", session=mock_session)
mgr._loop = MagicMock()
mgr._tool_map["mcp__test__ping"] = ("test", "ping")
mock_future = MagicMock()
@@ -1782,7 +1802,7 @@ class TestCircuitBreaker:
mgr = MCPClientManager({"test": {"type": "stdio", "command": "echo"}})
mock_session = MagicMock()
mock_session.call_tool = MagicMock(return_value="sentinel")
mgr._sessions["test"] = mock_session
_seed_static_state(mgr, "test", session=mock_session)
mgr._loop = MagicMock()
mgr._tool_map["mcp__test__ping"] = ("test", "ping")
# Pre-set a failure
@@ -1800,7 +1820,10 @@ class TestCircuitBreaker:
mgr = MCPClientManager({"test": {"type": "stdio", "command": "echo"}})
mock_session = MagicMock()
mock_session.call_tool = MagicMock(return_value="sentinel")
mgr._sessions["test"] = mock_session
# Seed both session and stack so the test can verify stack survives.
old_stack = MagicMock()
old_streams = (MagicMock(), MagicMock())
_seed_static_state(mgr, "test", session=mock_session, stack=old_stack, streams=old_streams)
mgr._loop = MagicMock()
mgr._tool_map["mcp__test__ping"] = ("test", "ping")
mock_future = MagicMock()
@@ -1810,7 +1833,12 @@ class TestCircuitBreaker:
pytest.raises(BrokenPipeError),
):
mgr.call_tool_sync("mcp__test__ping", {}, timeout=5)
assert "test" not in mgr._sessions
# Session evicted, but stack/streams remain for the stale-and-stack
# guard in _connect_one to clean up on next reconnect attempt.
state = mgr._static_servers["test"]
assert state.session is None
assert state.stack is old_stack
assert state.streams is old_streams
def test_independent_circuits_per_server(self):
mgr = MCPClientManager({})
@@ -1829,7 +1857,7 @@ class TestCircuitBreaker:
mgr = MCPClientManager({"test": {"type": "stdio", "command": "echo"}})
mock_session = MagicMock()
mock_session.call_tool = MagicMock(return_value="sentinel")
mgr._sessions["test"] = mock_session
_seed_static_state(mgr, "test", session=mock_session)
mgr._loop = MagicMock()
mgr._tool_map["mcp__test__ping"] = ("test", "ping")
mock_future = MagicMock()
@@ -1855,7 +1883,7 @@ class TestSafeTransportStreams:
mgr = MCPClientManager({})
stream_a = MagicMock()
stream_b = MagicMock()
mgr._server_streams["srv"] = (stream_a, stream_b)
_seed_static_state(mgr, "srv", streams=(stream_a, stream_b))
async def _run():
await mgr._pre_close_streams("srv")
@@ -1863,7 +1891,8 @@ class TestSafeTransportStreams:
asyncio.run(_run())
stream_a.aclose.assert_called_once()
stream_b.aclose.assert_called_once()
assert "srv" not in mgr._server_streams
# 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({})
@@ -1878,7 +1907,7 @@ class TestSafeTransportStreams:
stream_a = MagicMock()
stream_a.aclose.side_effect = RuntimeError("boom")
stream_b = MagicMock()
mgr._server_streams["srv"] = (stream_a, stream_b)
_seed_static_state(mgr, "srv", streams=(stream_a, stream_b))
async def _run():
await mgr._pre_close_streams("srv")
@@ -1888,9 +1917,9 @@ class TestSafeTransportStreams:
def test_shutdown_clears_stream_refs(self):
mgr = MCPClientManager({})
mgr._server_streams["srv"] = (MagicMock(), MagicMock())
_seed_static_state(mgr, "srv", streams=(MagicMock(), MagicMock()))
mgr.shutdown()
assert len(mgr._server_streams) == 0
assert len(mgr._static_servers) == 0
# ---------------------------------------------------------------------------
@@ -1950,7 +1979,7 @@ class TestReconnectSync:
mgr, _loop, _thread = running_loop_mgr
async def _fake_connect_one(name: str, _cfg: dict[str, Any]) -> None:
mgr._sessions[name] = MagicMock()
_seed_static_state(mgr, name, session=MagicMock())
# Pre-trip the breaker
for _ in range(3):
@@ -1982,11 +2011,16 @@ class TestReconnectSync:
async def _connect_one(name: str, _cfg: dict[str, Any]) -> None:
order.append("connect_one")
mgr._sessions[name] = MagicMock()
_seed_static_state(mgr, name, session=MagicMock())
mgr._sessions["srv"] = MagicMock() # populated old session
mgr._per_server_stacks["srv"] = old_stack
mgr._server_streams["srv"] = (MagicMock(), MagicMock())
# Seed the old session/stack/streams that the guard should clear.
_seed_static_state(
mgr,
"srv",
session=MagicMock(),
stack=old_stack,
streams=(MagicMock(), MagicMock()),
)
with (
patch.object(mgr, "_pre_close_streams", side_effect=_pre_close),
@@ -1996,7 +2030,8 @@ class TestReconnectSync:
result = mgr.reconnect_sync("srv")
assert result["connected"] is True
assert order == ["pre_close", "safe_close", "connect_one"]
assert "srv" not in mgr._per_server_stacks
# The old stack reference should have been cleared from state.
assert mgr._static_servers["srv"].stack is not old_stack
def test_reconnect_failure_returns_error_dict(self, running_loop_mgr):
mgr, _loop, _thread = running_loop_mgr
@@ -2019,9 +2054,13 @@ class TestReconnectSync:
mgr, _loop, _thread = running_loop_mgr
# Seed catalog state from a previous successful connect.
mgr._per_server_tools["srv"] = [_fake_openai_tool("mcp__srv__t")]
mgr._per_server_resources["srv"] = [_fake_resource_dict(server="srv")]
mgr._per_server_prompts["srv"] = [_fake_prompt_dict(server="srv")]
_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()
@@ -2035,12 +2074,53 @@ class TestReconnectSync:
):
result = mgr.reconnect_sync("srv")
assert result["connected"] is False
# Per-server catalog and merged maps should both be empty for srv.
assert "srv" not in mgr._per_server_tools
assert "srv" not in mgr._per_server_resources
assert "srv" not in mgr._per_server_prompts
# 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", 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", 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
@@ -2064,7 +2144,7 @@ class TestCBAutoReconnectRefresh:
refresh_event = _threading.Event()
async def _connect_one(name: str, _cfg: dict[str, Any]) -> None:
mgr._sessions[name] = new_session
_seed_static_state(mgr, name, session=new_session)
async def _refresh(name: str) -> tuple[list[str], list[str]]:
refresh_event.set()
@@ -2088,7 +2168,7 @@ class TestCBAutoReconnectRefresh:
refresh_started = _threading.Event()
async def _connect_one(name: str, _cfg: dict[str, Any]) -> None:
mgr._sessions[name] = new_session
_seed_static_state(mgr, name, session=new_session)
async def _refresh_failing(name: str) -> tuple[list[str], list[str]]:
refresh_started.set()
+28 -27
View File
@@ -4,6 +4,7 @@ from __future__ import annotations
from typing import Any
from tests.conftest import _seed_static_state
from turnstone.core.mcp_client import MCPClientManager
# ---------------------------------------------------------------------------
@@ -105,14 +106,18 @@ class TestRemoveServerSync:
"""remove_server_sync cleans up all per-server state dicts."""
mgr = MCPClientManager({"test": {"command": "echo"}})
# Simulate state as if the server was connected
mgr._per_server_tools["test"] = [_fake_openai_tool()]
mgr._per_server_resources["test"] = [_fake_resource_dict()]
mgr._per_server_prompts["test"] = [_fake_prompt_dict()]
mgr._supports_list_changed["test"] = True
mgr._supports_resources["test"] = True
mgr._supports_resource_list_changed["test"] = True
mgr._supports_prompts["test"] = True
mgr._supports_prompt_list_changed["test"] = True
_seed_static_state(
mgr,
"test",
tools=[_fake_openai_tool()],
resources=[_fake_resource_dict()],
prompts=[_fake_prompt_dict()],
supports_list_changed=True,
supports_resources=True,
supports_resource_list_changed=True,
supports_prompts=True,
supports_prompt_list_changed=True,
)
mgr._rebuild_tools()
mgr._rebuild_resources()
mgr._rebuild_prompts()
@@ -127,14 +132,7 @@ class TestRemoveServerSync:
assert len(mgr.get_tools()) == 0
assert mgr.resource_count == 0
assert mgr.prompt_count == 0
assert "test" not in mgr._per_server_tools
assert "test" not in mgr._per_server_resources
assert "test" not in mgr._per_server_prompts
assert "test" not in mgr._supports_list_changed
assert "test" not in mgr._supports_resources
assert "test" not in mgr._supports_resource_list_changed
assert "test" not in mgr._supports_prompts
assert "test" not in mgr._supports_prompt_list_changed
assert "test" not in mgr._static_servers
def test_removes_config_to_prevent_reconnect(self) -> None:
"""remove_server_sync removes from _server_configs to prevent reconnect."""
@@ -146,8 +144,8 @@ class TestRemoveServerSync:
def test_preserves_other_servers(self) -> None:
"""Removing one server does not affect another server's state."""
mgr = MCPClientManager({"srv_a": {}, "srv_b": {}})
mgr._per_server_tools["srv_a"] = [_fake_openai_tool("mcp__srv_a__foo")]
mgr._per_server_tools["srv_b"] = [_fake_openai_tool("mcp__srv_b__bar")]
_seed_static_state(mgr, "srv_a", tools=[_fake_openai_tool("mcp__srv_a__foo")])
_seed_static_state(mgr, "srv_b", tools=[_fake_openai_tool("mcp__srv_b__bar")])
mgr._rebuild_tools()
assert len(mgr.get_tools()) == 2
@@ -179,13 +177,17 @@ class TestGetServerStatus:
"""Status of a connected server reports correct tool/resource/prompt counts."""
mgr = MCPClientManager({"test": {}})
# Simulate connected state
mgr._sessions["test"] = object() # any truthy value
mgr._per_server_tools["test"] = [
_fake_openai_tool("mcp__test__a"),
_fake_openai_tool("mcp__test__b"),
]
mgr._per_server_resources["test"] = [_fake_resource_dict()]
mgr._per_server_prompts["test"] = [_fake_prompt_dict()]
_seed_static_state(
mgr,
"test",
session=object(), # any truthy value
tools=[
_fake_openai_tool("mcp__test__a"),
_fake_openai_tool("mcp__test__b"),
],
resources=[_fake_resource_dict()],
prompts=[_fake_prompt_dict()],
)
status = mgr.get_server_status("test")
assert status["connected"] is True
@@ -225,8 +227,7 @@ class TestGetAllServerStatus:
def test_mixed_connected_and_disconnected(self) -> None:
"""Status correctly reflects a mix of connected and disconnected servers."""
mgr = MCPClientManager({"up": {}, "down": {}})
mgr._sessions["up"] = object()
mgr._per_server_tools["up"] = [_fake_openai_tool("mcp__up__x")]
_seed_static_state(mgr, "up", session=object(), tools=[_fake_openai_tool("mcp__up__x")])
statuses = mgr.get_all_server_status()
assert statuses["up"]["connected"] is True
+81 -66
View File
@@ -14,6 +14,7 @@ from unittest.mock import AsyncMock, MagicMock
import pytest
from tests.conftest import _seed_static_state
from turnstone.core.mcp_client import MCPClientManager
from turnstone.core.storage._sqlite import SQLiteBackend
@@ -109,13 +110,15 @@ class TestFullLifecycleResourcesPrompts:
def test_rebuild_resources_produces_merged_state(self, mgr: MCPClientManager) -> None:
"""_rebuild_resources merges per-server resources into a unified list."""
mgr._per_server_resources["alpha"] = [
_make_resource("file:///a.txt", "a", "alpha"),
_make_resource("file:///b.txt", "b", "alpha"),
]
mgr._per_server_resources["beta"] = [
_make_resource("file:///c.txt", "c", "beta"),
]
_seed_static_state(
mgr,
"alpha",
resources=[
_make_resource("file:///a.txt", "a", "alpha"),
_make_resource("file:///b.txt", "b", "alpha"),
],
)
_seed_static_state(mgr, "beta", resources=[_make_resource("file:///c.txt", "c", "beta")])
mgr._rebuild_resources()
@@ -130,13 +133,19 @@ class TestFullLifecycleResourcesPrompts:
def test_rebuild_prompts_produces_merged_state(self, mgr: MCPClientManager) -> None:
"""_rebuild_prompts merges per-server prompts into a unified list."""
mgr._per_server_prompts["alpha"] = [
_make_prompt("mcp__alpha__greet", "greet", "alpha", "Say hello"),
]
mgr._per_server_prompts["beta"] = [
_make_prompt("mcp__beta__summarize", "summarize", "beta", "Summarize text"),
_make_prompt("mcp__beta__translate", "translate", "beta", "Translate text"),
]
_seed_static_state(
mgr,
"alpha",
prompts=[_make_prompt("mcp__alpha__greet", "greet", "alpha", "Say hello")],
)
_seed_static_state(
mgr,
"beta",
prompts=[
_make_prompt("mcp__beta__summarize", "summarize", "beta", "Summarize text"),
_make_prompt("mcp__beta__translate", "translate", "beta", "Translate text"),
],
)
mgr._rebuild_prompts()
@@ -164,10 +173,12 @@ class TestFullLifecycleResourcesPrompts:
try:
# Populate session and resource map
session = _make_mock_session()
mgr._sessions["alpha"] = session
mgr._per_server_resources["alpha"] = [
_make_resource("file:///readme.md", "readme", "alpha"),
]
_seed_static_state(
mgr,
"alpha",
session=session,
resources=[_make_resource("file:///readme.md", "readme", "alpha")],
)
mgr._rebuild_resources()
result = mgr.read_resource_sync("file:///readme.md", timeout=5)
@@ -194,18 +205,22 @@ class TestFullLifecycleResourcesPrompts:
try:
session = _make_mock_session()
mgr._sessions["alpha"] = session
# Register a template resource (no concrete resources)
mgr._per_server_resources["alpha"] = [
{
"uri": "db://tables/{table}/rows/{id}",
"name": "row",
"description": "Fetch a row",
"mimeType": "application/json",
"server": "alpha",
"template": True,
},
]
_seed_static_state(
mgr,
"alpha",
session=session,
resources=[
{
"uri": "db://tables/{table}/rows/{id}",
"name": "row",
"description": "Fetch a row",
"mimeType": "application/json",
"server": "alpha",
"template": True,
},
],
)
mgr._rebuild_resources()
# Template should not be in _resource_map
@@ -230,10 +245,12 @@ class TestFullLifecycleResourcesPrompts:
try:
session = _make_mock_session()
mgr._sessions["alpha"] = session
mgr._per_server_prompts["alpha"] = [
_make_prompt("mcp__alpha__greet", "greet", "alpha", "Say hello"),
]
_seed_static_state(
mgr,
"alpha",
session=session,
prompts=[_make_prompt("mcp__alpha__greet", "greet", "alpha", "Say hello")],
)
mgr._rebuild_prompts()
messages = mgr.get_prompt_sync(
@@ -314,40 +331,42 @@ class TestFullLifecycleResourcesPrompts:
def test_shutdown_clears_all_state(self, mgr: MCPClientManager) -> None:
"""shutdown() clears sessions, tools, resources, prompts, and listeners."""
# Populate state
mgr._sessions["alpha"] = MagicMock()
mgr._per_server_tools["alpha"] = [
{
"type": "function",
"function": {
"name": "mcp__alpha__search",
"description": "Search",
"parameters": {},
_seed_static_state(
mgr,
"alpha",
session=MagicMock(),
tools=[
{
"type": "function",
"function": {
"name": "mcp__alpha__search",
"description": "Search",
"parameters": {},
},
}
],
resources=[
_make_resource("file:///a.txt", "a", "alpha"),
{
"uri": "db://tables/{table}",
"name": "table",
"description": "",
"mimeType": "",
"server": "alpha",
"template": True,
},
}
]
],
prompts=[_make_prompt("mcp__alpha__greet", "greet", "alpha")],
)
mgr._rebuild_tools()
mgr._per_server_resources["alpha"] = [
_make_resource("file:///a.txt", "a", "alpha"),
{
"uri": "db://tables/{table}",
"name": "table",
"description": "",
"mimeType": "",
"server": "alpha",
"template": True,
},
]
mgr._rebuild_resources()
mgr._per_server_prompts["alpha"] = [
_make_prompt("mcp__alpha__greet", "greet", "alpha"),
]
mgr._rebuild_prompts()
mgr._listeners.append(lambda: None)
mgr._resource_listeners.append(lambda: None)
mgr._prompt_listeners.append(lambda: None)
# Verify populated
assert len(mgr._sessions) == 1
assert len(mgr._static_servers) == 1
assert len(mgr._tools) == 1
assert len(mgr._resources) == 2 # 1 concrete + 1 template
assert len(mgr._template_prefixes) == 1
@@ -355,7 +374,7 @@ class TestFullLifecycleResourcesPrompts:
mgr.shutdown()
assert len(mgr._sessions) == 0
assert len(mgr._static_servers) == 0
assert len(mgr._tools) == 0
assert len(mgr._tool_map) == 0
assert len(mgr._resources) == 0
@@ -376,19 +395,15 @@ class TestFullLifecycleResourcesPrompts:
mgr.add_resource_listener(lambda: resource_fired.append(1))
mgr.add_prompt_listener(lambda: prompt_fired.append(1))
mgr._per_server_tools["alpha"] = []
_seed_static_state(mgr, "alpha", tools=[])
mgr._rebuild_tools()
assert len(tool_fired) == 1
mgr._per_server_resources["alpha"] = [
_make_resource("file:///x.txt", "x", "alpha"),
]
_seed_static_state(mgr, "alpha", resources=[_make_resource("file:///x.txt", "x", "alpha")])
mgr._rebuild_resources()
assert len(resource_fired) == 1
mgr._per_server_prompts["alpha"] = [
_make_prompt("mcp__alpha__p1", "p1", "alpha"),
]
_seed_static_state(mgr, "alpha", prompts=[_make_prompt("mcp__alpha__p1", "p1", "alpha")])
mgr._rebuild_prompts()
assert len(prompt_fired) == 1
+238 -147
View File
@@ -25,6 +25,7 @@ import threading
import time
import uuid
from contextlib import AsyncExitStack
from dataclasses import dataclass, field
from pathlib import Path
from typing import TYPE_CHECKING, Any
@@ -68,6 +69,57 @@ def _mcp_to_openai(server_name: str, tool: Any) -> dict[str, Any]:
}
# ---------------------------------------------------------------------------
# Per-server state containers
# ---------------------------------------------------------------------------
@dataclass
class StaticServerState:
"""Per-server state for auth_type ∈ {none, static}. Name-keyed only.
Phase 5 introduces PoolEntryState as the (user, server)-keyed sibling
for auth_type=oauth_user. Together with the typed map declarations
(dict[str, StaticServerState] vs dict[tuple[str, str], PoolEntryState]),
this makes accidental cross-keying lookups easier to catch in review
and rejected by mypy.
"""
name: str
session: Any | None = None
stack: AsyncExitStack | None = None
streams: tuple[Any, Any] | None = None
tools: list[dict[str, Any]] = field(default_factory=list)
resources: list[dict[str, Any]] = field(default_factory=list)
prompts: list[dict[str, Any]] = field(default_factory=list)
supports_list_changed: bool = False
supports_resources: bool = False
supports_prompts: bool = False
supports_resource_list_changed: bool = False
supports_prompt_list_changed: bool = False
@dataclass
class PoolEntryState:
"""Per-(user, server) state for auth_type = oauth_user.
Defined for Phase 5 use; not instantiated anywhere in Phase 0.
open_lock has no default — RFC §2.0 invariant 2 forbids allocating an
asyncio.Lock outside the mcp-loop. Phase 5 allocates lazily inside
connect coroutines.
"""
key: tuple[str, str] # (user_id, server_name)
open_lock: asyncio.Lock
session: Any | None = None
stack: AsyncExitStack | None = None
streams: tuple[Any, Any] | None = None
tools: list[dict[str, Any]] = field(default_factory=list)
resources: list[dict[str, Any]] = field(default_factory=list)
prompts: list[dict[str, Any]] = field(default_factory=list)
last_used: float = 0.0
# ---------------------------------------------------------------------------
# Client manager
# ---------------------------------------------------------------------------
@@ -88,9 +140,13 @@ class MCPClientManager:
self._loop: asyncio.AbstractEventLoop | None = None
self._thread: threading.Thread | None = None
self._exit_stack: AsyncExitStack | None = None
self._per_server_stacks: dict[str, AsyncExitStack] = {}
self._sessions: dict[str, Any] = {}
# Per-server state for auth_type ∈ {none, static}. Each entry holds
# session/stack/streams/catalog/capability flags for one name-keyed
# connection. Phase 5 introduces a sibling pool-entry map for
# auth_type=oauth_user; static entries always live here.
self._static_servers: dict[str, StaticServerState] = {}
self._tools: list[dict[str, Any]] = []
# prefixed_name -> (server_name, original_tool_name)
self._tool_map: dict[str, tuple[str, str]] = {}
@@ -104,30 +160,19 @@ class MCPClientManager:
self._last_error: dict[str, str] = {}
self._MAX_ERROR_LEN = 256
# Per-server tool storage for surgical refresh
self._per_server_tools: dict[str, list[dict[str, Any]]] = {}
# Tracks which servers support push notifications
self._supports_list_changed: dict[str, bool] = {}
# Listener infrastructure (tool-change callbacks for ChatSession)
self._listeners: list[Callable[[], None]] = []
self._listeners_lock = threading.Lock()
# Resources — parallel to tools
self._per_server_resources: dict[str, list[dict[str, Any]]] = {}
# Merged resource catalog
self._resources: list[dict[str, Any]] = []
self._resource_map: dict[str, tuple[str, str]] = {} # uri → (server, uri)
self._supports_resources: dict[str, bool] = {} # server has resources capability
self._supports_resource_list_changed: dict[str, bool] = {}
self._resource_listeners: list[Callable[[], None]] = []
self._resource_listeners_lock = threading.Lock()
# Prompts — parallel to tools
self._per_server_prompts: dict[str, list[dict[str, Any]]] = {}
# Merged prompt catalog
self._prompts: list[dict[str, Any]] = []
self._prompt_map: dict[str, tuple[str, str]] = {} # prefixed → (server, original)
self._supports_prompts: dict[str, bool] = {} # server has prompts capability
self._supports_prompt_list_changed: dict[str, bool] = {}
self._prompt_listeners: list[Callable[[], None]] = []
self._prompt_listeners_lock = threading.Lock()
@@ -143,13 +188,21 @@ class MCPClientManager:
self._circuit_open_until: dict[str, float] = {} # monotonic timestamp
self._circuit_trip_count: dict[str, int] = {} # backoff exponent
# Safe transport stream refs (pre-close before stack teardown to avoid
# the anyio cancel-scope CPU busy-loop — MCP SDK #2147)
self._server_streams: dict[str, tuple[Any, Any]] = {}
# Notification debounce (per-server)
self._last_notification_refresh: dict[str, float] = {}
def _ensure_static_state(self, name: str) -> StaticServerState:
"""Get or create the StaticServerState for ``name``.
Returns an empty state on first access; subsequent fields are populated
as connect proceeds.
"""
state = self._static_servers.get(name)
if state is None:
state = StaticServerState(name=name)
self._static_servers[name] = state
return state
# -- lifecycle -----------------------------------------------------------
def start(self) -> None:
@@ -259,34 +312,43 @@ class MCPClientManager:
# -- safe transport helpers ------------------------------------------------
async def _pre_close_streams(self, name: str) -> None:
async def _pre_close_streams(self, key: str) -> None:
"""Close MCP transport streams before stack teardown.
Pre-closing unblocks anyio transport tasks stuck on zero-buffer
``send()`` calls, preventing the CPU busy-loop from SDK #2147.
"""
streams = self._server_streams.pop(name, None)
if streams:
for s in streams:
with contextlib.suppress(Exception):
await s.aclose()
async def _tcp_probe(self, name: str, url: str) -> None:
Parameter is ``str`` today; Phase 5 widens to ``str | tuple[str, str]``
once ``PoolEntryState`` is wired.
"""
state = self._static_servers.get(key)
if state is None or state.streams is None:
return
streams = state.streams
state.streams = None # take-and-clear pattern
for s in streams:
with contextlib.suppress(Exception):
await s.aclose()
async def _tcp_probe(self, key: str, url: str) -> None:
"""Fast TCP connect check before entering the MCP transport context.
Fails fast when the server is unreachable, avoiding the anyio
cancel-scope orphan bug that causes 100% CPU spin.
Parameter is ``str`` today; Phase 5 widens to ``str | tuple[str, str]``
once ``PoolEntryState`` is wired.
"""
from urllib.parse import urlparse
parsed = urlparse(url)
host = parsed.hostname
if not host:
raise ConnectionError(f"MCP server '{name}' has invalid URL (no hostname): {url}")
raise ConnectionError(f"MCP server '{key}' has invalid URL (no hostname): {url}")
try:
port = parsed.port or (443 if parsed.scheme == "https" else 80)
except ValueError:
raise ConnectionError(f"MCP server '{name}' has invalid port in URL: {url}") from None
raise ConnectionError(f"MCP server '{key}' has invalid port in URL: {url}") from None
try:
_, writer = await asyncio.wait_for(
asyncio.open_connection(host, port),
@@ -296,7 +358,7 @@ class MCPClientManager:
await writer.wait_closed()
except (TimeoutError, OSError) as exc:
raise ConnectionError(
f"MCP server '{name}' unreachable at {host}:{port}: {exc}"
f"MCP server '{key}' unreachable at {host}:{port}: {exc}"
) from None
@staticmethod
@@ -319,14 +381,21 @@ class MCPClientManager:
log.error("MCP server name '%s' contains '__' (reserved delimiter), skipping", name)
return
# Operate on a single state object throughout: get-or-create up front
# so the stale-entry guard and the post-handshake field assignments
# touch the same instance (PR #296 invariant 5: identity stability).
state = self._ensure_static_state(name)
# Guard: tear down stale session/stack so we don't leak. Checks both
# _sessions and _per_server_stacks because transport errors in the sync
# dispatch methods evict the session but leave the stack behind.
if name in self._sessions or name in self._per_server_stacks:
self._sessions.pop(name, None)
# session and stack because transport errors in the sync dispatch
# methods evict the session but leave the stack behind. On a brand
# new entry both fields are None, so this branch is skipped.
if state.session is not None or state.stack is not None:
state.session = None
await self._pre_close_streams(name)
old_stack = self._per_server_stacks.pop(name, None)
if old_stack:
old_stack = state.stack
state.stack = None
if old_stack is not None:
await self._safe_close_stack(old_stack)
# Per-server exit stack for clean per-server lifecycle management
@@ -351,7 +420,7 @@ class MCPClientManager:
)
# Stash stream refs so _pre_close_streams can unblock anyio
# transport tasks before the cancel scope fires (SDK #2147).
self._server_streams[name] = (read, write)
state.streams = (read, write)
else:
# Default: stdio transport
command = cfg.get("command", "")
@@ -368,7 +437,7 @@ class MCPClientManager:
env=env,
)
read, write = await stack.enter_async_context(stdio_client(params))
self._server_streams[name] = (read, write)
state.streams = (read, write)
except asyncio.CancelledError:
# Stray CancelledError from broken anyio cancel scope -- treat as
# connection failure. But if the task is genuinely being cancelled
@@ -441,11 +510,11 @@ class MCPClientManager:
await self._safe_close_stack(stack)
raise
self._per_server_stacks[name] = stack
state.stack = stack
try:
await asyncio.wait_for(session.initialize(), timeout=self._CONNECT_TIMEOUT)
except asyncio.CancelledError:
self._per_server_stacks.pop(name, None)
state.stack = None
task = asyncio.current_task()
if task is not None and task.cancelling():
await self._pre_close_streams(name)
@@ -455,32 +524,30 @@ class MCPClientManager:
await self._safe_close_stack(stack)
raise TimeoutError(f"MCP handshake failed for '{name}'") from None
except TimeoutError:
self._per_server_stacks.pop(name, None)
state.stack = None
await self._pre_close_streams(name)
await self._safe_close_stack(stack)
raise TimeoutError(f"MCP handshake timed out after {self._CONNECT_TIMEOUT}s") from None
except Exception:
self._per_server_stacks.pop(name, None)
state.stack = None
await self._pre_close_streams(name)
await self._safe_close_stack(stack)
raise
self._sessions[name] = session
state.session = session
# Check push notification support for each capability
caps = session.get_server_capabilities()
tools_cap = getattr(caps, "tools", None) if caps else None
self._supports_list_changed[name] = bool(getattr(tools_cap, "listChanged", False))
state.supports_list_changed = bool(getattr(tools_cap, "listChanged", False))
resources_cap = getattr(caps, "resources", None) if caps else None
self._supports_resources[name] = resources_cap is not None
self._supports_resource_list_changed[name] = bool(
getattr(resources_cap, "listChanged", False)
)
state.supports_resources = resources_cap is not None
state.supports_resource_list_changed = bool(getattr(resources_cap, "listChanged", False))
prompts_cap = getattr(caps, "prompts", None) if caps else None
self._supports_prompts[name] = prompts_cap is not None
self._supports_prompt_list_changed[name] = bool(getattr(prompts_cap, "listChanged", False))
state.supports_prompts = prompts_cap is not None
state.supports_prompt_list_changed = bool(getattr(prompts_cap, "listChanged", False))
# Discover tools
result = await session.list_tools()
@@ -488,7 +555,7 @@ class MCPClientManager:
for tool in result.tools:
server_tools.append(_mcp_to_openai(name, tool))
self._per_server_tools[name] = server_tools
state.tools = server_tools
self._rebuild_tools()
# Discover resources
@@ -521,7 +588,7 @@ class MCPClientManager:
}
)
resource_count = len(server_resources)
self._per_server_resources[name] = server_resources
state.resources = server_resources
self._rebuild_resources()
# Discover prompts
@@ -547,15 +614,15 @@ class MCPClientManager:
}
)
prompt_count = len(server_prompts)
self._per_server_prompts[name] = server_prompts
state.prompts = server_prompts
self._rebuild_prompts()
push_parts: list[str] = []
if self._supports_list_changed[name]:
if state.supports_list_changed:
push_parts.append("tools")
if self._supports_resource_list_changed[name]:
if state.supports_resource_list_changed:
push_parts.append("resources")
if self._supports_prompt_list_changed[name]:
if state.supports_prompt_list_changed:
push_parts.append("prompts")
push_status = f" (push: {','.join(push_parts)})" if push_parts else ""
log.info(
@@ -586,8 +653,8 @@ class MCPClientManager:
"""
new_tools: list[dict[str, Any]] = []
new_map: dict[str, tuple[str, str]] = {}
for srv_name, srv_tools in self._per_server_tools.items():
for tool in srv_tools:
for srv_name, srv_state in self._static_servers.items():
for tool in srv_state.tools:
prefixed: str = tool["function"]["name"]
new_tools.append(tool)
# Extract original name from the mcp__server__original pattern
@@ -599,17 +666,21 @@ class MCPClientManager:
async def _refresh_server_tools(self, name: str) -> tuple[list[str], list[str]]:
"""Re-fetch tools for one server. Returns ``(added, removed)`` names."""
session = self._sessions.get(name)
if session is None:
state = self._static_servers.get(name)
if state is None or state.session is None:
raise RuntimeError(f"MCP server '{name}' is not connected")
# Capture session locally — a concurrent transport-error eviction in
# call_tool_sync can clear state.session; reads after an await would
# raise AttributeError without this snapshot.
session = state.session
old_names = {t["function"]["name"] for t in self._per_server_tools.get(name, [])}
old_names = {t["function"]["name"] for t in state.tools}
result = await session.list_tools()
server_tools = [_mcp_to_openai(name, tool) for tool in result.tools]
new_names = {t["function"]["name"] for t in server_tools}
self._per_server_tools[name] = server_tools
state.tools = server_tools
self._rebuild_tools()
added = sorted(new_names - old_names)
@@ -651,16 +722,18 @@ class MCPClientManager:
for name in targets:
try:
if name not in self._sessions:
state = self._static_servers.get(name)
if state is None or state.session is None:
# Attempt reconnect
cfg = self._server_configs.get(name)
if cfg:
log.info("Reconnecting MCP server '%s'", name)
await self._connect_one(name, cfg)
self._cb_record_success(name)
new_names = [
t["function"]["name"] for t in self._per_server_tools.get(name, [])
]
post = self._static_servers.get(name)
new_names = (
[t["function"]["name"] for t in post.tools] if post is not None else []
)
results[name] = (new_names, [])
continue
added, removed = await self._refresh_server(name)
@@ -703,8 +776,8 @@ class MCPClientManager:
"""
new_resources: list[dict[str, Any]] = []
new_map: dict[str, tuple[str, str]] = {}
for srv_name, srv_resources in self._per_server_resources.items():
for res in srv_resources:
for srv_name, srv_state in self._static_servers.items():
for res in srv_state.resources:
uri: str = res["uri"]
new_resources.append(res)
if res.get("template"):
@@ -719,8 +792,8 @@ class MCPClientManager:
new_map[uri] = (srv_name, uri)
# Build template prefix map for URI expansion fallback
new_prefixes: dict[str, tuple[str, str]] = {}
for srv_name, srv_resources in self._per_server_resources.items():
for res in srv_resources:
for srv_name, srv_state in self._static_servers.items():
for res in srv_state.resources:
if res.get("template"):
tmpl_uri = res["uri"]
brace = tmpl_uri.find("{")
@@ -755,11 +828,15 @@ class MCPClientManager:
async def _refresh_server_resources(self, name: str) -> None:
"""Re-fetch resources for one server."""
if not self._supports_resources.get(name, False):
state = self._static_servers.get(name)
if state is None or not state.supports_resources:
return
session = self._sessions.get(name)
if session is None:
if state.session is None:
return
# Capture session locally — a concurrent transport-error eviction in
# call_tool_sync can clear state.session between awaits, which would
# turn the second list_resource_templates() call into AttributeError.
session = state.session
server_resources: list[dict[str, Any]] = []
res_result = await session.list_resources()
@@ -786,7 +863,7 @@ class MCPClientManager:
}
)
self._per_server_resources[name] = server_resources
state.resources = server_resources
self._rebuild_resources()
# -- prompt refresh ------------------------------------------------------
@@ -798,8 +875,8 @@ class MCPClientManager:
"""
new_prompts: list[dict[str, Any]] = []
new_map: dict[str, tuple[str, str]] = {}
for srv_name, srv_prompts in self._per_server_prompts.items():
for prompt in srv_prompts:
for srv_name, srv_state in self._static_servers.items():
for prompt in srv_state.prompts:
prefixed: str = prompt["name"]
new_prompts.append(prompt)
new_map[prefixed] = (srv_name, prompt["original_name"])
@@ -809,11 +886,15 @@ class MCPClientManager:
async def _refresh_server_prompts(self, name: str) -> None:
"""Re-fetch prompts for one server."""
if not self._supports_prompts.get(name, False):
state = self._static_servers.get(name)
if state is None or not state.supports_prompts:
return
session = self._sessions.get(name)
if session is None:
if state.session is None:
return
# Capture session locally — see _refresh_server_resources for the
# concurrent-eviction race this guards against. Single-await today,
# multi-await tomorrow; consistent capture-once idiom.
session = state.session
server_prompts: list[dict[str, Any]] = []
prompt_result = await session.list_prompts()
@@ -835,7 +916,7 @@ class MCPClientManager:
}
)
self._per_server_prompts[name] = server_prompts
state.prompts = server_prompts
self._rebuild_prompts()
# Sync discovered prompts into governance storage
@@ -1030,14 +1111,15 @@ class MCPClientManager:
def shutdown(self) -> None:
"""Close all MCP sessions and stop the background loop."""
# Close all per-server stacks (transports + sessions)
if self._loop and self._per_server_stacks:
if self._loop and self._static_servers:
async def _close_all_stacks() -> None:
# Pre-close streams to prevent anyio CPU busy-loop during teardown
for srv_name in list(self._server_streams):
for srv_name in list(self._static_servers):
await self._pre_close_streams(srv_name)
for stack in self._per_server_stacks.values():
await self._safe_close_stack(stack)
for srv_state in self._static_servers.values():
if srv_state.stack is not None:
await self._safe_close_stack(srv_state.stack)
future = asyncio.run_coroutine_threadsafe(_close_all_stacks(), self._loop)
try:
@@ -1059,24 +1141,15 @@ class MCPClientManager:
self._thread.join(timeout=5)
# Clear all state
self._sessions.clear()
self._per_server_stacks.clear()
self._static_servers.clear()
self._db_managed.clear()
self._tools = []
self._tool_map = {}
self._per_server_tools.clear()
self._supports_list_changed.clear()
self._resources = []
self._resource_map = {}
self._template_prefixes = {}
self._per_server_resources.clear()
self._supports_resources.clear()
self._supports_resource_list_changed.clear()
self._prompts = []
self._prompt_map = {}
self._per_server_prompts.clear()
self._supports_prompts.clear()
self._supports_prompt_list_changed.clear()
# Clear listener lists to release callback references
self._listeners.clear()
self._resource_listeners.clear()
@@ -1085,7 +1158,6 @@ class MCPClientManager:
self._consecutive_failures.clear()
self._circuit_open_until.clear()
self._circuit_trip_count.clear()
self._server_streams.clear()
self._last_notification_refresh.clear()
log.info("MCP client shut down")
@@ -1125,11 +1197,12 @@ class MCPClientManager:
self._server_configs.pop(name, None)
return {"connected": False, "tools": 0, "resources": 0, "prompts": 0, "error": str(exc)}
state = self._static_servers.get(name)
return {
"connected": name in self._sessions,
"tools": len(self._per_server_tools.get(name, [])),
"resources": len(self._per_server_resources.get(name, [])),
"prompts": len(self._per_server_prompts.get(name, [])),
"connected": state is not None and state.session is not None,
"tools": len(state.tools) if state else 0,
"resources": len(state.resources) if state else 0,
"prompts": len(state.prompts) if state else 0,
"error": "",
}
@@ -1162,20 +1235,25 @@ class MCPClientManager:
async def _reconnect() -> None:
self._cb_clear(name)
self._sessions.pop(name, None)
await self._pre_close_streams(name)
stack = self._per_server_stacks.pop(name, None)
if stack is not None:
await self._safe_close_stack(stack)
state = self._static_servers.get(name)
if state is not None:
state.session = None
await self._pre_close_streams(name)
old_stack = state.stack
state.stack = None
if old_stack is not None:
await self._safe_close_stack(old_stack)
try:
await self._connect_one(name, cfg)
except Exception:
# Connect failed mid-reconnect — drop the stale per-server
# catalog so the merged tool/resource/prompt maps don't keep
# advertising entries with no live session behind them.
self._per_server_tools.pop(name, None)
self._per_server_resources.pop(name, None)
self._per_server_prompts.pop(name, None)
fail_state = self._static_servers.get(name)
if fail_state is not None:
fail_state.tools = []
fail_state.resources = []
fail_state.prompts = []
self._rebuild_tools()
self._rebuild_resources()
self._rebuild_prompts()
@@ -1196,11 +1274,12 @@ class MCPClientManager:
except Exception as exc:
return {"connected": False, "tools": 0, "resources": 0, "prompts": 0, "error": str(exc)}
state = self._static_servers.get(name)
return {
"connected": name in self._sessions,
"tools": len(self._per_server_tools.get(name, [])),
"resources": len(self._per_server_resources.get(name, [])),
"prompts": len(self._per_server_prompts.get(name, [])),
"connected": state is not None and state.session is not None,
"tools": len(state.tools) if state else 0,
"resources": len(state.resources) if state else 0,
"prompts": len(state.prompts) if state else 0,
"error": "",
}
@@ -1212,7 +1291,8 @@ class MCPClientManager:
Returns True if the server was connected and successfully removed.
"""
was_connected = name in self._sessions
existing = self._static_servers.get(name)
was_connected = existing is not None and existing.session is not None
# Remove from config to prevent reconnection
self._server_configs.pop(name, None)
@@ -1221,20 +1301,16 @@ class MCPClientManager:
async def _remove() -> None:
# Close session + transport via per-server stack
self._sessions.pop(name, None)
await self._pre_close_streams(name)
stack = self._per_server_stacks.pop(name, None)
if stack is not None:
await self._safe_close_stack(stack)
state = self._static_servers.get(name)
if state is not None:
state.session = None
await self._pre_close_streams(name)
stack = state.stack
state.stack = None
if stack is not None:
await self._safe_close_stack(stack)
# Clean up per-server state (on the event loop thread)
self._per_server_tools.pop(name, None)
self._per_server_resources.pop(name, None)
self._per_server_prompts.pop(name, None)
self._supports_list_changed.pop(name, None)
self._supports_resources.pop(name, None)
self._supports_resource_list_changed.pop(name, None)
self._supports_prompts.pop(name, None)
self._supports_prompt_list_changed.pop(name, None)
self._static_servers.pop(name, None)
self._last_error.pop(name, None)
self._last_notification_refresh.pop(name, None)
self._cb_clear(name)
@@ -1250,16 +1326,7 @@ class MCPClientManager:
log.warning("Error removing MCP server '%s'", name, exc_info=True)
else:
# No event loop (tests / pre-start) — mutate directly
self._sessions.pop(name, None)
self._server_streams.pop(name, None)
self._per_server_tools.pop(name, None)
self._per_server_resources.pop(name, None)
self._per_server_prompts.pop(name, None)
self._supports_list_changed.pop(name, None)
self._supports_resources.pop(name, None)
self._supports_resource_list_changed.pop(name, None)
self._supports_prompts.pop(name, None)
self._supports_prompt_list_changed.pop(name, None)
self._static_servers.pop(name, None)
self._last_error.pop(name, None)
self._last_notification_refresh.pop(name, None)
self._cb_clear(name)
@@ -1283,16 +1350,21 @@ class MCPClientManager:
def get_server_status(self, name: str) -> dict[str, Any]:
"""Return live status for a single server, including config details."""
connected = name in self._sessions
state = self._static_servers.get(name)
connected = state is not None and state.session is not None
cfg = self._server_configs.get(name, {})
transport = cfg.get("type", "stdio")
cb_deadline = self._circuit_open_until.get(name)
cb_open = cb_deadline is not None and time.monotonic() < cb_deadline
# Inline predicate (instead of reusing ``connected``) so mypy narrows
# ``state`` for the attribute reads — a separate boolean wouldn't.
return {
"connected": connected,
"tools": len(self._per_server_tools.get(name, [])) if connected else 0,
"resources": len(self._per_server_resources.get(name, [])) if connected else 0,
"prompts": len(self._per_server_prompts.get(name, [])) if connected else 0,
"tools": len(state.tools) if state is not None and state.session is not None else 0,
"resources": (
len(state.resources) if state is not None and state.session is not None else 0
),
"prompts": len(state.prompts) if state is not None and state.session is not None else 0,
"error": self._last_error.get(name, ""),
"transport": transport,
"command": cfg.get("command", "") if transport == "stdio" else "",
@@ -1414,7 +1486,7 @@ class MCPClientManager:
@property
def server_count(self) -> int:
return len(self._sessions)
return sum(1 for s in self._static_servers.values() if s.session is not None)
@property
def error_count(self) -> int:
@@ -1471,7 +1543,8 @@ class MCPClientManager:
except Exception as exc:
self._cb_record_failure(server_name)
raise RuntimeError(f"MCP server '{server_name}' reconnect failed: {exc}") from None
session = self._sessions.get(server_name)
state = self._static_servers.get(server_name)
session = state.session if state is not None else None
if session is None:
self._cb_record_failure(server_name)
raise RuntimeError(f"MCP server '{server_name}' reconnect produced no session")
@@ -1511,7 +1584,8 @@ class MCPClientManager:
self._cb_gate(server_name)
session = self._sessions.get(server_name)
state = self._static_servers.get(server_name)
session = state.session if state is not None else None
if session is None:
session = self._cb_auto_reconnect(server_name)
assert self._loop is not None
@@ -1531,7 +1605,12 @@ class MCPClientManager:
if not isinstance(exc, McpError):
self._cb_record_failure(server_name)
if isinstance(exc, (BrokenPipeError, ConnectionResetError, EOFError)):
self._sessions.pop(server_name, None)
# Evict the session only — leave stack/streams behind so the
# stale-session-and-stack guard in _connect_one cleans them up
# on the next connect attempt.
evict = self._static_servers.get(server_name)
if evict is not None:
evict.session = None
raise
self._cb_record_success(server_name)
@@ -1586,7 +1665,8 @@ class MCPClientManager:
self._cb_gate(server_name)
session = self._sessions.get(server_name)
state = self._static_servers.get(server_name)
session = state.session if state is not None else None
if session is None:
session = self._cb_auto_reconnect(server_name)
assert self._loop is not None
@@ -1602,7 +1682,12 @@ class MCPClientManager:
if not isinstance(exc, McpError):
self._cb_record_failure(server_name)
if isinstance(exc, (BrokenPipeError, ConnectionResetError, EOFError)):
self._sessions.pop(server_name, None)
# Evict the session only — leave stack/streams behind so the
# stale-session-and-stack guard in _connect_one cleans them up
# on the next connect attempt.
evict = self._static_servers.get(server_name)
if evict is not None:
evict.session = None
raise
self._cb_record_success(server_name)
@@ -1636,7 +1721,8 @@ class MCPClientManager:
self._cb_gate(server_name)
session = self._sessions.get(server_name)
state = self._static_servers.get(server_name)
session = state.session if state is not None else None
if session is None:
session = self._cb_auto_reconnect(server_name)
assert self._loop is not None
@@ -1654,7 +1740,12 @@ class MCPClientManager:
if not isinstance(exc, McpError):
self._cb_record_failure(server_name)
if isinstance(exc, (BrokenPipeError, ConnectionResetError, EOFError)):
self._sessions.pop(server_name, None)
# Evict the session only — leave stack/streams behind so the
# stale-session-and-stack guard in _connect_one cleans them up
# on the next connect attempt.
evict = self._static_servers.get(server_name)
if evict is not None:
evict.session = None
raise
self._cb_record_success(server_name)