fix(mcp): keep persistent upstream session alive across caller timeouts and reuse it for OBO calls

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Devin AI 2026-09-23 15:28:43 +00:00
parent 074e48425b
commit dc10e0f4ea
4 changed files with 139 additions and 13 deletions

View file

@ -1205,6 +1205,7 @@ class PersistentMCPSession:
self._client: Final = client
self._queue: Final[asyncio.Queue[_PendingOperation]] = asyncio.Queue(maxsize=_MAX_PENDING_OPERATIONS)
self._ready: Final[asyncio.Future[None]] = loop.create_future()
self._active: asyncio.Future[object] | None = None
self._task: Final = loop.create_task(self._serve())
@property
@ -1216,25 +1217,37 @@ class PersistentMCPSession:
self._ready.set_result(None)
while True:
operation, future = await self._queue.get()
if future.done():
continue
self._active = future
try:
future.set_result(await operation(session))
result: Final = await operation(session)
except Exception as e:
future.set_exception(e)
if not future.done():
future.set_exception(e)
if isinstance(e, (ValueError, httpx2.HTTPError, OSError, MCPError)):
return
else:
if not future.done():
future.set_result(result)
self._active = None
try:
await self._client.run_with_session(drain, quiet_on_error=True)
except BaseException as e:
except Exception as e:
if not self._ready.done():
self._ready.set_exception(e)
if not isinstance(e, Exception):
raise
finally:
while not self._queue.empty():
_, future = self._queue.get_nowait()
if not future.done():
future.set_exception(RuntimeError("upstream MCP session closed"))
self._fail_waiters()
def _fail_waiters(self) -> None:
pending: Final = (self._ready, self._active, *(future for _, future in self._drained()))
for future in pending:
if future is not None and not future.done():
future.set_exception(RuntimeError("upstream MCP session closed"))
def _drained(self) -> tuple[_PendingOperation, ...]:
return tuple(self._queue.get_nowait() for _ in range(self._queue.qsize()))
async def run(
self,

View file

@ -5810,6 +5810,10 @@ class MCPServerManager:
for key in tuple(key for key in self._upstream_sessions if key[0] == gateway_session_id):
self._upstream_sessions.pop(key).close()
def _drop_upstream_session(self, session: PersistentMCPSession | None) -> None:
for key in tuple(key for key, value in self._upstream_sessions.items() if value is session):
self._upstream_sessions.pop(key).close()
async def _obo_call_tool_with_retry(
self,
*,
@ -5824,6 +5828,7 @@ class MCPServerManager:
user_api_key_auth: UserAPIKeyAuth | None,
raw_headers: Mapping[str, str] | None = None,
client_ip: str | None = None,
persistent_session: PersistentMCPSession | None = None,
) -> CallToolResult:
"""Call a token_exchange (OBO) tool; on an upstream 401/403 re-mint the token once and retry.
@ -5834,11 +5839,15 @@ class MCPServerManager:
"""
try:
return await client.call_tool(
call_tool_params, host_progress_callback=host_progress_callback, raise_on_error=True
call_tool_params,
host_progress_callback=host_progress_callback,
raise_on_error=True,
persistent_session=persistent_session,
)
except Exception as exc:
if _extract_upstream_auth_failure(exc) is None:
return MCPClient.error_tool_result(exc)
self._drop_upstream_session(persistent_session)
spec: Final = to_server_spec(mcp_server)
if spec is not None:
await self._cred_provider.invalidate_credentials(to_subject(user_api_key_auth, subject_token), spec)
@ -5852,7 +5861,11 @@ class MCPServerManager:
raw_headers=raw_headers,
client_ip=client_ip,
)
return await retry_client.call_tool(call_tool_params, host_progress_callback=host_progress_callback)
return await retry_client.call_tool(
call_tool_params,
host_progress_callback=host_progress_callback,
persistent_session=await self._upstream_session_for(retry_client, mcp_server, raw_headers),
)
async def _call_regular_mcp_tool(
self,
@ -6041,6 +6054,7 @@ class MCPServerManager:
user_api_key_auth=user_api_key_auth,
raw_headers=raw_headers,
client_ip=client_ip,
persistent_session=persistent_session,
)
tool_call_coro = _obo_call_tool_limited()

View file

@ -42,6 +42,7 @@ from pydantic import TypeAdapter, ValidationError
import litellm.experimental_mcp_client.client as mcp_client_module
from litellm.experimental_mcp_client.client import (
MCPClient,
PersistentMCPSession,
_first_non_cancelled_cause,
_TransportContext,
as_mcp_read_timeout,
@ -2960,3 +2961,55 @@ async def test_persistent_session_keeps_upstream_state_across_tool_calls():
assert created.content[0].text == "a/b"
await asyncio.wait_for(session.wait_closed(), 5)
assert session.closed
def _client_with_session(app) -> tuple[_StatefulUpstreamClient, PersistentMCPSession]:
client: Final = _StatefulUpstreamClient(app, server_url="http://upstream/mcp", transport_type=MCPTransport.http)
return client, client.open_persistent_session()
@pytest.mark.asyncio
async def test_persistent_session_survives_a_caller_timeout_on_one_operation():
app: Final = _stateful_upstream()
async with app.router.lifespan_context(app):
client, session = _client_with_session(app)
try:
with pytest.raises(asyncio.TimeoutError):
await asyncio.wait_for(
client.call_tool(
CallToolRequestParams(name="select_project", arguments={"name": "a"}),
persistent_session=session,
),
timeout=0,
)
selected: Final = await client.call_tool(
CallToolRequestParams(name="select_project", arguments={"name": "c"}), persistent_session=session
)
created: Final = await client.call_tool(
CallToolRequestParams(name="create_feature", arguments={"title": "d"}), persistent_session=session
)
finally:
session.close()
assert selected.is_error is False, "the session must outlive a timed out call"
assert created.content[0].text == "c/d"
await asyncio.wait_for(session.wait_closed(), 5)
@pytest.mark.asyncio
async def test_closing_persistent_session_mid_operation_fails_the_waiter_instead_of_hanging():
app: Final = _stateful_upstream()
async with app.router.lifespan_context(app):
_, session = _client_with_session(app)
started: Final = asyncio.Event()
async def slow_operation(_: object) -> str:
started.set()
await asyncio.sleep(30)
return "never"
waiter: Final = asyncio.ensure_future(session.run(slow_operation))
await asyncio.wait_for(started.wait(), 5)
session.close()
with pytest.raises(RuntimeError, match="upstream MCP session closed"):
await asyncio.wait_for(waiter, 5)
await asyncio.wait_for(session.wait_closed(), 5)

View file

@ -7,6 +7,7 @@ import os
import sys
from datetime import datetime
from pathlib import Path
from types import SimpleNamespace
from typing import Any, Dict, Final, Literal, Optional, cast
from unittest.mock import AsyncMock, MagicMock, patch
@ -10033,7 +10034,7 @@ class _RetryFakeClient:
self._MCPClient = MCPClient
self.attempts = 0
async def call_tool(self, params, host_progress_callback=None, raise_on_error=False):
async def call_tool(self, params, host_progress_callback=None, raise_on_error=False, persistent_session=None):
self.attempts += 1
if self._raises is not None:
if raise_on_error:
@ -10245,7 +10246,7 @@ class TestOBOConcurrencyLimit:
inflight = {"current": 0, "peak": 0}
class _ConcurrencyRecordingClient:
async def call_tool(self, params, host_progress_callback=None, raise_on_error=False):
async def call_tool(self, params, host_progress_callback=None, raise_on_error=False, persistent_session=None):
inflight["current"] += 1
inflight["peak"] = max(inflight["peak"], inflight["current"])
try:
@ -10293,6 +10294,51 @@ class TestOBOConcurrencyLimit:
assert inflight["current"] == 0
assert all(result.is_error is False for result in results)
@pytest.mark.asyncio
async def test_obo_dispatch_reuses_the_gateway_sessions_persistent_upstream_session(self):
server = MCPServer(
server_id="obo-stateful",
name="obo",
url="https://upstream.example/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2_token_exchange,
token_exchange_endpoint="https://idp.example.com/token",
client_id="cid",
client_secret="csec",
)
sessions_seen = []
class _SessionRecordingClient:
async def discovery_auth_fingerprint(self):
return "same-token"
def open_persistent_session(self):
return SimpleNamespace(closed=False, close=lambda: None)
async def call_tool(self, params, host_progress_callback=None, raise_on_error=False, persistent_session=None):
sessions_seen.append(persistent_session)
return CallToolResult(content=[], isError=False)
manager = MCPServerManager()
manager._create_mcp_client = AsyncMock(return_value=_SessionRecordingClient())
for tool in ("select_project", "create_feature"):
result = await manager._call_regular_mcp_tool(
mcp_server=server,
original_tool_name=tool,
arguments={},
tasks=[],
mcp_auth_header=None,
mcp_server_auth_headers=None,
oauth2_headers={"Authorization": "Bearer subject-jwt"},
raw_headers={"mcp-session-id": "gateway-1"},
proxy_logging_obj=None,
)
assert result.is_error is False
assert len(sessions_seen) == 2 and None not in sessions_seen, sessions_seen
assert sessions_seen[0] is sessions_seen[1], "OBO calls in one gateway session must share one upstream session"
class TestOBOEndpointDiscovery:
"""An oauth2_token_exchange server with no configured token endpoint discovers it (RFC 9728 ->