mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
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:
parent
595ced358a
commit
8b59890789
5 changed files with 42 additions and 2 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue