mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-25 01:02:15 +00:00
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:
parent
2ea43214c9
commit
0d0f73cd14
10 changed files with 900 additions and 380 deletions
|
|
@ -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:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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]]
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
########################################################
|
||||
|
|
|
|||
|
|
@ -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",)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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).
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue