diff --git a/litellm/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py index 1ccae8de35f..3ac39710dc9 100644 --- a/litellm/experimental_mcp_client/client.py +++ b/litellm/experimental_mcp_client/client.py @@ -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) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 20c114a2f3e..58ec737e13f 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index fe508fc22c7..815eb923863 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -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 diff --git a/tests/integration/_support/mcp.py b/tests/integration/_support/mcp.py index bdf60becbaa..2b055036852 100644 --- a/tests/integration/_support/mcp.py +++ b/tests/integration/_support/mcp.py @@ -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} diff --git a/tests/integration/contracts.json b/tests/integration/contracts.json index 605f78b71be..c5b7958cbf7 100644 --- a/tests/integration/contracts.json +++ b/tests/integration/contracts.json @@ -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" diff --git a/tests/integration/mcp/test_mcp_lifecycle.py b/tests/integration/mcp/test_mcp_lifecycle.py index b32cf97605f..b2547d0a8d7 100644 --- a/tests/integration/mcp/test_mcp_lifecycle.py +++ b/tests/integration/mcp/test_mcp_lifecycle.py @@ -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 diff --git a/tests/test_litellm/experimental_mcp_client/test_mcp_client.py b/tests/test_litellm/experimental_mcp_client/test_mcp_client.py index b260d240f29..fd49b6d2045 100644 --- a/tests/test_litellm/experimental_mcp_client/test_mcp_client.py +++ b/tests/test_litellm/experimental_mcp_client/test_mcp_client.py @@ -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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 7725aca1948..84595f8d917 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -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 + )