fix(mcp): keep one upstream MCP session per gateway session across tool 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 14:15:33 +00:00
parent 860bc7811d
commit 71aa41699a
8 changed files with 332 additions and 11 deletions

View file

@ -672,6 +672,9 @@ class MCPClient:
await anyio.lowlevel.checkpoint_if_cancelled()
return result
def open_persistent_session(self) -> "PersistentMCPSession":
return PersistentMCPSession(self)
def update_auth_value(self, mcp_auth_value: str | dict[str, str]) -> None:
"""
Set the authentication header for the MCP client.
@ -838,10 +841,15 @@ class MCPClient:
call_tool_request_params: MCPCallToolRequestParams,
host_progress_callback: Callable | None = None,
raise_on_error: bool = False,
persistent_session: "PersistentMCPSession | None" = None,
) -> MCPCallToolResult:
"""
Call an MCP Tool.
persistent_session runs the call inside an already-open upstream session (one
upstream mcp-session-id shared by every call of a stateful gateway session) instead of
the per-call initialize + teardown that run_with_session performs.
Args:
raise_on_error: When True, re-raise the underlying exception instead of returning an
``isError=True`` result. The token-exchange (OBO) tool-call path uses this to detect
@ -871,8 +879,9 @@ class MCPClient:
progress_callback=on_progress,
)
run: Final = self.run_with_session if persistent_session is None else persistent_session.run
try:
tool_result: Final = await self.run_with_session(_call_tool_operation, quiet_on_error=raise_on_error)
tool_result: Final = await run(_call_tool_operation, quiet_on_error=raise_on_error)
verbose_logger.info("MCP client tool call '%s' completed successfully", call_tool_request_params.name)
return tool_result
except asyncio.CancelledError:
@ -1178,3 +1187,70 @@ class MCPClient:
"the MCP server may have crashed, disconnected, or timed out."
)
raise
_PendingOperation: TypeAlias = tuple[Callable[[ClientSession], Awaitable[object]], "asyncio.Future[object]"]
class PersistentMCPSession:
"""One upstream MCP session kept open across operations.
The SDK transport owns anyio cancel scopes that must be entered and exited by the same
task, so a dedicated task holds the session open and serves queued operations in order.
"""
def __init__(self, client: MCPClient) -> None:
loop: Final = asyncio.get_running_loop()
self._client: Final = client
self._queue: Final[asyncio.Queue[_PendingOperation | None]] = asyncio.Queue()
self._ready: Final[asyncio.Future[None]] = loop.create_future()
self._task: Final = loop.create_task(self._serve())
@property
def closed(self) -> bool:
return self._task.done()
async def _serve(self) -> None:
async def drain(session: ClientSession) -> None:
self._ready.set_result(None)
while (pending := await self._queue.get()) is not None:
operation, future = pending
try:
future.set_result(await operation(session))
except Exception as e:
future.set_exception(e)
if isinstance(e, (ValueError, httpx2.HTTPError, OSError, MCPError)):
return
try:
await self._client.run_with_session(drain, quiet_on_error=True)
except BaseException as e:
if not self._ready.done():
self._ready.set_exception(e)
if not isinstance(e, Exception):
raise
finally:
while not self._queue.empty():
pending = self._queue.get_nowait()
if pending is not None and not pending[1].done():
pending[1].set_exception(RuntimeError("upstream MCP session closed"))
async def run(
self,
operation: Callable[[ClientSession], Awaitable[TSessionResult]],
*,
quiet_on_error: bool = False,
) -> TSessionResult:
del quiet_on_error
await self._ready
if self.closed:
raise RuntimeError("upstream MCP session closed")
future: Final[asyncio.Future[object]] = asyncio.get_running_loop().create_future()
self._queue.put_nowait((operation, future))
return cast(TSessionResult, await future)
def close(self) -> None:
self._queue.put_nowait(None)
async def wait_closed(self) -> None:
await asyncio.gather(self._task, return_exceptions=True)

View file

@ -62,7 +62,13 @@ from litellm.constants import (
MCP_TOOL_LISTING_TIMEOUT,
)
from litellm.exceptions import BlockedPiiEntityError, GuardrailRaisedException
from litellm.experimental_mcp_client.client import MCPClient, MCPSigV4Auth, strip_auth_scheme, to_basic_credentials
from litellm.experimental_mcp_client.client import (
MCPClient,
MCPSigV4Auth,
PersistentMCPSession,
strip_auth_scheme,
to_basic_credentials,
)
from litellm.integrations.custom_guardrail import (
_sync_guardrail_info_to_logging_obj, # pyright: ignore[reportPrivateUsage] - the same bridge @log_guardrail_information uses; reimplementing it here would fork the metadata-key logic
)
@ -1861,6 +1867,10 @@ class MCPServerManager:
discovery_ttl, discovery_clock, TypeAdapter(tuple[ResourceTemplate, ...])
)
self.registry: dict[str, MCPServer] = {}
# (gateway mcp-session-id, server_id, upstream auth fingerprint) -> the one upstream
# session every tool call of that gateway session reuses, so stateful upstreams keep
# their per-session state between calls. Released with the gateway session.
self._upstream_sessions: dict[tuple[str, str, str], PersistentMCPSession] = {}
self._openapi_health_probes: Callable[[str], _OpenAPIHealthProbe] = lru_cache(maxsize=128)(_OpenAPIHealthProbe)
self.config_mcp_servers: dict[str, MCPServer] = {}
"""
@ -5776,6 +5786,30 @@ class MCPServerManager:
async with semaphore:
yield
async def _upstream_session_for(
self,
client: MCPClient,
mcp_server: MCPServer,
raw_headers: Mapping[str, str] | None,
) -> PersistentMCPSession | None:
gateway_session_id: Final = next(
(value for key, value in (raw_headers or {}).items() if key.lower() == "mcp-session-id" and value),
None,
)
if gateway_session_id is None 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)
if existing is not None and not existing.closed:
return existing
opened: Final = client.open_persistent_session()
self._upstream_sessions[key] = opened
return opened
def release_upstream_sessions(self, gateway_session_id: str) -> None:
for key in tuple(key for key in self._upstream_sessions if key[0] == gateway_session_id):
self._upstream_sessions.pop(key).close()
async def _obo_call_tool_with_retry(
self,
*,
@ -5982,6 +6016,7 @@ class MCPServerManager:
raw_headers=raw_headers,
client_ip=client_ip,
)
persistent_session: Final = await self._upstream_session_for(client, mcp_server, raw_headers)
call_tool_params: Final = MCPCallToolRequestParams(
name=original_tool_name,
@ -6019,7 +6054,9 @@ class MCPServerManager:
async def _call_tool_via_client(client, params):
async with self._limit_outbound_concurrency(mcp_server):
if not relays_upstream_auth:
return await client.call_tool(params, host_progress_callback=host_progress_callback)
return await client.call_tool(
params, host_progress_callback=host_progress_callback, persistent_session=persistent_session
)
# The client-forwarded modes carry the caller's own upstream token, so an upstream
# 401 (expired/invalid token) is the caller's to resolve: relay it as
# MCPUpstreamAuthError so single-server REST callers turn it into a 401 +
@ -6031,7 +6068,10 @@ class MCPServerManager:
# the same isError degradation the default path produces.
try:
return await client.call_tool(
params, host_progress_callback=host_progress_callback, raise_on_error=True
params,
host_progress_callback=host_progress_callback,
raise_on_error=True,
persistent_session=persistent_session,
)
except Exception as e:
auth_info: Final = _extract_upstream_auth_failure(e)

View file

@ -628,6 +628,7 @@ if MCP_AVAILABLE:
_stateful_session_locks.pop(session_id, None)
_stateful_session_active_request_counts.pop(session_id, None)
_stateful_session_client_info.pop(session_id, None)
operations.global_mcp_server_manager.release_upstream_sessions(session_id)
# Keep this alias so existing references to session_manager still work
session_manager: Final = session_manager_stateless

View file

@ -1,6 +1,6 @@
import json
import queue
from collections.abc import Iterator
from collections.abc import Generator, Iterator
from contextlib import contextmanager
from dataclasses import dataclass
from typing import Final
@ -9,10 +9,11 @@ import httpx
from integration._support.asgi import asgi_server
from integration._support.client import Gateway, Scenario
from integration._support.database import read_rows
from mcp.server.mcpserver import MCPServer
from mcp.server.mcpserver import Context, MCPServer
from mcp.server.transport_security import TransportSecuritySettings
from mcp_tests.mcp_e2e_upstream_server import add, multiply
from starlette.requests import Request
from starlette.responses import Response
from starlette.types import Message, Receive, Scope, Send
@ -65,6 +66,55 @@ def mcp_peer() -> Iterator[McpPeer]:
yield McpPeer(url + "/mcp", observed)
@contextmanager
def stateful_mcp_peer() -> Generator[McpPeer]:
service: Final = MCPServer("integration-stateful")
selected: Final[dict[str, str]] = {} # mutable-ok: per-session state the upstream keeps across tool calls
def upstream_session(ctx: Context) -> str:
assert ctx.headers is not None, "stateful upstream requires HTTP request headers"
return ctx.headers["mcp-session-id"]
@service.tool()
def select_project(name: str, ctx: Context) -> str:
selected[upstream_session(ctx)] = name
return f"selected {name}"
@service.tool()
def create_feature(title: str, ctx: Context) -> str:
project: Final = selected.get(upstream_session(ctx))
if project is None:
raise ValueError("no project selected in this session")
return f"{project}/{title}"
app: Final = service.streamable_http_app(
stateless_http=False,
json_response=True,
transport_security=TransportSecuritySettings(enable_dns_rebinding_protection=False),
)
observed: Final[queue.Queue[dict[str, object]]] = queue.Queue()
async def capture(scope: Scope, receive: Receive, send: Send) -> None:
if scope["type"] == "http" and scope["method"] == "GET":
await Response(status_code=405)(scope, receive, send)
return
if scope["type"] != "http" or scope["method"] != "POST":
await app(scope, receive, send)
return
body: Final = await Request(scope, receive).body()
observed.put({"body": json.loads(body) if body else None, "headers": dict(scope["headers"])})
message: Final[Message] = {"type": "http.request", "body": body, "more_body": False}
pending: Final = iter((message,))
async def replay() -> Message:
return next(pending, {"type": "http.disconnect"})
await app(scope, replay, send)
with asgi_server(capture) as url:
yield McpPeer(url + "/mcp", observed)
def register_mcp(scenario: Scenario, peer: McpPeer, alias: str, **fields: object) -> str:
response: Final = scenario.gateway.request(
"POST", "/v1/mcp/server", {"server_name": alias, "alias": alias, "url": peer.url, "transport": "http", **fields}

View file

@ -250,6 +250,9 @@
"tests/integration/mcp/test_mcp_lifecycle.py::test_generated_mcp_edits_preserve_actual_headers_and_tool_results": [
"other.mcp.lifecycle.generated_save_reload_preserves_effective_headers"
],
"tests/integration/mcp/test_mcp_lifecycle.py::test_stateful_upstream_keeps_selection_across_tool_calls_through_gateway": [
"mcp.call_tool.session.stateful_upstream_state_survives_between_tool_calls"
],
"tests/integration/observability/test_callback_delivery.py::test_concurrent_success_and_failure_join_callbacks_and_rows_without_credentials": [
"other.observability.callbacks.credentials_stay_out_of_event_bodies",
"other.observability.callbacks.concurrent_results_join_complete_events_and_rows"

View file

@ -1,19 +1,24 @@
import asyncio
import json
import uuid
from contextlib import ExitStack
from pathlib import Path
from typing import Final
from typing import Final, cast
import httpx
import httpx2
import pytest
import yaml
from hypothesis import strategies as st
from hypothesis.stateful import RuleBasedStateMachine, invariant, rule, run_state_machine_as_test
from integration._support.client import Gateway
from integration._support.database import read_rows
from integration._support.generation import LIFECYCLE_SETTINGS, bounded_http_requests
from integration._support.mcp import call_tool, mcp_peer, register_mcp, stateful_mcp_peer, tool_names
from integration._support.process import owned_proxy
from integration._support.mcp import call_tool, mcp_peer, register_mcp, tool_names
from mcp import ClientSession
from mcp.client.streamable_http import streamable_http_client
from mcp.types import CallToolResult, TextContent
@pytest.mark.covers("mcp.call_tool.saved_headers.reach_actual_transport")
@ -304,3 +309,40 @@ def test_same_url_server_grants_scope_discovery_and_direct_or_virtual_execution(
assert calls[0]["headers"][b"x-integration-server"] == aliases[server_index].encode()
expected_auth: Final = f"Bearer synthetic-{aliases[server_index]}".encode() if authenticated else None
assert all(item["headers"].get(b"authorization") == expected_auth for item in observed)
async def _select_then_create(url: str, key: str) -> tuple[CallToolResult, CallToolResult]:
async with httpx2.AsyncClient(headers={"x-litellm-api-key": key}, timeout=30, trust_env=False) as http:
async with streamable_http_client(url, http_client=http) as (read, write):
async with ClientSession(read, write) as session:
await session.initialize()
names: Final = {tool.name.rsplit("-", 1)[-1]: tool.name for tool in (await session.list_tools()).tools}
selected: Final = await session.call_tool(names["select_project"], {"name": "alpha"})
created: Final = await session.call_tool(names["create_feature"], {"title": "login"})
return selected, created
def _upstream_session_ids(observed: tuple[dict[str, object], ...]) -> set[str | None]:
headers: Final = tuple(cast(dict[bytes, bytes], item["headers"]) for item in observed)
return {raw.decode() if isinstance(raw := item.get(b"mcp-session-id"), bytes) else None for item in headers}
@pytest.mark.covers("mcp.call_tool.session.stateful_upstream_state_survives_between_tool_calls")
def test_stateful_upstream_keeps_selection_across_tool_calls_through_gateway(gateway: Gateway) -> None:
with stateful_mcp_peer() as peer, gateway.scenario() as scenario:
alias: Final = "integration" + uuid.uuid4().hex
identity: Final = register_mcp(scenario, peer, alias)
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
peer.drain()
selected, created = asyncio.run(_select_then_create(f"{gateway.client.base_url}/{alias}/mcp", key))
observed: Final = peer.drain()
for session_id in _upstream_session_ids(observed):
if session_id is not None:
httpx.delete(peer.url, headers={"mcp-session-id": session_id}, trust_env=False)
call_sessions: Final = _upstream_session_ids(
tuple(item for item in observed if cast(dict[str, object], item["body"]).get("method") == "tools/call")
)
assert selected.is_error is False, selected.model_dump_json()
assert created.is_error is False, created.model_dump_json()
assert isinstance(created.content[0], TextContent) and created.content[0].text == "alpha/login"
assert len(call_sessions) == 1 and None not in call_sessions, call_sessions

View file

@ -15,6 +15,9 @@ import httpx2
import pytest
from mcp import MCPError
from mcp.client.streamable_http import streamable_http_client
from mcp.server.mcpserver import Context
from mcp.server.mcpserver import MCPServer as UpstreamServer
from mcp.server.transport_security import TransportSecuritySettings
from mcp.shared.message import SessionMessage
from mcp.types import (
CONNECTION_CLOSED,
@ -2896,3 +2899,64 @@ async def test_cancellation_delivers_termination_over_tcp(
closed: Final = await asyncio.wait_for(asyncio.gather(*connections, return_exceptions=True), 2)
assert all(result is None or isinstance(result, asyncio.CancelledError) for result in closed), closed
await asyncio.wait_for(listener.wait_closed(), 2)
class _StatefulUpstreamClient(MCPClient):
"""An MCPClient whose streamable-HTTP transport talks to an in-process stateful MCP server."""
def __init__(self, app, **kwargs):
super().__init__(**kwargs)
self._app = app
def _create_transport_context(self) -> tuple[_TransportContext, httpx2.AsyncClient]:
http_client: Final = self._create_httpx_client_factory(transport=httpx2.ASGITransport(app=self._app))(
headers=self._get_auth_headers(), timeout=httpx2.Timeout(self.timeout)
)
return streamable_http_client(self.server_url, http_client=http_client), http_client
def _stateful_upstream():
service: Final = UpstreamServer("stateful")
selected: Final[dict[str, str]] = {}
@service.tool()
def select_project(name: str, ctx: Context) -> str:
selected[ctx.headers["mcp-session-id"]] = name
return f"selected {name}"
@service.tool()
def create_feature(title: str, ctx: Context) -> str:
return f"{selected[ctx.headers['mcp-session-id']]}/{title}"
return service.streamable_http_app(
stateless_http=False,
json_response=True,
transport_security=TransportSecuritySettings(enable_dns_rebinding_protection=False),
)
@pytest.mark.asyncio
async def test_persistent_session_keeps_upstream_state_across_tool_calls():
app: Final = _stateful_upstream()
async with app.router.lifespan_context(app):
client: Final = _StatefulUpstreamClient(app, server_url="http://upstream/mcp", transport_type=MCPTransport.http)
per_call: Final = await client.call_tool(CallToolRequestParams(name="select_project", arguments={"name": "a"}))
assert per_call.is_error is False
fresh: Final = await client.call_tool(CallToolRequestParams(name="create_feature", arguments={"title": "b"}))
assert fresh.is_error is True, "a fresh upstream session per call must not see the earlier selection"
session: Final = client.open_persistent_session()
try:
selected: Final = await client.call_tool(
CallToolRequestParams(name="select_project", arguments={"name": "a"}), persistent_session=session
)
created: Final = await client.call_tool(
CallToolRequestParams(name="create_feature", arguments={"title": "b"}), persistent_session=session
)
finally:
session.close()
assert selected.is_error is False
assert created.is_error is False
assert created.content[0].text == "a/b"
await asyncio.wait_for(session.wait_closed(), 5)
assert session.closed

View file

@ -7,7 +7,7 @@ import os
import sys
from datetime import datetime
from pathlib import Path
from typing import Any, Dict, Final, Literal, Optional
from typing import Any, Dict, Final, Literal, Optional, cast
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
@ -27,7 +27,7 @@ import contextlib
import httpx
import httpx2
from mcp import ReadResourceResult, Resource
from mcp import ClientSession, ReadResourceResult, Resource
from mcp.types import (
CallToolResult,
GetPromptResult,
@ -39,6 +39,7 @@ from mcp.types import Tool as MCPTool
from pydantic import AnyUrl, TypeAdapter
from litellm.constants import MCP_METADATA_TIMEOUT
from litellm.experimental_mcp_client.client import MCPClient
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
MCPServerManager,
_deserialize_json_dict,
@ -14516,3 +14517,47 @@ async def test_client_sampling_does_not_fill_explicit_context_from_another_ambie
assert captured["client_ip"] is None
finally:
auth_context_var.reset(token)
class _OfflineMCPClient(MCPClient):
"""An MCPClient whose sessions never touch the network: operations run against a stand-in session."""
async def run_with_session(self, operation, *, quiet_on_error: bool = False):
del quiet_on_error
return await operation(cast(ClientSession, object()))
@pytest.mark.asyncio
async def test_upstream_session_is_shared_per_gateway_session_and_released_with_it():
manager: Final = MCPServerManager()
server: Final = MCPServer(server_id="s1", name="stateful", url="http://upstream/mcp", transport=MCPTransport.http)
client: Final = _OfflineMCPClient(server_url=server.url, transport_type=MCPTransport.http)
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
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"})
assert first is not None and first is second, "every call of one gateway session must share one upstream session"
assert other is not None and other is not first, "distinct gateway sessions must not share upstream state"
manager.release_upstream_sessions("gw-1")
await asyncio.wait_for(first.wait_closed(), 5)
assert first.closed and not other.closed
replacement: Final = await manager._upstream_session_for(client, server, {"mcp-session-id": "gw-1"})
assert replacement is not first and not replacement.closed
other.close()
await asyncio.wait_for(other.wait_closed(), 5)
reopened: Final = await manager._upstream_session_for(client, server, {"mcp-session-id": "gw-2"})
assert reopened is not other and not reopened.closed, "a dead upstream session must be replaced, not reused"
stdio: Final = MCPServer(server_id="s2", name="local", command="cat", transport=MCPTransport.stdio)
assert await manager._upstream_session_for(client, stdio, {"mcp-session-id": "gw-1"}) is None
finally:
for gateway_session_id in ("gw-1", "gw-2"):
manager.release_upstream_sessions(gateway_session_id)
await asyncio.wait_for(
asyncio.gather(*(s.wait_closed() for s in manager._upstream_sessions.values()), return_exceptions=True), 5
)