mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
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:
parent
074e48425b
commit
dc10e0f4ea
4 changed files with 139 additions and 13 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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 ->
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue