diff --git a/tests/conftest.py b/tests/conftest.py index db828a71..9ba652a2 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -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. diff --git a/tests/test_mcp_client.py b/tests/test_mcp_client.py index 5e607ead..b462ec31 100644 --- a/tests/test_mcp_client.py +++ b/tests/test_mcp_client.py @@ -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() diff --git a/tests/test_mcp_hot_reload.py b/tests/test_mcp_hot_reload.py index c75ced37..6aaa9b29 100644 --- a/tests/test_mcp_hot_reload.py +++ b/tests/test_mcp_hot_reload.py @@ -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 diff --git a/tests/test_mcp_integration.py b/tests/test_mcp_integration.py index ffbb3d36..4460e8b0 100644 --- a/tests/test_mcp_integration.py +++ b/tests/test_mcp_integration.py @@ -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 diff --git a/turnstone/core/mcp_client.py b/turnstone/core/mcp_client.py index 01b98a5c..40ad9c90 100644 --- a/turnstone/core/mcp_client.py +++ b/turnstone/core/mcp_client.py @@ -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)