fix(mcp): restore legacy SSE and bounded cancellation cleanup (#42382)

* fix(mcp): restore legacy SSE and bounded cancellation cleanup

* fix(mcp): drain termination despite repeated task cancellation

* fix(mcp): admit virtual keys on legacy SSE message routes

* fix(mcp): preserve cancellation through late transport failures

* fix(mcp): preserve process exits during cleanup

* test(mcp): drain idle sockets before server shutdown

---------

Co-authored-by: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com>
This commit is contained in:
joshua-berri 2026-09-22 18:27:50 +00:00 • committed by GitHub
parent 2ea43214c9
commit 0d0f73cd14
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
10 changed files with 900 additions and 380 deletions

View file

@ -7,11 +7,11 @@ import base64
import hashlib
import json
import os
from collections.abc import Awaitable, Callable, Generator, Sequence
from collections.abc import AsyncIterator, Awaitable, Callable, Generator, Sequence
from contextlib import AbstractAsyncContextManager
from functools import partial
from types import MappingProxyType
from typing import Any, Final, TypeAlias, TypeVar
from typing import Final, TypeAlias, TypeVar, cast
import anyio
import httpx2
@ -121,14 +121,14 @@ def _strip_header_whitespace(headers: dict[str, str]) -> dict[str, str]:
}
def _first_non_cancelled_cause(exc: BaseException) -> BaseException | None:
def _first_non_cancelled_cause(exc: BaseException, cleanup_errors: tuple[Exception, ...] = ()) -> BaseException | None:
queue: Final[list[BaseException]] = [exc]
while queue:
current = queue.pop(0)
nested = getattr(current, "exceptions", None)
if nested:
queue.extend(nested)
elif not isinstance(current, asyncio.CancelledError):
elif not isinstance(current, asyncio.CancelledError) and not any(current is error for error in cleanup_errors):
return current
return None
@ -159,7 +159,59 @@ _ListPage = TypeVar("_ListPage", bound=PaginatedResult)
_ListItem = TypeVar("_ListItem")
async def _run_bounded_cleanup(operation: Callable[[], Awaitable[TSessionResult]], deadline: float) -> TSessionResult:
async def run() -> TSessionResult:
with anyio.fail_after(max(0, deadline - anyio.current_time()), shield=True):
return await operation()
# A cancelled asyncio.gather repeatedly forwards Task.cancel, bypassing AnyIO shields.
# Isolate only cleanup, and drain it before propagating the caller's cancellation.
task: Final = asyncio.create_task(run())
interrupted: asyncio.CancelledError | None = None
with anyio.CancelScope(shield=True):
while not task.done():
try:
await asyncio.shield(task)
except asyncio.CancelledError as exc:
interrupted = exc
except Exception:
break
if interrupted is not None:
if not task.cancelled():
task.exception()
raise interrupted
return task.result()
class _MCPResponseStream(httpx2.AsyncByteStream):
def __init__(self, stream: httpx2.AsyncByteStream, record_error: Callable[[Exception], None]) -> None:
self._stream: Final = stream
self._record_error: Final = record_error
async def __aiter__(self) -> AsyncIterator[bytes]:
try:
async for chunk in self._stream:
yield chunk
except Exception as error:
self._record_error(error)
raise
async def aclose(self) -> None:
try:
await self._stream.aclose()
except Exception as error:
self._record_error(error)
raise
class _MCPHTTPClient(httpx2.AsyncClient):
cleanup_scope: anyio.CancelScope | None = None
cleanup_errors: tuple[Exception, ...] = ()
def _record_cleanup_error(self, error: Exception) -> None:
if self.cleanup_scope is not None and self.cleanup_scope.shield:
self.cleanup_errors += (error,)
async def send(
self,
request: httpx2.Request,
@ -168,11 +220,29 @@ class _MCPHTTPClient(httpx2.AsyncClient):
auth: AuthTypes | UseClientDefault | None = httpx2.USE_CLIENT_DEFAULT,
follow_redirects: bool | UseClientDefault = httpx2.USE_CLIENT_DEFAULT,
) -> httpx2.Response:
response: Final = await super().send(request, stream=stream, auth=auth, follow_redirects=follow_redirects)
if request.method == "POST" and response.is_error and response.status_code != 404:
await response.aclose()
response.raise_for_status()
return response
if request.method == "DELETE" and self.cleanup_scope is not None:
async def terminate() -> httpx2.Response:
termination: Final = await super(_MCPHTTPClient, self).send(
request, stream=stream, auth=auth, follow_redirects=follow_redirects
)
await termination.aread()
return termination
return await _run_bounded_cleanup(terminate, self.cleanup_scope.deadline)
try:
response: Final = await super().send(request, stream=stream, auth=auth, follow_redirects=follow_redirects)
if request.method == "POST" and response.is_error and response.status_code != 404:
await response.aclose()
response.raise_for_status()
if stream:
response.stream = _MCPResponseStream(
cast(httpx2.AsyncByteStream, response.stream), self._record_cleanup_error
)
return response
except Exception as error:
self._record_cleanup_error(error)
raise
class MCPSigV4Auth(httpx2.Auth):
@ -458,6 +528,7 @@ class MCPClient:
self,
transport_ctx: _TransportContext,
operation: Callable[[ClientSession], Awaitable[TSessionResult]],
http_client: httpx2.AsyncClient | None = None,
) -> TSessionResult:
"""
Execute an operation within a transport and session context.
@ -466,69 +537,97 @@ class MCPClient:
so that upstream MCP servers can request LLM inference (sampling),
user input (elicitation), or send log messages.
"""
transport: Final = await transport_ctx.__aenter__()
in_flight_error: BaseException | None = None
try:
read_stream: Final = transport[0]
write_stream: Final = transport[1]
stream_error: Final[asyncio.Future[Exception]] = asyncio.get_running_loop().create_future()
async def receive_message(
message: ServerNotification | Exception,
) -> None:
if not isinstance(message, (ValueError, httpx2.HTTPError, OSError)):
return
if not stream_error.done():
stream_error.set_result(message)
# The SDK closes pending requests when its message handler raises.
raise RuntimeError("MCP response stream failed")
# Build session kwargs with optional callbacks
session_kwargs: Final[dict[str, Any]] = {}
if self._sampling_callback is not None:
session_kwargs["sampling_callback"] = self._sampling_callback
if self._elicitation_callback is not None:
session_kwargs["elicitation_callback"] = self._elicitation_callback
if self._logging_callback is not None:
session_kwargs["logging_callback"] = self._logging_callback
# The SDK drops a response stream that ends without a JSON-RPC reply, so nothing else
# ever fails the request.
session_ctx: Final = ClientSession(
read_stream,
write_stream,
read_timeout_seconds=self.timeout,
message_handler=receive_message,
**session_kwargs,
)
session: Final = await session_ctx.__aenter__()
with anyio.CancelScope() as cleanup_scope:
if isinstance(http_client, _MCPHTTPClient):
http_client.cleanup_scope = cleanup_scope
try:
init_result: Final = await session.initialize()
self._last_initialize_instructions = None
if init_result is not None:
ins: Final = getattr(init_result, "instructions", None)
if isinstance(ins, str) and ins.strip():
self._last_initialize_instructions = ins.strip()
return await operation(session)
except MCPError:
if stream_error.done():
raise stream_error.result()
raise
finally:
transport: Final = await transport_ctx.__aenter__()
try:
await session_ctx.__aexit__(None, None, None)
read_stream: Final = transport[0]
write_stream: Final = transport[1]
stream_error: Final[asyncio.Future[Exception]] = asyncio.get_running_loop().create_future()
async def receive_message(
message: ServerNotification | Exception,
) -> None:
if not isinstance(message, (ValueError, httpx2.HTTPError, OSError)):
return
if not stream_error.done():
stream_error.set_result(message)
# The SDK closes pending requests when its message handler raises.
raise RuntimeError("MCP response stream failed")
session_kwargs: Final = {
name: callback
for name, callback in (
("sampling_callback", self._sampling_callback),
("elicitation_callback", self._elicitation_callback),
("logging_callback", self._logging_callback),
)
if callback is not None
}
# The SDK drops a response stream that ends without a JSON-RPC reply, so nothing else
# ever fails the request.
session_ctx: Final = ClientSession(
read_stream,
write_stream,
read_timeout_seconds=self.timeout,
message_handler=receive_message,
**session_kwargs,
)
session: Final = await session_ctx.__aenter__()
try:
init_result: Final = await session.initialize()
instructions: Final = getattr(init_result, "instructions", None)
self._last_initialize_instructions = (
instructions.strip() or None if isinstance(instructions, str) else None
)
result: Final = await operation(session)
except BaseException as operation_error:
in_flight_error = operation_error
if isinstance(operation_error, MCPError) and stream_error.done():
raise stream_error.result()
raise
finally:
cleanup_scope.shield = True
cleanup_scope.deadline = anyio.current_time() + 5
try:
await session_ctx.__aexit__(None, None, None)
except (Exception, asyncio.CancelledError) as e:
verbose_logger.debug("Error during session context exit: %s", e)
if in_flight_error is None and isinstance(e, asyncio.CancelledError):
raise
except BaseException as e:
verbose_logger.debug("Error during session context exit: %s", e)
except BaseException as e:
in_flight_error = e
raise
finally:
try:
await transport_ctx.__aexit__(None, None, None)
except BaseException as exit_error:
verbose_logger.debug("Error during transport context exit: %s", exit_error)
root_cause: Final = _first_non_cancelled_cause(exit_error)
if root_cause is not None and isinstance(in_flight_error, asyncio.CancelledError):
raise root_cause from in_flight_error
in_flight_error = e
raise
finally:
cleanup_scope.shield = True
cleanup_scope.deadline = min(cleanup_scope.deadline, anyio.current_time() + 5)
try:
await transport_ctx.__aexit__(None, None, None)
except (Exception, asyncio.CancelledError) as exit_error:
verbose_logger.debug("Error during transport context exit: %s", exit_error)
if in_flight_error is None and isinstance(exit_error, asyncio.CancelledError):
raise
root_cause: Final = _first_non_cancelled_cause(
exit_error, http_client.cleanup_errors if isinstance(http_client, _MCPHTTPClient) else ()
)
if root_cause is not None and isinstance(in_flight_error, asyncio.CancelledError):
raise root_cause from in_flight_error
finally:
cleanup_scope.shield = False
if isinstance(http_client, _MCPHTTPClient):
http_client.cleanup_errors = ()
http_client.cleanup_scope = None
await anyio.lowlevel.checkpoint_if_cancelled()
if cleanup_scope.cancel_called:
raise (
in_flight_error
if in_flight_error is not None
else asyncio.CancelledError("MCP session cleanup timed out")
)
return result
async def run_with_session(
self,
@ -542,10 +641,11 @@ class MCPClient:
(call_tool / list_tools under raise_on_error), so an expected pass-through re-auth does
not emit a warning per call; every other caller keeps the operator-visible warning."""
http_client: httpx2.AsyncClient | None = None
close_cancellation: asyncio.CancelledError | None = None
try:
self._last_initialize_instructions = None
transport_ctx, http_client = self._create_transport_context()
return await self._execute_session_operation(transport_ctx, operation)
result: Final = await self._execute_session_operation(transport_ctx, operation, http_client=http_client)
except Exception as e:
read_timeout: Final = as_mcp_read_timeout(e)
if read_timeout is not None:
@ -561,9 +661,16 @@ class MCPClient:
finally:
if http_client is not None:
try:
await http_client.aclose()
except BaseException as e:
await _run_bounded_cleanup(http_client.aclose, anyio.current_time() + 1)
except (Exception, asyncio.CancelledError) as e:
verbose_logger.debug("Error during http_client cleanup: %s", e)
if isinstance(e, asyncio.CancelledError):
close_cancellation = e
if close_cancellation is not None:
raise close_cancellation
await anyio.lowlevel.checkpoint_if_cancelled()
return result
def update_auth_value(self, mcp_auth_value: str | dict[str, str]) -> None:
"""

View file

@ -663,6 +663,8 @@ class MCPRequestHandler:
# with ``server.py::_get_mcp_servers_in_path``, which also accepts the
# un-rewritten form (some entry points may skip the
# ``dynamic_mcp_route`` rewrite).
if path.rstrip("/") in ("/mcp/sse", "/mcp/sse/messages"):
return []
segments: Final = [s for s in path.split("/") if s]
if len(segments) >= 2 and segments[1] == "mcp" and segments[0] != "mcp":
return [segments[0]]

View file

@ -21,6 +21,7 @@ from fastapi import FastAPI, HTTPException
from pydantic import ConfigDict, TypeAdapter, ValidationError
from starlette.requests import Request as StarletteRequest
from starlette.responses import JSONResponse
from starlette.routing import Route
from starlette.types import Message, Receive, Scope, Send
from litellm._logging import verbose_logger
@ -474,6 +475,8 @@ if MCP_AVAILABLE:
AuthContextMiddleware,
auth_context_var,
)
from mcp.server.auth.middleware.bearer_auth import AuthenticatedUser
from mcp.server.auth.provider import AccessToken
from mcp.server.context import ServerRequestContext
from mcp.server.lowlevel.server import NotificationOptions
from mcp.server.models import InitializationOptions
@ -573,7 +576,7 @@ if MCP_AVAILABLE:
version=LITELLM_MCP_SERVER_VERSION,
)
server.create_initialization_options = types.MethodType(_gateway_create_initialization_options, server)
sse: Final[SseServerTransport] = SseServerTransport("/mcp/sse/messages")
sse: Final[SseServerTransport] = SseServerTransport("/sse/messages")
# Create session managers
session_manager_stateless: Final = StreamableHTTPSessionManager(
@ -629,18 +632,9 @@ if MCP_AVAILABLE:
# Keep this alias so existing references to session_manager still work
session_manager: Final = session_manager_stateless
# Create SSE session manager
sse_session_manager: Final = StreamableHTTPSessionManager(
app=server,
event_store=None,
json_response=False, # Use SSE responses for this endpoint
stateless=True,
)
# Context managers for proper lifecycle management
_session_manager_cm = None
_session_manager_stateful_cm = None
_sse_session_manager_cm = None
_stateful_auth_context_cleanup_task: asyncio.Task | None = None
async def _purge_expired_stateful_session_auth_contexts(
@ -732,7 +726,6 @@ if MCP_AVAILABLE:
_SESSION_MANAGERS_INITIALIZED, \
_session_manager_cm, \
_session_manager_stateful_cm, \
_sse_session_manager_cm, \
_stateful_auth_context_cleanup_task
# Use async lock to prevent concurrent initialization
@ -745,12 +738,10 @@ if MCP_AVAILABLE:
# Start the session managers with context managers
_session_manager_cm = session_manager_stateless.run()
_session_manager_stateful_cm = session_manager_stateful.run()
_sse_session_manager_cm = sse_session_manager.run()
# Enter the context managers
await _session_manager_cm.__aenter__()
await _session_manager_stateful_cm.__aenter__()
await _sse_session_manager_cm.__aenter__()
_stateful_auth_context_cleanup_task = asyncio.create_task(_cleanup_expired_stateful_session_auth_contexts())
_SESSION_MANAGERS_INITIALIZED = True
@ -762,7 +753,6 @@ if MCP_AVAILABLE:
_SESSION_MANAGERS_INITIALIZED, \
_session_manager_cm, \
_session_manager_stateful_cm, \
_sse_session_manager_cm, \
_stateful_auth_context_cleanup_task
if _SESSION_MANAGERS_INITIALIZED:
@ -773,8 +763,6 @@ if MCP_AVAILABLE:
_stateful_auth_context_cleanup_task.cancel()
with contextlib.suppress(asyncio.CancelledError):
await _stateful_auth_context_cleanup_task
if _sse_session_manager_cm:
await _sse_session_manager_cm.__aexit__(None, None, None)
if _session_manager_stateful_cm:
await _session_manager_stateful_cm.__aexit__(None, None, None)
if _session_manager_cm:
@ -784,7 +772,6 @@ if MCP_AVAILABLE:
_session_manager_cm = None
_session_manager_stateful_cm = None
_sse_session_manager_cm = None
_stateful_auth_context_cleanup_task = None
_SESSION_MANAGERS_INITIALIZED = False
@ -1054,6 +1041,8 @@ if MCP_AVAILABLE:
"""
import re
if path.rstrip("/") in ("/mcp/sse", "/mcp/sse/messages"):
return None
mcp_servers_from_path: list[str] | None = None
segments: Final = [s for s in path.split("/") if s]
if len(segments) >= 2 and segments[1] == "mcp" and segments[0] != "mcp":
@ -2286,7 +2275,24 @@ if MCP_AVAILABLE:
async def handle_sse_mcp(scope: Scope, receive: Receive, send: Send) -> None:
"""Handle MCP requests through SSE."""
try:
path: Final[str] = scope.get("path", "")
bad_version: Final = unsupported_protocol_version(scope)
if bad_version is not None:
supported: Final = ", ".join(sorted(HANDSHAKE_PROTOCOL_VERSIONS))
await JSONResponse(
status_code=400,
content={ # mutable-ok: JSON-RPC error payload
"jsonrpc": "2.0",
"id": None,
"error": {
"code": INVALID_REQUEST,
"message": f"Unsupported MCP-Protocol-Version {bad_version}; supported: {supported}",
},
},
)(scope, receive, send)
return
from litellm.proxy.auth.auth_utils import get_request_route
path: Final = get_request_route(StarletteRequest(scope))
(
user_api_key_auth,
mcp_auth_header,
@ -2353,9 +2359,14 @@ if MCP_AVAILABLE:
client_ip=_sse_client_ip,
)
if not _SESSION_MANAGERS_INITIALIZED:
await initialize_session_managers()
await asyncio.sleep(0.1)
owner: Final = _owner_fingerprint_for(user_api_key_auth, oauth2_headers, _sse_client_ip)
transport_scope: Final[Scope] = {
**scope,
"user": AuthenticatedUser(AccessToken(token=owner, client_id=owner, scopes=[])),
}
if scope["method"] == "POST":
await sse.handle_post_message(transport_scope, receive, send)
return
async with _gateway_initialize_instructions_request_scope(
user_api_key_auth,
@ -2364,7 +2375,8 @@ if MCP_AVAILABLE:
scoped_server_endpoint=scoped_server_endpoint,
is_initialize=scope.get("method") == "GET",
):
await sse_session_manager.handle_request(scope, receive, send)
async with sse.connect_sse(transport_scope, receive, send) as (read_stream, write_stream):
await server.run(read_stream, write_stream, server.create_initialization_options())
except MCPUpstreamAuthError as e:
# Upstream delegated auth returned 401; surface it to the client so
# standards-compliant MCP clients trigger the upstream OAuth flow.
@ -2387,7 +2399,6 @@ if MCP_AVAILABLE:
# Try to send a graceful error response for non-HTTP exceptions
try:
# Send a proper HTTP error response instead of letting the exception bubble up
from starlette.responses import JSONResponse
from starlette.status import HTTP_500_INTERNAL_SERVER_ERROR
error_response: Final = JSONResponse(
@ -2418,11 +2429,22 @@ if MCP_AVAILABLE:
"""
return {"enabled": MCP_AVAILABLE}
class _LegacySseEndpoint:
async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
await handle_sse_mcp(scope, receive, send)
for sse_path, sse_method in (
("/sse", "GET"),
("/sse/", "GET"),
("/sse/messages", "POST"),
("/sse/messages/", "POST"),
):
app.router.routes.append(Route(sse_path, endpoint=_LegacySseEndpoint(), methods=[sse_method]))
# Mount the MCP handlers
app.mount("/", handle_streamable_http_mcp)
app.mount("/mcp", handle_streamable_http_mcp)
app.mount("/{mcp_server_name}/mcp", handle_streamable_http_mcp)
app.mount("/sse", handle_sse_mcp)
app.add_middleware(AuthContextMiddleware)
########################################################

View file

@ -1,138 +1,3 @@
"""
This is a modification of code from: https://github.com/SecretiveShell/MCP-Bridge/blob/master/mcp_bridge/mcp_server/sse_transport.py
from mcp.server.sse import SseServerTransport
Credit to the maintainers of SecretiveShell for their SSE Transport implementation
"""
from contextlib import asynccontextmanager
from typing import Any, Final
from urllib.parse import quote
from uuid import UUID, uuid4
import anyio
from anyio.streams.memory import MemoryObjectReceiveStream, MemoryObjectSendStream
from fastapi.requests import Request
from fastapi.responses import Response
from mcp import types
from pydantic import ValidationError
from sse_starlette import EventSourceResponse
from starlette.types import Receive, Scope, Send
from litellm._logging import verbose_logger
class SseServerTransport:
"""
SSE server transport for MCP. This class provides _two_ ASGI applications,
suitable to be used with a framework like Starlette and a server like Hypercorn:
1. connect_sse() is an ASGI application which receives incoming GET requests,
and sets up a new SSE stream to send server messages to the client.
2. handle_post_message() is an ASGI application which receives incoming POST
requests, which should contain client messages that link to a
previously-established SSE session.
"""
_endpoint: str
_read_stream_writers: dict[UUID, MemoryObjectSendStream[types.JSONRPCMessage | Exception]]
def __init__(self, endpoint: str) -> None:
"""
Creates a new SSE server transport, which will direct the client to POST
messages to the relative or absolute URL given.
"""
super().__init__()
self._endpoint = endpoint
self._read_stream_writers = {}
verbose_logger.debug("SseServerTransport initialized with endpoint: %s", endpoint)
@asynccontextmanager
async def connect_sse(self, request: Request):
if request.scope["type"] != "http":
verbose_logger.error("connect_sse received non-HTTP request")
raise ValueError("connect_sse can only handle HTTP requests")
verbose_logger.debug("Setting up SSE connection")
read_stream: MemoryObjectReceiveStream[types.JSONRPCMessage | Exception]
read_stream_writer: MemoryObjectSendStream[types.JSONRPCMessage | Exception]
write_stream: MemoryObjectSendStream[types.JSONRPCMessage]
write_stream_reader: MemoryObjectReceiveStream[types.JSONRPCMessage]
read_stream_writer, read_stream = anyio.create_memory_object_stream(0)
write_stream, write_stream_reader = anyio.create_memory_object_stream(0)
session_id: Final = uuid4()
session_uri: Final = f"{quote(self._endpoint)}?session_id={session_id.hex}"
self._read_stream_writers[session_id] = read_stream_writer
verbose_logger.debug("Created new session with ID: %s", session_id)
sse_stream_writer: MemoryObjectSendStream[dict[str, Any]]
sse_stream_reader: MemoryObjectReceiveStream[dict[str, Any]]
sse_stream_writer, sse_stream_reader = anyio.create_memory_object_stream(0, dict[str, Any])
async def sse_writer():
verbose_logger.debug("Starting SSE writer")
async with sse_stream_writer, write_stream_reader:
await sse_stream_writer.send({"event": "endpoint", "data": session_uri})
verbose_logger.debug("Sent endpoint event: %s", session_uri)
async for message in write_stream_reader:
verbose_logger.debug("Sending message via SSE: %s", message)
await sse_stream_writer.send(
{
"event": "message",
"data": message.model_dump_json(by_alias=True, exclude_none=True),
}
)
async with anyio.create_task_group() as tg:
response: Final = EventSourceResponse(content=sse_stream_reader, data_sender_callable=sse_writer)
verbose_logger.debug("Starting SSE response task")
tg.start_soon(response, request.scope, request.receive, request._send)
verbose_logger.debug("Yielding read and write streams")
yield (read_stream, write_stream)
async def handle_post_message(self, scope: Scope, receive: Receive, send: Send) -> Response:
verbose_logger.debug("Handling POST message")
request: Final = Request(scope, receive)
session_id_param: Final = request.query_params.get("session_id")
if session_id_param is None:
verbose_logger.warning("Received request without session_id")
response = Response("session_id is required", status_code=400)
return response
try:
session_id: Final = UUID(hex=session_id_param)
verbose_logger.debug("Parsed session ID: %s", session_id)
except ValueError:
verbose_logger.warning("Received invalid session ID: %s", session_id_param)
response = Response("Invalid session ID", status_code=400)
return response
writer: Final = self._read_stream_writers.get(session_id)
if not writer:
verbose_logger.warning("Could not find session for ID: %s", session_id)
response = Response("Could not find session", status_code=404)
return response
json: Final = await request.json()
verbose_logger.debug("Received JSON: %s", json)
try:
message: Final = types.JSONRPCMessage.model_validate(json)
verbose_logger.debug("Validated client message: %s", message)
except ValidationError as err:
verbose_logger.error("Failed to parse message: %s", err)
response = Response("Could not parse message", status_code=400)
await writer.send(err)
return response
verbose_logger.debug("Sending message to writer: %s", message)
response = Response("Accepted", status_code=202)
await writer.send(message)
return response
__all__ = ("SseServerTransport",)

View file

@ -533,6 +533,9 @@ class LiteLLMRoutes(enum.Enum):
mcp_inference_routes = [
"/mcp",
"/mcp/",
"/mcp/sse/",
"/mcp/sse/messages",
"/mcp/sse/messages/",
"/mcp/proxy",
"/mcp/{subpath}",
"/mcp/tools",

View file

@ -456,9 +456,15 @@ async def test_sse_mcp_handler_mock():
"""Test the SSE MCP handler functionality"""
from litellm.proxy._types import UserAPIKeyAuth
# Mock the SSE session manager and its methods
mock_sse_session_manager = AsyncMock()
mock_sse_session_manager.handle_request = AsyncMock()
read_stream, write_stream = MagicMock(), MagicMock()
@asynccontextmanager
async def connect_sse(scope, receive, send):
yield read_stream, write_stream
mock_sse = MagicMock()
mock_sse.connect_sse.side_effect = connect_sse
run = AsyncMock()
# Mock scope, receive, send with proper ASGI scope format
mock_scope = {
@ -483,13 +489,14 @@ async def test_sse_mcp_handler_mock():
)
with (
patch("litellm.proxy._experimental.mcp_server.server.server.run", run),
patch(
"litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED",
True,
),
patch(
"litellm.proxy._experimental.mcp_server.server.sse_session_manager",
mock_sse_session_manager,
"litellm.proxy._experimental.mcp_server.server.sse",
mock_sse,
),
patch(
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
@ -504,10 +511,8 @@ async def test_sse_mcp_handler_mock():
# Call the handler
await handle_sse_mcp(mock_scope, mock_receive, mock_send)
# Verify SSE session manager handle_request was called
mock_sse_session_manager.handle_request.assert_called_once_with(
mock_scope, mock_receive, mock_send
)
assert run.await_args.args[:2] == (read_stream, write_stream)
assert mock_sse.connect_sse.call_args.args[0]["path"] == "/mcp/sse"
@pytest.mark.asyncio
@ -545,7 +550,10 @@ async def test_sse_mcp_handler_propagates_passthrough_401():
True,
),
patch(
"litellm.proxy._experimental.mcp_server.server.sse_session_manager",
"litellm.proxy._experimental.mcp_server.server.sse",
) as transport,
patch(
"litellm.proxy._experimental.mcp_server.server.server.run",
AsyncMock(),
),
patch(
@ -569,6 +577,7 @@ async def test_sse_mcp_handler_propagates_passthrough_401():
with pytest.raises(HTTPException) as excinfo:
await handle_sse_mcp(mock_scope, mock_receive, mock_send)
transport.connect_sse.assert_not_called()
assert excinfo.value.status_code == 401
assert excinfo.value.headers and "WWW-Authenticate" in excinfo.value.headers

View file

@ -552,6 +552,50 @@ class TestExecuteSessionOperationSurfacesTransportError:
with pytest.raises(asyncio.CancelledError):
await client._execute_session_operation(transport_ctx, _op)
@pytest.mark.asyncio
@pytest.mark.parametrize("failure_phase", ("early", "late", "mixed"))
@patch("litellm.experimental_mcp_client.client.ClientSession")
async def test_response_close_preserves_cancellation_and_original_errors(self, session_class, failure_phase):
closed: Final = asyncio.Event()
close_error: Final = httpx2.ReadError("response close failed")
connect_error: Final = httpx2.ConnectError("another request failed before cancellation")
cancelled: Final = asyncio.CancelledError("caller cancelled")
class FailingCloseStream(httpx2.AsyncByteStream):
async def __aiter__(self) -> AsyncIterator[bytes]:
yield b"pending"
async def aclose(self) -> None:
closed.set()
raise close_error
client: Final = MCPClient(server_url="https://example.com/mcp")
async with client._create_httpx_client_factory(
transport=httpx2.MockTransport(lambda _: httpx2.Response(200, stream=FailingCloseStream()))
)() as http_client:
response: Final = await http_client.send(http_client.build_request("POST", client.server_url), stream=True)
async def initialize():
if failure_phase == "early":
await response.aclose()
raise cancelled
async def close_transport(*args):
if failure_phase == "early":
return
try:
await response.aclose()
except httpx2.ReadError as error:
failures: Final = [error, connect_error] if failure_phase == "mixed" else [error]
raise _FakeExceptionGroup("transport", [_FakeExceptionGroup("reader", failures)])
self._make_session(session_class, initialize)
expected: Final = close_error if failure_phase == "early" else connect_error if failure_phase == "mixed" else cancelled
with pytest.raises(type(expected)) as caught:
await client._execute_session_operation(self._make_transport(close_transport), AsyncMock(), http_client)
assert caught.value is expected
assert closed.is_set()
@pytest.mark.asyncio
@patch("litellm.experimental_mcp_client.client.ClientSession")
async def test_cleanup_error_after_success_is_swallowed(self, mock_session_cls):
@ -568,6 +612,94 @@ class TestExecuteSessionOperationSurfacesTransportError:
assert result == "done"
@pytest.mark.asyncio
@patch("litellm.experimental_mcp_client.client.ClientSession")
async def test_session_entry_failure_still_closes_transport(self, session_class):
failure: Final = RuntimeError("session dispatcher did not start")
session_class.return_value.__aenter__ = AsyncMock(side_effect=failure)
closed: Final = asyncio.Event()
async def close_transport(*args):
await anyio.lowlevel.checkpoint()
closed.set()
transport: Final = self._make_transport(close_transport)
client: Final = MCPClient(server_url="https://example.com/mcp")
with pytest.raises(RuntimeError) as caught:
await client._execute_session_operation(transport, AsyncMock())
assert caught.value is failure
assert closed.is_set()
@pytest.mark.asyncio
@pytest.mark.parametrize("original_error", (False, True))
@patch("litellm.experimental_mcp_client.client.ClientSession")
async def test_session_exit_cancellation_preserves_original_failure(self, session_class, original_error):
self._make_session(session_class, AsyncMock(return_value=None))
cancelled: Final = asyncio.CancelledError("cancelled while closing session")
session_class.return_value.__aexit__ = AsyncMock(side_effect=cancelled)
original: Final = RuntimeError("operation failed")
transport: Final = self._make_transport(None)
client: Final = MCPClient(server_url="https://example.com/mcp")
async def operation(session):
if original_error:
raise original
return "done"
with pytest.raises(RuntimeError if original_error else asyncio.CancelledError) as caught:
await client._execute_session_operation(transport, operation)
assert caught.value is (original if original_error else cancelled)
transport.__aexit__.assert_awaited_once()
@pytest.mark.asyncio
@pytest.mark.parametrize("signal_type", (SystemExit, KeyboardInterrupt))
@pytest.mark.parametrize("phase", ("session", "transport"))
@patch("litellm.experimental_mcp_client.client.ClientSession")
async def test_cleanup_preserves_process_exit(self, session_class, phase, signal_type):
self._make_session(session_class, AsyncMock(return_value=None))
signal: Final = signal_type("process stopping")
if phase == "session":
session_class.return_value.__aexit__ = AsyncMock(side_effect=signal)
transport: Final = self._make_transport(signal if phase == "transport" else None)
client: Final = MCPClient(server_url="https://example.com/mcp")
with pytest.raises(signal_type) as caught:
await client._execute_session_operation(transport, AsyncMock(return_value="done"))
assert caught.value is signal
transport.__aexit__.assert_awaited_once()
@pytest.mark.asyncio
@patch("litellm.experimental_mcp_client.client.ClientSession")
async def test_session_and_termination_share_one_cleanup_deadline(self, session_class):
self._make_session(session_class, AsyncMock(return_value=None))
deleting: Final = asyncio.Event()
async def close_session(*args):
await anyio.sleep(1)
async def respond(request: httpx2.Request) -> httpx2.Response:
deleting.set()
await anyio.sleep_forever()
raise AssertionError("termination unexpectedly resumed")
client: Final = MCPClient(server_url="https://example.com/mcp")
http_client: Final = client._create_httpx_client_factory(transport=httpx2.MockTransport(respond))()
session_class.return_value.__aexit__ = AsyncMock(side_effect=close_session)
async def close_transport(*args):
await http_client.delete(client.server_url)
before: Final = anyio.current_time()
try:
with pytest.raises(asyncio.CancelledError):
await client._execute_session_operation(
self._make_transport(close_transport), AsyncMock(return_value="completed"), http_client=http_client
)
assert deleting.is_set()
assert 4.8 <= anyio.current_time() - before < 5.8
finally:
await http_client.aclose()
class TestMCPClientResolvedAuth:
"""A pre-resolved httpx2.Auth is attached to the upstream client's auth= slot."""
@ -736,7 +868,7 @@ async def test_run_with_session_quiet_on_error_demotes_warning_to_debug():
async def _op(_session):
raise boom
async def _fake_exec(_transport_ctx, _operation):
async def _fake_exec(_transport_ctx, _operation, http_client=None):
raise boom
with patch.object(client, "_create_transport_context", return_value=(object(), None)):
@ -1486,7 +1618,9 @@ async def test_http_response_handler_preserves_success_and_http_errors(status_co
client: Final = MCPClient(server_url="https://example.com/mcp", timeout=30)
async with client._create_httpx_client_factory(transport=httpx2.MockTransport(respond))() as http_client:
operation: Final = client._execute_session_operation(
streamable_http_client(client.server_url, http_client=http_client), lambda session: session.list_tools()
streamable_http_client(client.server_url, http_client=http_client),
lambda session: session.list_tools(),
http_client=http_client,
)
if status_code == 200:
result: Final = await asyncio.wait_for(operation, timeout=3)
@ -1816,13 +1950,14 @@ async def test_interrupted_http_response_preserves_the_transport_failure() -> No
def respond(request: httpx2.Request) -> httpx2.Response:
return httpx2.Response(200, headers={"Content-Type": "application/json"}, stream=_InterruptedHTTPBody())
async with httpx2.AsyncClient(transport=httpx2.MockTransport(respond)) as http_client:
client: Final = MCPClient(server_url="https://example.com/mcp", timeout=30)
client: Final = MCPClient(server_url="https://example.com/mcp", timeout=30)
async with client._create_httpx_client_factory(transport=httpx2.MockTransport(respond))() as http_client:
with pytest.raises(httpx2.RemoteProtocolError, match="secret-incomplete-response"):
await asyncio.wait_for(
client._execute_session_operation(
streamable_http_client(client.server_url, http_client=http_client),
lambda session: session.list_tools(),
http_client=http_client,
),
timeout=3,
)
@ -1906,7 +2041,7 @@ async def test_optional_discovery_capabilities_and_errors(
"jsonrpc": "2.0",
"id": payload.id,
"result": {
"protocolVersion": payload.params["protocolVersion"],
"protocolVersion": (payload.params or {})["protocolVersion"],
"capabilities": {}
if outcome == "absent"
else {advertised if outcome == "other_capability" else capability: {}},
@ -1985,7 +2120,7 @@ async def test_optional_discovery_uses_each_sessions_capabilities(supports_first
return httpx2.Response(202)
result: Final = (
{
"protocolVersion": payload.params["protocolVersion"],
"protocolVersion": (payload.params or {})["protocolVersion"],
"capabilities": next(capabilities),
"serverInfo": {"name": "changing", "version": "1"},
}
@ -2030,7 +2165,7 @@ async def test_optional_discovery_preserves_cancellation(method: str) -> None:
"jsonrpc": "2.0",
"id": payload.id,
"result": {
"protocolVersion": payload.params["protocolVersion"],
"protocolVersion": (payload.params or {})["protocolVersion"],
"capabilities": {"resources": {}, "prompts": {}},
"serverInfo": {"name": "pending", "version": "1"},
},
@ -2112,7 +2247,7 @@ async def test_optional_discovery_collects_all_pages(method: str, session_id: st
"jsonrpc": "2.0",
"id": payload.id,
"result": {
"protocolVersion": payload.params["protocolVersion"],
"protocolVersion": (payload.params or {})["protocolVersion"],
"capabilities": {"prompts": {}, "resources": {}},
"serverInfo": {"name": "paged", "version": "1"},
},
@ -2200,7 +2335,7 @@ async def test_optional_discovery_rejects_incomplete_walks(
"jsonrpc": "2.0",
"id": payload.id,
"result": {
"protocolVersion": payload.params["protocolVersion"],
"protocolVersion": (payload.params or {})["protocolVersion"],
"capabilities": {"prompts": {}, "resources": {}},
"serverInfo": {"name": "interrupted", "version": "1"},
},
@ -2281,7 +2416,7 @@ async def test_optional_discovery_allows_exhaustion_at_page_cap(method: str, mon
return httpx2.Response(202)
if payload.method == "initialize":
result: Final = {
"protocolVersion": payload.params["protocolVersion"],
"protocolVersion": (payload.params or {})["protocolVersion"],
"capabilities": {"prompts": {}, "resources": {}},
"serverInfo": {"name": "empty-pages", "version": "1"},
}
@ -2455,3 +2590,309 @@ def test_public_mcp_import_preserves_incompatible_sdk_error() -> None:
assert not isinstance(caught.value, ModuleNotFoundError)
assert caught.value.__cause__ is None
assert "litellm[mcp]" not in str(caught.value)
@pytest.mark.asyncio
@pytest.mark.parametrize("grouped", (False, True))
@pytest.mark.parametrize("raise_on_error", (False, True))
@pytest.mark.parametrize("termination", ("ok", "failure", "hang"))
async def test_outer_deadline_delivers_session_termination(termination: str, grouped: bool, raise_on_error: bool) -> None:
deleted: Final = asyncio.Event()
started: Final = asyncio.Event()
async def respond(request: httpx2.Request) -> httpx2.Response:
await anyio.lowlevel.checkpoint()
if request.method == "DELETE":
first_termination: Final = not deleted.is_set()
deleted.set()
if termination == "hang" and first_termination:
await anyio.sleep_forever()
return httpx2.Response(500 if termination == "failure" else 200)
if request.method == "GET":
return httpx2.Response(405)
payload: Final = _JSONRPC_MESSAGE_ADAPTER.validate_json(request.content)
if not isinstance(payload, JSONRPCRequest):
return httpx2.Response(202)
if payload.method == "initialize":
return httpx2.Response(
200,
headers={"mcp-session-id": "cancel-owned-session"},
json={
"jsonrpc": "2.0",
"id": payload.id,
"result": {
"protocolVersion": "2025-11-25",
"capabilities": {"tools": {}},
"serverInfo": {"name": "cancellation-peer", "version": "1"},
},
},
)
if payload.method == "tools/list":
return httpx2.Response(200, json={"jsonrpc": "2.0", "id": payload.id, "result": {"tools": []}})
started.set()
await anyio.sleep_forever()
raise AssertionError("cancelled request resumed")
client: Final = _MockTransportClient(respond, server_url="https://example.com/mcp", timeout=30)
async def invoke():
with anyio.fail_after(0.2):
pending: Final = client.call_tool(CallToolRequestParams(name="slow", arguments={}), raise_on_error=raise_on_error)
if grouped:
await asyncio.gather(pending)
else:
await pending
before: Final = anyio.current_time()
with pytest.raises(TimeoutError):
await invoke()
assert started.is_set()
assert deleted.is_set(), "Cancellation must deliver DELETE before returning to the caller"
assert anyio.current_time() - before < 6.5
assert await client.list_tools(raise_on_error=True) == []
@pytest.mark.asyncio
@pytest.mark.parametrize("original_error", (False, True))
async def test_task_cancellation_during_cleanup_preserves_failure(original_error: bool) -> None:
deleting: Final = asyncio.Event()
drained: Final = asyncio.Event()
original: Final = RuntimeError("operation failed before teardown")
async def respond(request: httpx2.Request) -> httpx2.Response:
if request.method == "DELETE":
deleting.set()
try:
await anyio.sleep_forever()
finally:
drained.set()
if request.method == "GET":
return httpx2.Response(405)
payload: Final = _JSONRPC_MESSAGE_ADAPTER.validate_json(request.content)
if not isinstance(payload, JSONRPCRequest):
return httpx2.Response(202)
return httpx2.Response(
200,
headers={"mcp-session-id": "cleanup-session"},
json={
"jsonrpc": "2.0",
"id": payload.id,
"result": {
"protocolVersion": (payload.params or {})["protocolVersion"],
"capabilities": {},
"serverInfo": {"name": "cleanup-peer", "version": "1"},
},
},
)
async def operation(session: mcp_client_module.ClientSession) -> str:
if original_error:
raise original
return "completed"
client: Final = _MockTransportClient(respond, server_url="https://example.com/mcp", timeout=30)
task: Final = asyncio.create_task(client.run_with_session(operation))
await asyncio.wait_for(deleting.wait(), 2)
for _ in range(3):
task.cancel()
await asyncio.sleep(0)
with pytest.raises(RuntimeError if original_error else asyncio.CancelledError) as caught:
await task
assert drained.is_set(), "Caller must wait for termination cleanup to finish"
if original_error:
assert caught.value is original
else:
assert task.cancelled()
@pytest.mark.asyncio
@pytest.mark.parametrize("original_error", (False, True))
@pytest.mark.parametrize("cancel_mode", ("task", "scope"))
async def test_http_close_cancellation_cannot_turn_into_success(original_error: bool, cancel_mode: str) -> None:
closing: Final = asyncio.Event()
drained: Final = asyncio.Event()
original: Final = RuntimeError("failed before HTTP close")
class ClosingHTTPClient(httpx2.AsyncClient):
async def aclose(self) -> None:
closing.set()
try:
await anyio.sleep_forever()
finally:
drained.set()
class ClosingMCPClient(MCPClient):
def _create_transport_context(self):
http_client: Final = ClosingHTTPClient(transport=httpx2.MockTransport(lambda _: httpx2.Response(200)))
return streamable_http_client(self.server_url, http_client=http_client), http_client
async def _execute_session_operation(self, transport_ctx, operation, http_client=None):
if original_error:
raise original
return "completed"
client: Final = ClosingMCPClient(server_url="https://example.com/mcp")
async def invoke() -> str:
with anyio.fail_after(0.05 if cancel_mode == "scope" else None):
return await client.run_with_session(AsyncMock())
task: Final = asyncio.create_task(invoke())
await asyncio.wait_for(closing.wait(), 2)
if cancel_mode == "task":
task.cancel()
cancellation_type: Final = asyncio.CancelledError if cancel_mode == "task" else TimeoutError
with pytest.raises(RuntimeError if original_error else cancellation_type) as caught:
await task
assert drained.is_set(), "Caller must wait for HTTP closure to finish"
if original_error:
assert caught.value is original
elif cancel_mode == "task":
assert task.cancelled()
@pytest.mark.asyncio
@pytest.mark.parametrize("cancel_mode", ("scope", "task", "wait_for", "read_timeout"))
@pytest.mark.parametrize("concurrency", (1, 5))
@pytest.mark.parametrize("termination", ("ok", "hang", "hang_body"))
@pytest.mark.parametrize("raise_on_error", (False, True))
async def test_cancellation_delivers_termination_over_tcp(
cancel_mode: str, concurrency: int, termination: str, raise_on_error: bool
) -> None:
started: Final = asyncio.Event()
terminations: Final[list[bytes]] = []
starts: Final[list[bytes]] = []
stop: Final = asyncio.Event()
connections: Final[list[asyncio.Task[None]]] = []
async def handle_connection(reader: asyncio.StreamReader, writer: asyncio.StreamWriter) -> None:
connection: Final = asyncio.current_task()
assert connection is not None
connections.append(connection)
try:
request_line: Final = await reader.readline()
if not request_line:
return
method: Final = request_line.split()[0]
headers: Final = await reader.readuntil(b"\r\n\r\n")
length: Final = next(
(
int(line.split(b":", 1)[1])
for line in headers.splitlines()
if line.lower().startswith(b"content-length:")
),
0,
)
try:
body: Final = await reader.readexactly(length)
except asyncio.IncompleteReadError:
return
if method == b"DELETE":
terminations.append(body)
if termination != "ok":
await stop.wait()
return
writer.write(b"HTTP/1.1 200 OK\r\nContent-Length: 0\r\nConnection: close\r\n\r\n")
elif method == b"GET":
writer.write(b"HTTP/1.1 405 Method Not Allowed\r\nContent-Length: 0\r\nConnection: close\r\n\r\n")
else:
payload: Final = json.loads(body)
if payload["method"] == "tools/call":
if termination == "hang_body":
writer.write(
b"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\n"
b"Content-Length: 200\r\nConnection: close\r\n\r\n"
)
await writer.drain()
starts.append(body)
if len(starts) == concurrency:
started.set()
await stop.wait()
return
if payload["method"] == "initialize":
response: Final = json.dumps(
{
"jsonrpc": "2.0",
"id": payload["id"],
"result": {
"protocolVersion": "2025-06-18",
"capabilities": {"tools": {}},
"serverInfo": {"name": "tcp-peer", "version": "1"},
},
}
).encode()
writer.write(
b"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nMcp-Session-Id: tcp-session\r\n"
+ f"Content-Length: {len(response)}\r\nConnection: close\r\n\r\n".encode()
+ response
)
else:
writer.write(b"HTTP/1.1 202 Accepted\r\nContent-Length: 0\r\nConnection: close\r\n\r\n")
await writer.drain()
finally:
writer.close()
await writer.wait_closed()
listener: Final = await asyncio.start_server(handle_connection, "127.0.0.1", 0)
port: Final = listener.sockets[0].getsockname()[1]
client: Final = MCPClient(
server_url=f"http://127.0.0.1:{port}/mcp", timeout=2 if cancel_mode == "read_timeout" else 0.5 if termination != "ok" else 30
)
async def calls():
results: Final = await asyncio.gather(
*(
client.call_tool(CallToolRequestParams(name="slow", arguments={}), raise_on_error=raise_on_error)
for _ in range(concurrency)
),
return_exceptions=cancel_mode == "read_timeout",
)
if cancel_mode == "read_timeout":
if raise_on_error:
assert all(isinstance(result, TimeoutError) for result in results)
else:
assert all(isinstance(result, CallToolResult) and result.is_error for result in results)
return results
async def invoke():
if cancel_mode == "scope":
with anyio.fail_after(0.2):
return await calls()
return await calls()
try:
task: Final = asyncio.create_task(invoke())
await asyncio.wait_for(started.wait(), 3)
if cancel_mode == "task":
task.cancel()
expected_error: Final = (
TimeoutError
if cancel_mode == "read_timeout"
else asyncio.CancelledError
if cancel_mode == "task"
else TimeoutError
)
if cancel_mode == "read_timeout":
done, _ = await asyncio.wait((task,), timeout=8)
assert task in done, "Read timeout and bounded cleanup must complete without external cancellation"
await task
elif cancel_mode == "wait_for":
with pytest.raises(expected_error):
await asyncio.wait_for(task, 0.2)
else:
with pytest.raises(expected_error):
await task
assert len(starts) == concurrency
assert len(terminations) == concurrency, "Each cancelled call must send DELETE over a fresh TCP connection"
finally:
stop.set()
if not task.done():
task.cancel()
await asyncio.wait((task,), timeout=8)
listener.close()
for connection in connections:
connection.cancel()
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)

View file

@ -3057,6 +3057,11 @@ class TestMCPDelegateAuthToUpstream:
)
cases = [
("/mcp/sse", []),
("/mcp/sse/", []),
("/mcp/sse/messages", []),
("/mcp/sse/messages/", []),
("/sse/mcp", ["sse"]),
# Single server, single segment.
("/mcp/foo", ["foo"]),
# Server name with one embedded slash (two segments).

View file

@ -1147,109 +1147,39 @@ async def test_mcp_read_resource_success():
assert result is read_result
def test_normalize_resource_contents_passes_metadata():
"""Test that _normalize_resource_contents preserves meta from ResourceContents (MCP 1.26.0+)."""
try:
from litellm.proxy._experimental.mcp_server.server import (
_normalize_resource_contents,
)
except ImportError:
pytest.skip("MCP server not available")
@pytest.mark.asyncio
@pytest.mark.parametrize(
"kind,metadata",
(("text", {"version": "1.0", "source": "test"}), ("blob", {"encoding": "base64"}), ("text", {}), ("text", None)),
)
async def test_read_resource_preserves_content_metadata(_mcp_request_ctx, kind, metadata):
from mcp.types import ReadResourceRequestParams
from litellm.proxy._experimental.mcp_server import operations, server
meta = {"version": "1.0", "source": "test"}
contents = [
TextResourceContents(
uri="https://example.com/resource",
text="hello world",
mimeType="text/plain",
meta=meta,
)
]
uri: Final = "https://example.com/resource"
caller: Final = UserAPIKeyAuth(user_id="resource-caller")
upstream_server: Final = MCPServer(server_id="catalog", name="catalog", transport=MCPTransport.http)
content: Final = (
TextResourceContents(uri=uri, text="hello world", mimeType="text/plain", meta=metadata)
if kind == "text"
else BlobResourceContents(uri=uri, blob="aGVsbG8=", mimeType="image/png", meta=metadata)
)
with (
patch.object(server, "get_or_extract_auth_context", AsyncMock(return_value=(caller, None, ["catalog"], None, None, None, None))),
patch.object(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[upstream_server])),
patch.object(operations.global_mcp_server_manager, "read_resource_from_server", AsyncMock(return_value=ReadResourceResult(contents=[content]))),
):
result: Final = await server.read_resource(_mcp_request_ctx(), ReadResourceRequestParams(uri=uri))
result = _normalize_resource_contents(contents)
assert len(result) == 1
assert result[0].content == "hello world"
assert result[0].mime_type == "text/plain"
assert result[0].meta == meta
def test_normalize_resource_contents_blob_with_metadata():
"""Test that _normalize_resource_contents preserves meta for BlobResourceContents."""
try:
from litellm.proxy._experimental.mcp_server.server import (
_normalize_resource_contents,
)
except ImportError:
pytest.skip("MCP server not available")
meta = {"encoding": "base64"}
contents = [
BlobResourceContents(
uri="https://example.com/image.png",
blob="aGVsbG8=",
mimeType="image/png",
meta=meta,
)
]
result = _normalize_resource_contents(contents)
assert len(result) == 1
assert result[0].content == "aGVsbG8="
assert result[0].mime_type == "image/png"
assert result[0].meta == meta
def test_normalize_resource_contents_preserves_empty_metadata():
"""Test that empty dict meta is preserved (truthiness bug fix)."""
try:
from litellm.proxy._experimental.mcp_server.server import (
_normalize_resource_contents,
)
except ImportError:
pytest.skip("MCP server not available")
empty_meta: dict = {}
contents = [
TextResourceContents(
uri="https://example.com/resource",
text="hi",
mimeType="text/plain",
meta=empty_meta,
)
]
result = _normalize_resource_contents(contents)
assert len(result) == 1
assert result[0].meta == empty_meta
assert result[0].meta is not None
assert result[0].meta == {}
def test_normalize_resource_contents_without_metadata():
"""Test that _normalize_resource_contents works when meta is absent (backward compat)."""
try:
from litellm.proxy._experimental.mcp_server.server import (
_normalize_resource_contents,
)
except ImportError:
pytest.skip("MCP server not available")
contents = [
TextResourceContents(
uri="https://example.com/resource",
text="hello",
mimeType="text/plain",
)
]
result = _normalize_resource_contents(contents)
assert len(result) == 1
assert result[0].content == "hello"
assert result[0].meta is None
assert result.model_dump(mode="json", by_alias=True, exclude_none=True) == {
"cacheScope": "private", "resultType": "complete", "ttlMs": 0,
"contents": [{
"uri": uri,
"mimeType": "text/plain" if kind == "text" else "image/png",
"text" if kind == "text" else "blob": "hello world" if kind == "text" else "aGVsbG8=",
**({"_meta": metadata} if metadata is not None else {}),
}],
}
@pytest.mark.asyncio
@ -1859,14 +1789,12 @@ async def test_concurrent_initialize_session_managers():
original_initialized = mcp_server._SESSION_MANAGERS_INITIALIZED
original_session_cm = mcp_server._session_manager_cm
original_stateful_cm = mcp_server._session_manager_stateful_cm
original_sse_cm = mcp_server._sse_session_manager_cm
original_cleanup_task = mcp_server._stateful_auth_context_cleanup_task
try:
mcp_server._SESSION_MANAGERS_INITIALIZED = False
mcp_server._session_manager_cm = None
mcp_server._session_manager_stateful_cm = None
mcp_server._sse_session_manager_cm = None
# Create mock context managers for all three session managers
mock_cm_stateless = AsyncMock()
@ -1877,10 +1805,6 @@ async def test_concurrent_initialize_session_managers():
mock_cm_stateful.__aenter__ = AsyncMock()
mock_cm_stateful.__aexit__ = AsyncMock()
mock_cm_sse = AsyncMock()
mock_cm_sse.__aenter__ = AsyncMock()
mock_cm_sse.__aexit__ = AsyncMock()
with (
patch.object(
mcp_server.session_manager_stateless,
@ -1892,11 +1816,6 @@ async def test_concurrent_initialize_session_managers():
"run",
return_value=mock_cm_stateful,
) as mock_stateful_run,
patch.object(
mcp_server.sse_session_manager,
"run",
return_value=mock_cm_sse,
) as mock_sse_run,
patch("litellm.proxy._experimental.mcp_server.operations.verbose_logger"),
):
# Create multiple concurrent tasks that call initialize_session_managers
@ -1918,10 +1837,6 @@ async def test_concurrent_initialize_session_managers():
assert mock_stateful_run.call_count == 1, (
f"Expected 1 call to session_manager_stateful.run(), got {mock_stateful_run.call_count}"
)
assert mock_sse_run.call_count == 1, (
f"Expected 1 call to sse_session_manager.run(), got {mock_sse_run.call_count}"
)
# The context managers should only be entered once each
assert mock_cm_stateless.__aenter__.call_count == 1, (
f"Expected 1 call to stateless __aenter__, got {mock_cm_stateless.__aenter__.call_count}"
@ -1929,10 +1844,6 @@ async def test_concurrent_initialize_session_managers():
assert mock_cm_stateful.__aenter__.call_count == 1, (
f"Expected 1 call to stateful __aenter__, got {mock_cm_stateful.__aenter__.call_count}"
)
assert mock_cm_sse.__aenter__.call_count == 1, (
f"Expected 1 call to sse __aenter__, got {mock_cm_sse.__aenter__.call_count}"
)
# State should be properly set
assert mcp_server._SESSION_MANAGERS_INITIALIZED is True
@ -1948,7 +1859,6 @@ async def test_concurrent_initialize_session_managers():
mcp_server._SESSION_MANAGERS_INITIALIZED = original_initialized
mcp_server._session_manager_cm = original_session_cm
mcp_server._session_manager_stateful_cm = original_stateful_cm
mcp_server._sse_session_manager_cm = original_sse_cm
mcp_server._stateful_auth_context_cleanup_task = original_cleanup_task
@ -2263,7 +2173,7 @@ async def test_sse_endpoint_applies_the_same_client_allowlist(
new_callable=AsyncMock,
),
patch.object( # test-quality-ok: SSE manager is a module singleton; the downstream call is the observable
mcp_module.sse_session_manager, "handle_request", side_effect=handle_request
mcp_module.sse, "handle_post_message", side_effect=handle_request
),
):
if admitted:
@ -6651,17 +6561,22 @@ class TestGatewayCreateInitializationOptions:
)
captured = {}
async def record_request(scope, receive, send):
@contextlib.asynccontextmanager
async def connect_sse(scope, receive, send):
yield (None, None)
async def record_request(read_stream, write_stream, options):
captured["server_name"] = server.create_initialization_options().server_name
scope = {
"type": "http",
"method": "POST",
"method": "GET",
"path": "/mcp/grafana",
"headers": [],
}
with (
patch.object(mcp_server.sse, "connect_sse", connect_sse),
patch(
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
new_callable=AsyncMock,
@ -6697,8 +6612,8 @@ class TestGatewayCreateInitializationOptions:
True,
),
patch.object(
mcp_server.sse_session_manager,
"handle_request",
mcp_server.server,
"run",
side_effect=record_request,
),
):
@ -10260,7 +10175,10 @@ async def test_active_request_ctx_var_feeds_auth_resolution_recording(_mcp_reque
("1999-01-01", True),
],
)
async def test_streamable_http_rejects_modern_protocol_version(header_value: str, expected_rejected: bool) -> None:
@pytest.mark.parametrize("handler", ("handle_streamable_http_mcp", "handle_sse_mcp"))
async def test_streamable_http_rejects_modern_protocol_version(
header_value: str, expected_rejected: bool, handler: str
) -> None:
from litellm.proxy._experimental.mcp_server import server as mcp_module
from litellm.proxy._experimental.mcp_server.server import unsupported_protocol_version
@ -10283,7 +10201,7 @@ async def test_streamable_http_rejects_modern_protocol_version(header_value: str
async def send(message: Message) -> None:
sent.append(message)
await mcp_module.handle_streamable_http_mcp(scope, receive, send)
await getattr(mcp_module, handler)(scope, receive, send)
start = next(m for m in sent if m["type"] == "http.response.start")
assert start["status"] == 400
@ -10333,3 +10251,124 @@ async def test_tool_listing_preserves_permission_denial_when_failure_logging_fai
logger.post_call_failure_hook.assert_awaited_once()
assert logger.post_call_failure_hook.await_args.kwargs["original_exception"] is denial
assert logger.post_call_failure_hook.await_args.kwargs["user_api_key_dict"] == auth
@pytest.mark.asyncio
@pytest.mark.parametrize("prefix,suffix", (("", ""), ("/gateway", "/")))
async def test_legacy_sse_mount_emits_message_endpoint(prefix: str, suffix: str) -> None:
from starlette.applications import Starlette
from starlette.routing import Mount
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import AggregateToolListing
from litellm.proxy._experimental.mcp_server import server as mcp_server
app: Final = Starlette(routes=[Mount("/mcp", app=mcp_server.app)])
incoming: Final[asyncio.Queue[Message]] = asyncio.Queue()
outgoing: Final[asyncio.Queue[Message]] = asyncio.Queue()
await incoming.put({"type": "http.request", "body": b"", "more_body": False})
path: Final = f"{prefix}/mcp/sse{suffix}"
scope: Final[Scope] = {
"type": "http",
"asgi": {"version": "3.0"},
"http_version": "1.1",
"method": "GET",
"scheme": "http",
"path": path,
"raw_path": path.encode(),
"query_string": b"",
"root_path": prefix,
"server": ("localhost", 80),
"client": ("127.0.0.1", 1234),
"headers": [(b"accept", b"text/event-stream")],
}
auth: Final = UserAPIKeyAuth(api_key="test-owner")
with (
patch.object(
mcp_server, "extract_mcp_auth_context", AsyncMock(return_value=(auth, None, None, None, None, None))
),
patch.object(mcp_server, "_raise_preemptive_401_for_unauthenticated_servers", AsyncMock()),
patch.object(mcp_server, "_check_passthrough_upstream_auth", AsyncMock()),
patch.object(mcp_server.operations, "_raise_if_initialize_grants_no_mcp_servers", AsyncMock()),
patch.object(mcp_server, "_SESSION_MANAGERS_INITIALIZED", True),
):
task: Final = asyncio.create_task(app(scope, incoming.get, outgoing.put))
try:
start: Final = await asyncio.wait_for(outgoing.get(), 2)
assert start["type"] == "http.response.start"
assert start["status"] == 200
endpoint_frame: Final = await asyncio.wait_for(outgoing.get(), 2)
frame: Final = endpoint_frame["body"].decode()
assert "event: endpoint" in frame
endpoint: Final = frame.split("data: ", 1)[1].splitlines()[0]
assert endpoint.startswith(f"{prefix}/mcp/sse/messages?session_id=")
message_path, query = endpoint.split("?", 1)
async def post(body: bytes) -> int:
messages: Final[asyncio.Queue[Message]] = asyncio.Queue()
requests: Final[asyncio.Queue[Message]] = asyncio.Queue()
await requests.put({"type": "http.request", "body": body, "more_body": False})
post_scope: Final[Scope] = {
**scope,
"method": "POST",
"path": message_path + suffix,
"raw_path": (message_path + suffix).encode(),
"query_string": query.encode(),
"root_path": prefix,
"headers": [(b"content-type", b"application/json")],
}
await asyncio.wait_for(app(post_scope, requests.get, messages.put), 2)
return (await messages.get())["status"]
initialization: Final = json.dumps(
{
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": LATEST_HANDSHAKE_VERSION,
"capabilities": {},
"clientInfo": {"name": "legacy-client", "version": "1"},
},
}
).encode()
assert await post(initialization) == 202
reply: Final = (await asyncio.wait_for(outgoing.get(), 2))["body"].decode()
initialized: Final = json.loads(reply.split("data: ", 1)[1].splitlines()[0])
assert initialized["id"] == 1
assert initialized["result"]["serverInfo"]["name"] == "litellm-mcp-server"
assert await post(b'{"jsonrpc":"2.0","method":"notifications/initialized"}') == 202
for request_id, marker in ((2, "first-post"), (3, "second-post")):
post_auth: Final = UserAPIKeyAuth(api_key="test-owner", user_id=marker)
listing: Final = AsyncMock(return_value=AggregateToolListing(tools=[], outcomes={}))
with (
patch.object(
mcp_server,
"extract_mcp_auth_context",
AsyncMock(return_value=(post_auth, None, [marker], {marker: {"Authorization": marker}}, {"Authorization": marker}, {"x-request-marker": marker})),
),
patch.object(mcp_server.operations, "_get_tools_from_mcp_servers", listing),
):
assert (
await post(json.dumps({"jsonrpc": "2.0", "id": request_id, "method": "tools/list"}).encode())
== 202
)
listed_frame: Final = (await asyncio.wait_for(outgoing.get(), 2))["body"].decode()
listed: Final = json.loads(listed_frame.split("data: ", 1)[1].splitlines()[0])
assert listed["id"] == request_id
assert listed["result"]["tools"] == []
listing.assert_awaited_once()
assert listing.await_args.kwargs["user_api_key_auth"].user_id == marker
assert listing.await_args.kwargs["mcp_servers"] == [marker]
assert listing.await_args.kwargs["mcp_server_auth_headers"] == {marker: {"Authorization": marker}}
assert listing.await_args.kwargs["oauth2_headers"] == {"Authorization": marker}
assert listing.await_args.kwargs["raw_headers"] == {"x-request-marker": marker}
stranger: Final = UserAPIKeyAuth(api_key="different-owner")
with patch.object(
mcp_server, "extract_mcp_auth_context", AsyncMock(return_value=(stranger, None, None, None, None, None))
):
assert await post(initialization) == 404
finally:
await incoming.put({"type": "http.disconnect"})
await asyncio.wait_for(task, 2)
assert await post(initialization) == 404

View file

@ -653,6 +653,10 @@ def test_virtual_key_llm_api_routes_denies_spend_logs_v2():
"/mcp-rest/tools/call",
"/mcp/tools/list",
"/token",
"/mcp/sse",
"/mcp/sse/",
"/mcp/sse/messages",
"/mcp/sse/messages/",
],
)
def test_mcp_inference_routes_classified_as_llm_api(route):
@ -4310,3 +4314,26 @@ def test_project_delete_route_stays_proxy_admin_only():
valid_token=valid_token,
request_data={},
)
@pytest.mark.parametrize("route", ("/mcp/sse", "/mcp/sse/", "/mcp/sse/messages", "/mcp/sse/messages/"))
@pytest.mark.parametrize("route_group", ("mcp_routes", "llm_api_routes", "openai_routes"))
def test_legacy_sse_respects_virtual_key_route_permissions(route: str, route_group: str) -> None:
token: Final = UserAPIKeyAuth(
user_id="sse-caller", user_role=LitellmUserRoles.INTERNAL_USER, allowed_routes=[route_group]
)
request: Final = Request({"type": "http", "method": "POST" if "messages" in route else "GET", "path": route})
if route_group == "openai_routes":
with pytest.raises(HTTPException) as caught:
RouteChecks.is_virtual_key_allowed_to_call_route(route=route, valid_token=token, request=request)
assert caught.value.status_code == 403
return
RouteChecks.non_proxy_admin_allowed_routes_check(
user_obj=None,
_user_role=LitellmUserRoles.INTERNAL_USER,
route=route,
request=request,
valid_token=token,
request_data={},
)
assert RouteChecks.is_virtual_key_allowed_to_call_route(route=route, valid_token=token, request=request)