mirror of
https://github.com/turnstonelabs/turnstone.git
synced 2026-08-12 23:12:23 -06:00
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:
@@ -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
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user