diff --git a/litellm/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py index 49434befd4e..1ccae8de35f 100644 --- a/litellm/experimental_mcp_client/client.py +++ b/litellm/experimental_mcp_client/client.py @@ -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: """ diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index 60dc91a69cc..d6478421834 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -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]] diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 433b693fcae..fe508fc22c7 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -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) ######################################################## diff --git a/litellm/proxy/_experimental/mcp_server/sse_transport.py b/litellm/proxy/_experimental/mcp_server/sse_transport.py index 2a08a5f8f7a..839d2eebbbe 100644 --- a/litellm/proxy/_experimental/mcp_server/sse_transport.py +++ b/litellm/proxy/_experimental/mcp_server/sse_transport.py @@ -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",) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 8b2c81fea77..76df8790b92 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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", diff --git a/tests/mcp_tests/test_mcp_server.py b/tests/mcp_tests/test_mcp_server.py index 2b92367f186..d38a11ec927 100644 --- a/tests/mcp_tests/test_mcp_server.py +++ b/tests/mcp_tests/test_mcp_server.py @@ -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 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 6c20ef135ba..b260d240f29 100644 --- a/tests/test_litellm/experimental_mcp_client/test_mcp_client.py +++ b/tests/test_litellm/experimental_mcp_client/test_mcp_client.py @@ -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) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index 35e7055bbc0..0012be23428 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -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). diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index b715fe67e20..6588ab4c87b 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -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 diff --git a/tests/test_litellm/proxy/auth/test_route_checks.py b/tests/test_litellm/proxy/auth/test_route_checks.py index 8da93ba341b..e5179387f82 100644 --- a/tests/test_litellm/proxy/auth/test_route_checks.py +++ b/tests/test_litellm/proxy/auth/test_route_checks.py @@ -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)