fix(mcp): only keep upstream sessions for gateway-issued session ids and end them on dead streams

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Devin AI 2026-09-23 23:33:08 +00:00
parent 595ced358a
commit 8b59890789
5 changed files with 42 additions and 2 deletions

View file

@ -1191,6 +1191,15 @@ class MCPClient:
_PendingOperation: TypeAlias = "tuple[Callable[[ClientSession], Awaitable[object]], asyncio.Future[object]]"
_MAX_PENDING_OPERATIONS: Final = 64
_SESSION_ENDING_ERRORS: Final = (
ValueError,
httpx2.HTTPError,
OSError,
MCPError,
anyio.BrokenResourceError,
anyio.ClosedResourceError,
anyio.EndOfStream,
)
class UpstreamSessionClosedError(RuntimeError):
@ -1233,7 +1242,7 @@ class PersistentMCPSession:
future.set_exception(RuntimeError("upstream MCP operation was cancelled"))
else:
future.set_result(outcome)
if isinstance(outcome, (ValueError, httpx2.HTTPError, OSError, MCPError)):
if isinstance(outcome, _SESSION_ENDING_ERRORS):
return
self._active = None

View file

@ -1869,6 +1869,7 @@ class MCPServerManager:
)
self.registry: dict[str, MCPServer] = {}
self._upstream_sessions: dict[tuple[str, str, str], PersistentMCPSession] = {} # mutable-ok: session registry
self._live_gateway_sessions: frozenset[str] = frozenset()
self._openapi_health_probes: Callable[[str], _OpenAPIHealthProbe] = lru_cache(maxsize=128)(_OpenAPIHealthProbe)
self.config_mcp_servers: dict[str, MCPServer] = {}
"""
@ -5798,7 +5799,7 @@ class MCPServerManager:
),
None,
)
if gateway_session_id is None or mcp_server.transport == MCPTransport.stdio:
if gateway_session_id not in self._live_gateway_sessions or mcp_server.transport == MCPTransport.stdio:
return None
key: Final = (gateway_session_id, mcp_server.server_id, await client.discovery_auth_fingerprint())
existing: Final = self._upstream_sessions.get(key)
@ -5808,7 +5809,11 @@ class MCPServerManager:
self._upstream_sessions[key] = opened
return opened
def track_gateway_session(self, gateway_session_id: str) -> None:
self._live_gateway_sessions = self._live_gateway_sessions | frozenset((gateway_session_id,))
def release_upstream_sessions(self, gateway_session_id: str) -> None:
self._live_gateway_sessions = self._live_gateway_sessions - frozenset((gateway_session_id,))
for key in tuple(key for key in self._upstream_sessions if key[0] == gateway_session_id):
self._upstream_sessions.pop(key).close()

View file

@ -2569,6 +2569,7 @@ if MCP_AVAILABLE:
_stateful_session_auth_contexts[session_id] = auth_user
_stateful_session_auth_context_last_seen[session_id] = time.monotonic()
_stateful_session_owners[session_id] = owner_fingerprint
operations.global_mcp_server_manager.track_gateway_session(session_id)
if client_info is not None:
_stateful_session_client_info[session_id] = client_info
break

View file

@ -3060,3 +3060,18 @@ async def test_closing_persistent_session_fails_a_caller_blocked_on_a_full_queue
outcomes: Final = await asyncio.wait_for(asyncio.gather(*waiters, return_exceptions=True), 5)
assert all(isinstance(outcome, RuntimeError) for outcome in outcomes), outcomes
await asyncio.wait_for(session.wait_closed(), 5)
@pytest.mark.asyncio
async def test_persistent_session_ends_after_a_broken_stream_so_the_next_call_gets_a_fresh_one():
app: Final = _stateful_upstream()
async with app.router.lifespan_context(app):
_, session = _client_with_session(app)
async def broken_stream(_: object) -> str:
raise anyio.BrokenResourceError()
with pytest.raises(anyio.BrokenResourceError):
await asyncio.wait_for(session.run(broken_stream), 5)
await asyncio.wait_for(session.wait_closed(), 5)
assert session.closed, "a dead transport must end the session instead of being reused"

View file

@ -10383,6 +10383,7 @@ class TestOBOConcurrencyLimit:
manager = MCPServerManager()
manager._create_mcp_client = AsyncMock(return_value=_SessionRecordingClient())
manager.track_gateway_session("gateway-1")
for tool in ("select_project", "create_feature"):
result = await manager._call_regular_mcp_tool(
@ -14643,7 +14644,12 @@ async def test_upstream_session_is_shared_per_gateway_session_and_released_with_
try:
assert await manager._upstream_session_for(client, server, None) is None
assert await manager._upstream_session_for(client, server, {"accept": "application/json"}) is None
assert await manager._upstream_session_for(client, server, {"mcp-session-id": "forged"}) is None, (
"an mcp-session-id the gateway never issued must not open a long-lived upstream session"
)
manager.track_gateway_session("gw-1")
manager.track_gateway_session("gw-2")
first: Final = await manager._upstream_session_for(client, server, {"Mcp-Session-Id": "gw-1"})
second: Final = await manager._upstream_session_for(client, server, {"mcp-session-id": "gw-1"})
other: Final = await manager._upstream_session_for(client, server, {"mcp-session-id": "gw-2"})
@ -14653,6 +14659,10 @@ async def test_upstream_session_is_shared_per_gateway_session_and_released_with_
manager.release_upstream_sessions("gw-1")
await asyncio.wait_for(first.wait_closed(), 5)
assert first.closed and not other.closed
assert await manager._upstream_session_for(client, server, {"mcp-session-id": "gw-1"}) is None, (
"a released gateway session must not reopen upstream sessions"
)
manager.track_gateway_session("gw-1")
replacement: Final = await manager._upstream_session_for(client, server, {"mcp-session-id": "gw-1"})
assert replacement is not first and not replacement.closed