mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
feat(mcp): core sampling and elicitation flow with security hardening
- Add sampling_handler.py: full MCP sampling/createMessage flow with model selection (hint-based + priority-based), auth enforcement, budget checks, route restriction gates, and tag policy pre-auth - Add elicitation_handler.py: MCP elicitation/create relay with downstream client capability detection - Wire sampling/elicitation callbacks in mcp_server_manager.py gated behind allow_sampling/allow_elicitation config flags - Add allow_sampling/allow_elicitation fields to MCPServer type - Fix session lock deadlock: skip lock for JSON-RPC response POSTs (elicitation/sampling replies) with truncated-body heuristic - Extend client.py with sampling_callback and elicitation_callback - Security: RouteChecks gate, tag-budget bypass fix, x-forwarded-for spoofing fix, Latin-1 header encoding guard - Add 4 new test modules (model access, priority selection, request builder, tool conversion) + update existing MCP tests
This commit is contained in:
parent
d45e9e4d56
commit
695cde67c7
14 changed files with 3128 additions and 335 deletions
|
|
@ -4,6 +4,7 @@ LiteLLM Proxy uses this MCP Client to connnect to other MCP servers.
|
|||
|
||||
import asyncio
|
||||
import base64
|
||||
import os
|
||||
from typing import (
|
||||
Any,
|
||||
Awaitable,
|
||||
|
|
@ -16,7 +17,6 @@ from typing import (
|
|||
TypeVar,
|
||||
Union,
|
||||
)
|
||||
|
||||
import httpx
|
||||
from mcp import ClientSession, ReadResourceResult, Resource, StdioServerParameters
|
||||
from mcp.client.sse import sse_client
|
||||
|
|
@ -42,9 +42,8 @@ from mcp.types import (
|
|||
)
|
||||
from mcp.types import Tool as MCPTool
|
||||
from pydantic import AnyUrl
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import MCP_CLIENT_TIMEOUT
|
||||
from litellm.constants import MCP_CLIENT_TIMEOUT, MCP_NPM_CACHE_DIR
|
||||
from litellm.llms.custom_httpx.http_handler import get_ssl_configuration
|
||||
from litellm.types.llms.custom_http import VerifyTypes
|
||||
from litellm.types.mcp import (
|
||||
|
|
@ -67,7 +66,6 @@ TSessionResult = TypeVar("TSessionResult")
|
|||
class MCPSigV4Auth(httpx.Auth):
|
||||
"""
|
||||
httpx Auth class that signs each request with AWS SigV4.
|
||||
|
||||
This is used for MCP servers that require AWS SigV4 authentication,
|
||||
such as AWS Bedrock AgentCore MCP servers. httpx calls auth_flow()
|
||||
for every outgoing request, enabling per-request signature computation.
|
||||
|
|
@ -92,10 +90,8 @@ class MCPSigV4Auth(httpx.Auth):
|
|||
"Missing botocore to use AWS SigV4 authentication. "
|
||||
"Run 'pip install boto3'."
|
||||
)
|
||||
|
||||
self.service_name = aws_service_name or "bedrock-agentcore"
|
||||
self.region_name = aws_region_name or "us-east-1"
|
||||
|
||||
# Note: os.environ/ prefixed values are already resolved by
|
||||
# ProxyConfig._check_for_os_environ_vars() at config load time.
|
||||
# Values arrive here as plain strings.
|
||||
|
|
@ -143,20 +139,17 @@ class MCPSigV4Auth(httpx.Auth):
|
|||
session_name = (
|
||||
aws_session_name or f"litellm-mcp-{int(__import__('time').time())}"
|
||||
)
|
||||
|
||||
sts_kwargs: dict = {"region_name": aws_region_name}
|
||||
if aws_access_key_id and aws_secret_access_key:
|
||||
sts_kwargs["aws_access_key_id"] = aws_access_key_id
|
||||
sts_kwargs["aws_secret_access_key"] = aws_secret_access_key
|
||||
if aws_session_token:
|
||||
sts_kwargs["aws_session_token"] = aws_session_token
|
||||
|
||||
sts_client = boto3.client("sts", **sts_kwargs)
|
||||
sts_response = sts_client.assume_role(
|
||||
RoleArn=aws_role_name,
|
||||
RoleSessionName=session_name,
|
||||
)
|
||||
|
||||
sts_creds = sts_response["Credentials"]
|
||||
return Credentials(
|
||||
access_key=sts_creds["AccessKeyId"],
|
||||
|
|
@ -178,17 +171,14 @@ class MCPSigV4Auth(httpx.Auth):
|
|||
data=request.content,
|
||||
headers=dict(request.headers),
|
||||
)
|
||||
|
||||
# Sign the request — SigV4Auth.add_auth() adds Authorization,
|
||||
# X-Amz-Date, and X-Amz-Security-Token (if session token present).
|
||||
# Host header is derived automatically from the URL.
|
||||
sigv4 = SigV4Auth(self.credentials, self.service_name, self.region_name)
|
||||
sigv4.add_auth(aws_request)
|
||||
|
||||
# Copy SigV4 headers back to the httpx request
|
||||
for header_name, header_value in aws_request.headers.items():
|
||||
request.headers[header_name] = header_value
|
||||
|
||||
yield request
|
||||
|
||||
|
||||
|
|
@ -198,6 +188,8 @@ class MCPClient:
|
|||
SSE and HTTP transports
|
||||
Authentication via Bearer token, Basic Auth, or API Key
|
||||
Tool calling with error handling and result parsing
|
||||
Sampling callbacks for upstream server LLM requests
|
||||
Elicitation callbacks for upstream server user-input requests
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
|
|
@ -211,6 +203,9 @@ class MCPClient:
|
|||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
ssl_verify: Optional[VerifyTypes] = None,
|
||||
aws_auth: Optional[httpx.Auth] = None,
|
||||
sampling_callback: Optional[Callable] = None,
|
||||
elicitation_callback: Optional[Callable] = None,
|
||||
logging_callback: Optional[Callable] = None,
|
||||
):
|
||||
self.server_url: str = server_url
|
||||
self.transport_type: MCPTransport = transport_type
|
||||
|
|
@ -222,6 +217,9 @@ class MCPClient:
|
|||
self.ssl_verify: Optional[VerifyTypes] = ssl_verify
|
||||
self._aws_auth: Optional[httpx.Auth] = aws_auth
|
||||
self._last_initialize_instructions: Optional[str] = None
|
||||
self._sampling_callback: Optional[Callable] = sampling_callback
|
||||
self._elicitation_callback: Optional[Callable] = elicitation_callback
|
||||
self._logging_callback: Optional[Callable] = logging_callback
|
||||
# handle the basic auth value if provided
|
||||
if auth_value:
|
||||
self.update_auth_value(auth_value)
|
||||
|
|
@ -231,23 +229,20 @@ class MCPClient:
|
|||
) -> Tuple[Any, Optional[httpx.AsyncClient]]:
|
||||
"""
|
||||
Create the appropriate transport context based on transport type.
|
||||
|
||||
Returns:
|
||||
Tuple of (transport_context, http_client).
|
||||
http_client is only set for HTTP transport and needs cleanup.
|
||||
"""
|
||||
http_client: Optional[httpx.AsyncClient] = None
|
||||
|
||||
if self.transport_type == MCPTransport.stdio:
|
||||
if not self.stdio_config:
|
||||
raise ValueError("stdio_config is required for stdio transport")
|
||||
server_params = StdioServerParameters(
|
||||
command=self.stdio_config.get("command", ""),
|
||||
args=self.stdio_config.get("args", []),
|
||||
env=self.stdio_config.get("env", {}),
|
||||
env=self._get_safe_stdio_env(self.stdio_config.get("env")),
|
||||
)
|
||||
return stdio_client(server_params), None
|
||||
|
||||
if self.transport_type == MCPTransport.sse:
|
||||
headers = self._get_auth_headers()
|
||||
httpx_client_factory = self._create_httpx_client_factory()
|
||||
|
|
@ -260,14 +255,12 @@ class MCPClient:
|
|||
),
|
||||
None,
|
||||
)
|
||||
|
||||
# HTTP transport (default)
|
||||
if streamable_http_client is None:
|
||||
raise ImportError(
|
||||
"streamable_http_client is not available. "
|
||||
"Please install mcp with HTTP support."
|
||||
)
|
||||
|
||||
headers = self._get_auth_headers()
|
||||
httpx_client_factory = self._create_httpx_client_factory()
|
||||
verbose_logger.debug("litellm headers for streamable_http_client: %s", headers)
|
||||
|
|
@ -281,6 +274,54 @@ class MCPClient:
|
|||
)
|
||||
return transport_ctx, http_client
|
||||
|
||||
def _get_safe_stdio_env(
|
||||
self, provided_env: Optional[Dict[str, str]]
|
||||
) -> Optional[Dict[str, str]]:
|
||||
"""
|
||||
Return a safe environment for the stdio subprocess.
|
||||
|
||||
If provided_env is set, we use it as-is.
|
||||
If provided_env is None, we return a minimal allowlist from the parent environment
|
||||
to avoid leaking sensitive LiteLLM keys (OPENAI_API_KEY, etc.) to sub-processes.
|
||||
"""
|
||||
if provided_env is not None:
|
||||
return provided_env
|
||||
|
||||
# Minimal allowlist of safe/standard environment variables
|
||||
safe_keys = {
|
||||
"PATH",
|
||||
"HOME",
|
||||
"USER",
|
||||
"LOGNAME",
|
||||
"TMPDIR",
|
||||
"TMP",
|
||||
"TEMP",
|
||||
"SHELL",
|
||||
"LANG",
|
||||
"LC_ALL",
|
||||
# Node/Package manager caches
|
||||
"NPM_CONFIG_CACHE",
|
||||
"PNPM_HOME",
|
||||
"XDG_CACHE_HOME",
|
||||
"XDG_CONFIG_HOME",
|
||||
"XDG_DATA_HOME",
|
||||
# System info
|
||||
"SYSTEMROOT",
|
||||
"COMSPEC",
|
||||
"PATHEXT",
|
||||
"WINDIR",
|
||||
}
|
||||
|
||||
safe_env = {}
|
||||
for key in safe_keys:
|
||||
if key in os.environ:
|
||||
safe_env[key] = os.environ[key]
|
||||
|
||||
if "NPM_CONFIG_CACHE" not in safe_env:
|
||||
safe_env["NPM_CONFIG_CACHE"] = MCP_NPM_CACHE_DIR
|
||||
|
||||
return safe_env
|
||||
|
||||
async def _execute_session_operation(
|
||||
self,
|
||||
transport_ctx: Any,
|
||||
|
|
@ -288,13 +329,23 @@ class MCPClient:
|
|||
) -> TSessionResult:
|
||||
"""
|
||||
Execute an operation within a transport and session context.
|
||||
|
||||
Handles entering/exiting contexts and running the operation.
|
||||
Passes sampling/elicitation/logging callbacks to the ClientSession
|
||||
so that upstream MCP servers can request LLM inference (sampling),
|
||||
user input (elicitation), or send log messages.
|
||||
"""
|
||||
transport = await transport_ctx.__aenter__()
|
||||
try:
|
||||
read_stream, write_stream = transport[0], transport[1]
|
||||
session_ctx = ClientSession(read_stream, write_stream)
|
||||
# Build session kwargs with optional callbacks
|
||||
session_kwargs: 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
|
||||
session_ctx = ClientSession(read_stream, write_stream, **session_kwargs)
|
||||
session = await session_ctx.__aenter__()
|
||||
try:
|
||||
init_result = await session.initialize()
|
||||
|
|
@ -351,7 +402,6 @@ class MCPClient:
|
|||
def _get_auth_headers(self) -> dict:
|
||||
"""Generate authentication headers based on auth type."""
|
||||
headers = {}
|
||||
|
||||
if self._mcp_auth_value:
|
||||
if isinstance(self._mcp_auth_value, str):
|
||||
if self.auth_type == MCPAuth.bearer_token:
|
||||
|
|
@ -373,17 +423,14 @@ class MCPClient:
|
|||
# Note: aws_sigv4 auth is not handled here — SigV4 requires per-request
|
||||
# signing (including the body hash), so it uses httpx.Auth flow instead
|
||||
# of static headers. See MCPSigV4Auth and _create_httpx_client_factory().
|
||||
|
||||
# update the headers with the extra headers
|
||||
if self.extra_headers:
|
||||
headers.update(self.extra_headers)
|
||||
|
||||
return headers
|
||||
|
||||
def _create_httpx_client_factory(self) -> Callable[..., httpx.AsyncClient]:
|
||||
"""
|
||||
Create a custom httpx client factory that uses LiteLLM's SSL configuration.
|
||||
|
||||
This factory follows the same CA bundle path logic as http_handler.py:
|
||||
1. Check ssl_verify parameter (can be SSLContext, bool, or path to CA bundle)
|
||||
2. Check SSL_VERIFY environment variable
|
||||
|
|
@ -400,17 +447,14 @@ class MCPClient:
|
|||
"""Create an httpx.AsyncClient with LiteLLM's SSL configuration."""
|
||||
# Get unified SSL configuration using the same logic as http_handler.py
|
||||
ssl_config = get_ssl_configuration(self.ssl_verify)
|
||||
|
||||
verbose_logger.debug(
|
||||
f"MCP client using SSL configuration: {type(ssl_config).__name__}"
|
||||
)
|
||||
|
||||
# Use SigV4 auth if configured and no explicit auth provided.
|
||||
# The MCP SDK's sse_client and streamable_http_client call this
|
||||
# factory without passing auth=, so self._aws_auth is used.
|
||||
# For non-SigV4 clients, self._aws_auth is None — no behavior change.
|
||||
effective_auth = auth if auth is not None else self._aws_auth
|
||||
|
||||
return httpx.AsyncClient(
|
||||
headers=headers,
|
||||
timeout=timeout,
|
||||
|
|
@ -458,7 +502,6 @@ class MCPClient:
|
|||
f"Server: {self.server_url or 'stdio'}, "
|
||||
f"Transport: {self.transport_type}"
|
||||
)
|
||||
|
||||
# Check if it's a stream/connection error
|
||||
if "BrokenResourceError" in error_type or "Broken" in error_type:
|
||||
verbose_logger.error(
|
||||
|
|
@ -491,7 +534,6 @@ class MCPClient:
|
|||
f"MCP Tool '{call_tool_request_params.name}' progress: "
|
||||
f"{progress}/{total} ({percentage:.0f}%) - {message or ''}"
|
||||
)
|
||||
|
||||
# Forward to Host if callback provided
|
||||
if host_progress_callback:
|
||||
try:
|
||||
|
|
@ -521,7 +563,6 @@ class MCPClient:
|
|||
|
||||
error_trace = traceback.format_exc()
|
||||
verbose_logger.debug(f"MCP client tool call traceback:\n{error_trace}")
|
||||
|
||||
# Log detailed error information
|
||||
error_type = type(e).__name__
|
||||
verbose_logger.error(
|
||||
|
|
@ -532,14 +573,12 @@ class MCPClient:
|
|||
f"Server: {self.server_url or 'stdio'}, "
|
||||
f"Transport: {self.transport_type}"
|
||||
)
|
||||
|
||||
# Check if it's a stream/connection error
|
||||
if "BrokenResourceError" in error_type or "Broken" in error_type:
|
||||
verbose_logger.error(
|
||||
"MCP client detected broken connection/stream - "
|
||||
"the MCP server may have crashed, disconnected, or timed out."
|
||||
)
|
||||
|
||||
# Return a default error result instead of raising
|
||||
return MCPCallToolResult(
|
||||
content=[
|
||||
|
|
@ -577,14 +616,12 @@ class MCPClient:
|
|||
f"Server: {self.server_url or 'stdio'}, "
|
||||
f"Transport: {self.transport_type}"
|
||||
)
|
||||
|
||||
# Check if it's a stream/connection error
|
||||
if "BrokenResourceError" in error_type or "Broken" in error_type:
|
||||
verbose_logger.error(
|
||||
"MCP client detected broken connection/stream during list_tools - "
|
||||
"the MCP server may have crashed, disconnected, or timed out"
|
||||
)
|
||||
|
||||
# Return empty list instead of raising to allow graceful degradation
|
||||
return []
|
||||
|
||||
|
|
@ -617,7 +654,6 @@ class MCPClient:
|
|||
|
||||
error_trace = traceback.format_exc()
|
||||
verbose_logger.debug(f"MCP client get_prompt traceback:\n{error_trace}")
|
||||
|
||||
# Log detailed error information
|
||||
error_type = type(e).__name__
|
||||
verbose_logger.error(
|
||||
|
|
@ -628,14 +664,12 @@ class MCPClient:
|
|||
f"Server: {self.server_url or 'stdio'}, "
|
||||
f"Transport: {self.transport_type}"
|
||||
)
|
||||
|
||||
# Check if it's a stream/connection error
|
||||
if "BrokenResourceError" in error_type or "Broken" in error_type:
|
||||
verbose_logger.error(
|
||||
"MCP client detected broken connection/stream during get_prompt - "
|
||||
"the MCP server may have crashed, disconnected, or timed out."
|
||||
)
|
||||
|
||||
raise
|
||||
|
||||
async def list_resources(self) -> list[Resource]:
|
||||
|
|
@ -667,14 +701,12 @@ class MCPClient:
|
|||
f"Server: {self.server_url or 'stdio'}, "
|
||||
f"Transport: {self.transport_type}"
|
||||
)
|
||||
|
||||
# Check if it's a stream/connection error
|
||||
if "BrokenResourceError" in error_type or "Broken" in error_type:
|
||||
verbose_logger.error(
|
||||
"MCP client detected broken connection/stream during list_resources - "
|
||||
"the MCP server may have crashed, disconnected, or timed out"
|
||||
)
|
||||
|
||||
# Return empty list instead of raising to allow graceful degradation
|
||||
return []
|
||||
|
||||
|
|
@ -709,14 +741,12 @@ class MCPClient:
|
|||
f"Server: {self.server_url or 'stdio'}, "
|
||||
f"Transport: {self.transport_type}"
|
||||
)
|
||||
|
||||
# Check if it's a stream/connection error
|
||||
if "BrokenResourceError" in error_type or "Broken" in error_type:
|
||||
verbose_logger.error(
|
||||
"MCP client detected broken connection/stream during list_resource_templates - "
|
||||
"the MCP server may have crashed, disconnected, or timed out"
|
||||
)
|
||||
|
||||
# Return empty list instead of raising to allow graceful degradation
|
||||
return []
|
||||
|
||||
|
|
@ -742,7 +772,6 @@ class MCPClient:
|
|||
|
||||
error_trace = traceback.format_exc()
|
||||
verbose_logger.debug(f"MCP client read_resource traceback:\n{error_trace}")
|
||||
|
||||
# Log detailed error information
|
||||
error_type = type(e).__name__
|
||||
verbose_logger.error(
|
||||
|
|
@ -753,12 +782,10 @@ class MCPClient:
|
|||
f"Server: {self.server_url or 'stdio'}, "
|
||||
f"Transport: {self.transport_type}"
|
||||
)
|
||||
|
||||
# Check if it's a stream/connection error
|
||||
if "BrokenResourceError" in error_type or "Broken" in error_type:
|
||||
verbose_logger.error(
|
||||
"MCP client detected broken connection/stream during read_resource - "
|
||||
"the MCP server may have crashed, disconnected, or timed out."
|
||||
)
|
||||
|
||||
raise
|
||||
|
|
|
|||
163
litellm/proxy/_experimental/mcp_server/elicitation_handler.py
Normal file
163
litellm/proxy/_experimental/mcp_server/elicitation_handler.py
Normal file
|
|
@ -0,0 +1,163 @@
|
|||
"""
|
||||
MCP Elicitation Handler
|
||||
Handles `elicitation/create` requests from upstream MCP servers by either:
|
||||
1. Relaying them to the connected downstream MCP client (if it supports elicitation)
|
||||
2. Returning a decline/error response (if no downstream client or unsupported)
|
||||
Supports both Form mode (structured data collection) and URL mode (external URL
|
||||
navigation for sensitive interactions like OAuth).
|
||||
MCP Spec Reference:
|
||||
https://modelcontextprotocol.io/specification/2025-11-25/client/elicitation
|
||||
"""
|
||||
|
||||
from typing import Any, Optional, Union
|
||||
from litellm._logging import verbose_logger
|
||||
|
||||
# Guard imports that require the mcp package
|
||||
try:
|
||||
from mcp.types import (
|
||||
ElicitRequestFormParams,
|
||||
ElicitRequestParams,
|
||||
ElicitRequestURLParams,
|
||||
ElicitResult,
|
||||
ErrorData,
|
||||
)
|
||||
|
||||
MCP_ELICITATION_AVAILABLE = True
|
||||
except ImportError:
|
||||
MCP_ELICITATION_AVAILABLE = False
|
||||
|
||||
|
||||
async def handle_elicitation_request(
|
||||
context: Any,
|
||||
params: "ElicitRequestParams",
|
||||
downstream_session: Optional[Any] = None,
|
||||
downstream_capabilities: Optional[Any] = None,
|
||||
) -> Union["ElicitResult", "ErrorData"]:
|
||||
"""
|
||||
Handle an MCP elicitation/create request from an upstream MCP server.
|
||||
In Gateway mode (Mode A), we relay the elicitation request to the
|
||||
connected downstream client if they declared elicitation capabilities.
|
||||
In Tool Bridge mode (Mode B), there's no persistent downstream MCP
|
||||
client, so we return a decline response.
|
||||
Args:
|
||||
context: MCP RequestContext from the upstream server connection.
|
||||
params: The ElicitRequestParams (either form or URL mode).
|
||||
downstream_session: The ServerSession to the downstream client,
|
||||
if available (for relaying).
|
||||
downstream_capabilities: The downstream client's declared
|
||||
capabilities, used to check elicitation support.
|
||||
Returns:
|
||||
ElicitResult with the user's response, or ErrorData on failure.
|
||||
"""
|
||||
if not MCP_ELICITATION_AVAILABLE:
|
||||
return ErrorData(
|
||||
code=-1,
|
||||
message="MCP elicitation is not available (mcp package not installed)",
|
||||
)
|
||||
try:
|
||||
mode = getattr(params, "mode", "form")
|
||||
verbose_logger.info(
|
||||
"MCP elicitation: received request mode=%s, message=%s",
|
||||
mode,
|
||||
getattr(params, "message", ""),
|
||||
)
|
||||
# Check if we have a downstream session to relay to
|
||||
if downstream_session is not None:
|
||||
return await _relay_elicitation_to_downstream(
|
||||
params=params,
|
||||
downstream_session=downstream_session,
|
||||
downstream_capabilities=downstream_capabilities,
|
||||
)
|
||||
# No downstream session — we're in Tool Bridge mode
|
||||
# or the client doesn't support elicitation
|
||||
verbose_logger.info(
|
||||
"MCP elicitation: no downstream session available, declining"
|
||||
)
|
||||
return ElicitResult(
|
||||
action="decline",
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.exception("MCP elicitation handler failed: %s", e)
|
||||
return ErrorData(
|
||||
code=-1,
|
||||
message=f"Elicitation failed: {str(e)}",
|
||||
)
|
||||
|
||||
|
||||
async def _relay_elicitation_to_downstream(
|
||||
params: "ElicitRequestParams",
|
||||
downstream_session: Any,
|
||||
downstream_capabilities: Optional[Any] = None,
|
||||
) -> Union["ElicitResult", "ErrorData"]:
|
||||
"""
|
||||
Relay an elicitation request to the downstream MCP client.
|
||||
Uses the ServerSession's elicit_form() or elicit_url() methods to
|
||||
send the elicitation request back to the connected client.
|
||||
Args:
|
||||
params: The elicitation request parameters.
|
||||
downstream_session: The ServerSession connected to the downstream client.
|
||||
downstream_capabilities: Client capabilities to check support.
|
||||
Returns:
|
||||
ElicitResult from the downstream client.
|
||||
"""
|
||||
mode = getattr(params, "mode", "form")
|
||||
# Check if the downstream client supports the requested mode
|
||||
if downstream_capabilities is not None:
|
||||
elicit_caps = getattr(downstream_capabilities, "elicitation", None)
|
||||
if elicit_caps is None:
|
||||
verbose_logger.info(
|
||||
"MCP elicitation: downstream client does not support elicitation"
|
||||
)
|
||||
return ElicitResult(action="decline")
|
||||
if mode == "url":
|
||||
url_cap = getattr(elicit_caps, "url", None)
|
||||
if url_cap is None:
|
||||
verbose_logger.info(
|
||||
"MCP elicitation: downstream client does not support URL mode"
|
||||
)
|
||||
return ElicitResult(action="decline")
|
||||
if mode == "form":
|
||||
form_cap = getattr(elicit_caps, "form", None)
|
||||
if form_cap is None:
|
||||
verbose_logger.info(
|
||||
"MCP elicitation: downstream client does not support form mode"
|
||||
)
|
||||
return ElicitResult(action="decline")
|
||||
try:
|
||||
if mode == "url" and isinstance(params, ElicitRequestURLParams):
|
||||
# URL mode: relay URL to client for external navigation
|
||||
verbose_logger.info(
|
||||
"MCP elicitation: relaying URL mode to downstream, url=%s",
|
||||
getattr(params, "url", ""),
|
||||
)
|
||||
result = await downstream_session.elicit_url(
|
||||
message=params.message,
|
||||
url=params.url,
|
||||
elicitation_id=getattr(params, "elicitationId", None),
|
||||
)
|
||||
elif isinstance(params, ElicitRequestFormParams):
|
||||
# Form mode: relay structured form to client
|
||||
verbose_logger.info("MCP elicitation: relaying form mode to downstream")
|
||||
result = await downstream_session.elicit_form(
|
||||
message=params.message,
|
||||
requestedSchema=getattr(params, "requestedSchema", None),
|
||||
)
|
||||
else:
|
||||
# Fallback for generic ElicitRequestParams — pass an empty schema
|
||||
# since elicit() requires requestedSchema as a positional arg.
|
||||
verbose_logger.info(
|
||||
"MCP elicitation: relaying generic elicitation to downstream"
|
||||
)
|
||||
result = await downstream_session.elicit(
|
||||
message=getattr(params, "message", ""),
|
||||
requestedSchema=getattr(params, "requestedSchema", {}),
|
||||
)
|
||||
verbose_logger.info(
|
||||
"MCP elicitation: downstream responded with action=%s",
|
||||
getattr(result, "action", "unknown"),
|
||||
)
|
||||
return result
|
||||
except Exception as e:
|
||||
verbose_logger.warning("MCP elicitation: failed to relay to downstream: %s", e)
|
||||
# If relay fails, decline gracefully
|
||||
return ElicitResult(action="decline")
|
||||
|
|
@ -49,6 +49,12 @@ from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
|
|||
MCPRequestHandler,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.exceptions import MCPUpstreamAuthError
|
||||
from litellm.proxy._experimental.mcp_server.elicitation_handler import (
|
||||
MCP_ELICITATION_AVAILABLE,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.sampling_handler import (
|
||||
MCP_SAMPLING_AVAILABLE,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.oauth2_token_cache import resolve_mcp_auth
|
||||
from litellm.proxy._experimental.mcp_server.utils import (
|
||||
MCP_TOOL_PREFIX_SEPARATOR,
|
||||
|
|
@ -289,6 +295,82 @@ def _deserialize_json_dict(data: Any) -> Optional[Dict[str, str]]:
|
|||
return data
|
||||
|
||||
|
||||
def _create_sampling_callback(user_api_key_auth: Optional[Any] = None):
|
||||
"""
|
||||
Create a sampling callback for MCP ClientSession.
|
||||
Returns a callable that handles sampling/createMessage requests from
|
||||
upstream MCP servers by routing them through litellm.acompletion().
|
||||
"""
|
||||
if not MCP_SAMPLING_AVAILABLE:
|
||||
return None
|
||||
|
||||
async def _sampling_callback(context, params):
|
||||
from litellm.proxy._experimental.mcp_server.sampling_handler import (
|
||||
handle_sampling_create_message,
|
||||
)
|
||||
import litellm
|
||||
from litellm.proxy._experimental.mcp_server.server import (
|
||||
get_active_auth_context,
|
||||
)
|
||||
|
||||
auth_context = get_active_auth_context()
|
||||
resolved_auth = user_api_key_auth or (
|
||||
auth_context.user_api_key_auth if auth_context else None
|
||||
)
|
||||
# Forward original HTTP headers and client IP so that
|
||||
# header-dependent guardrails, tag-based routing, trace
|
||||
# correlation, and forward_llm_provider_auth_headers work
|
||||
# correctly for sampling sub-calls.
|
||||
_raw_headers = getattr(auth_context, "raw_headers", None)
|
||||
_client_ip = getattr(auth_context, "client_ip", None)
|
||||
|
||||
return await handle_sampling_create_message(
|
||||
context=context,
|
||||
params=params,
|
||||
default_model=getattr(litellm, "default_mcp_sampling_model", None),
|
||||
user_api_key_auth=resolved_auth,
|
||||
raw_headers=_raw_headers,
|
||||
client_ip=_client_ip,
|
||||
)
|
||||
|
||||
return _sampling_callback
|
||||
|
||||
|
||||
def _create_elicitation_callback():
|
||||
"""
|
||||
Create an elicitation callback for MCP ClientSession.
|
||||
Returns a callable that handles elicitation/create requests from
|
||||
upstream MCP servers. In gateway mode, this relays to the downstream
|
||||
client; in tool bridge mode, it returns a decline response.
|
||||
"""
|
||||
if not MCP_ELICITATION_AVAILABLE:
|
||||
return None
|
||||
|
||||
async def _elicitation_callback(context, params):
|
||||
from litellm.proxy._experimental.mcp_server.elicitation_handler import (
|
||||
handle_elicitation_request,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.server import get_active_mcp_session
|
||||
|
||||
# In Gateway mode, we relay the elicitation request to the downstream client
|
||||
# that triggered the current operation.
|
||||
downstream_session = get_active_mcp_session()
|
||||
downstream_capabilities = (
|
||||
getattr(downstream_session, "capabilities", None)
|
||||
if downstream_session
|
||||
else None
|
||||
)
|
||||
|
||||
return await handle_elicitation_request(
|
||||
context=context,
|
||||
params=params,
|
||||
downstream_session=downstream_session,
|
||||
downstream_capabilities=downstream_capabilities,
|
||||
)
|
||||
|
||||
return _elicitation_callback
|
||||
|
||||
|
||||
class MCPServerManager:
|
||||
_STDIO_ENV_TEMPLATE_PATTERN = re.compile(r"^\$\{(X-[^}]+)\}$")
|
||||
|
||||
|
|
@ -600,6 +682,8 @@ class MCPServerManager:
|
|||
"subject_token_type",
|
||||
"urn:ietf:params:oauth:token-type:access_token",
|
||||
),
|
||||
allow_sampling=bool(server_config.get("allow_sampling", False)),
|
||||
allow_elicitation=bool(server_config.get("allow_elicitation", False)),
|
||||
)
|
||||
self._assign_unique_short_prefix(new_server)
|
||||
_warn_internal_delegate_pkce_if_applicable(new_server, source="config")
|
||||
|
|
@ -699,8 +783,7 @@ class MCPServerManager:
|
|||
)
|
||||
|
||||
verbose_logger.debug(
|
||||
f"Using headers for OpenAPI tools (excluding sensitive values): "
|
||||
f"{list(headers.keys())}"
|
||||
f"Using headers for OpenAPI tools (excluding sensitive values): {list(headers.keys())}"
|
||||
)
|
||||
|
||||
# Extract and register tools from OpenAPI paths
|
||||
|
|
@ -1494,6 +1577,7 @@ class MCPServerManager:
|
|||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
stdio_env: Optional[Dict[str, str]] = None,
|
||||
subject_token: Optional[str] = None,
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
|
||||
) -> MCPClient:
|
||||
"""
|
||||
Create an MCPClient instance for the given server.
|
||||
|
|
@ -1510,6 +1594,7 @@ class MCPServerManager:
|
|||
extra_headers: Additional headers to forward.
|
||||
stdio_env: Environment variables for stdio transport.
|
||||
subject_token: Optional user JWT for token exchange (OBO) flow.
|
||||
user_api_key_auth: Optional auth context for sampling callbacks.
|
||||
|
||||
Returns:
|
||||
Configured MCP client instance.
|
||||
|
|
@ -1520,23 +1605,44 @@ class MCPServerManager:
|
|||
|
||||
transport = server.transport or MCPTransport.sse
|
||||
|
||||
# Create sampling and elicitation callbacks for this client
|
||||
sampling_cb = (
|
||||
_create_sampling_callback(user_api_key_auth=user_api_key_auth)
|
||||
if server.allow_sampling
|
||||
else None
|
||||
)
|
||||
elicitation_cb = (
|
||||
_create_elicitation_callback() if server.allow_elicitation else None
|
||||
)
|
||||
|
||||
# Handle stdio transport
|
||||
if transport == MCPTransport.stdio:
|
||||
resolved_env = (
|
||||
stdio_env if stdio_env is not None else dict(server.env or {})
|
||||
stdio_env
|
||||
if stdio_env is not None
|
||||
else (dict(server.env) if server.env is not None else None)
|
||||
)
|
||||
|
||||
# Ensure npm-based STDIO MCP servers have a writable cache dir.
|
||||
# In containers the default (~/.npm or /app/.npm) may not exist
|
||||
# or be read-only, causing npx to fail with ENOENT.
|
||||
if "NPM_CONFIG_CACHE" not in resolved_env:
|
||||
if resolved_env is not None and "NPM_CONFIG_CACHE" not in resolved_env:
|
||||
resolved_env["NPM_CONFIG_CACHE"] = MCP_NPM_CACHE_DIR
|
||||
# Defense-in-depth: block commands not in the allowlist.
|
||||
# The Pydantic validator blocks new servers; this catches legacy
|
||||
# config/DB records predating the allowlist.
|
||||
if server.command:
|
||||
base_command = os.path.basename(server.command)
|
||||
if base_command not in MCP_STDIO_ALLOWED_COMMANDS:
|
||||
# Strip .exe/.cmd/.bat/.com suffix for Windows compatibility
|
||||
base_command_no_ext = base_command.lower()
|
||||
for ext in [".exe", ".cmd", ".bat", ".com"]:
|
||||
if base_command.lower().endswith(ext):
|
||||
base_command_no_ext = base_command[: -len(ext)].lower()
|
||||
break
|
||||
if (
|
||||
base_command.lower() not in MCP_STDIO_ALLOWED_COMMANDS
|
||||
and base_command_no_ext not in MCP_STDIO_ALLOWED_COMMANDS
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail=f"MCP stdio command '{server.command}' is not in the allowlist ({sorted(MCP_STDIO_ALLOWED_COMMANDS)}). "
|
||||
|
|
@ -1559,6 +1665,8 @@ class MCPServerManager:
|
|||
timeout=MCP_CLIENT_TIMEOUT,
|
||||
stdio_config=stdio_config,
|
||||
extra_headers=extra_headers,
|
||||
sampling_callback=sampling_cb,
|
||||
elicitation_callback=elicitation_cb,
|
||||
)
|
||||
else:
|
||||
# For HTTP/SSE transports
|
||||
|
|
@ -1585,6 +1693,8 @@ class MCPServerManager:
|
|||
timeout=MCP_CLIENT_TIMEOUT,
|
||||
extra_headers=extra_headers,
|
||||
aws_auth=aws_auth,
|
||||
sampling_callback=sampling_cb,
|
||||
elicitation_callback=elicitation_cb,
|
||||
)
|
||||
|
||||
async def _get_tools_from_server(
|
||||
|
|
@ -1668,6 +1778,7 @@ class MCPServerManager:
|
|||
mcp_auth_header=mcp_auth_header,
|
||||
extra_headers=extra_headers,
|
||||
stdio_env=stdio_env,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
|
||||
## HANDLE OPENAPI TOOLS
|
||||
|
|
@ -3030,6 +3141,7 @@ class MCPServerManager:
|
|||
extra_headers=extra_headers,
|
||||
stdio_env=stdio_env,
|
||||
subject_token=subject_token,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
|
||||
call_tool_params = MCPCallToolRequestParams(
|
||||
|
|
@ -3260,7 +3372,6 @@ class MCPServerManager:
|
|||
)
|
||||
)
|
||||
else:
|
||||
# For regular MCP servers, use the MCP client
|
||||
return await self._call_regular_mcp_tool(
|
||||
mcp_server=mcp_server,
|
||||
original_tool_name=name,
|
||||
|
|
|
|||
1251
litellm/proxy/_experimental/mcp_server/sampling_handler.py
Normal file
1251
litellm/proxy/_experimental/mcp_server/sampling_handler.py
Normal file
File diff suppressed because it is too large
Load diff
|
|
@ -6,6 +6,7 @@ LiteLLM MCP Server Routes
|
|||
|
||||
import asyncio
|
||||
import contextlib
|
||||
import contextvars
|
||||
import hashlib
|
||||
import json
|
||||
import time
|
||||
|
|
@ -125,6 +126,18 @@ try:
|
|||
GetPromptResult,
|
||||
ResourceTemplate,
|
||||
TextResourceContents,
|
||||
Tool,
|
||||
)
|
||||
from mcp.server.session import ServerSession as _McpServerSession
|
||||
import weakref
|
||||
|
||||
# Robust auth lookup keyed by session_object.
|
||||
_session_obj_auth_storage: (
|
||||
"weakref.WeakKeyDictionary[Any, MCPAuthenticatedUser]"
|
||||
) = weakref.WeakKeyDictionary()
|
||||
|
||||
active_mcp_session_var: contextvars.ContextVar[Optional[_McpServerSession]] = (
|
||||
contextvars.ContextVar("active_mcp_session", default=None)
|
||||
)
|
||||
except ImportError as e:
|
||||
verbose_logger.debug(f"MCP module not found: {e}")
|
||||
|
|
@ -483,10 +496,18 @@ if MCP_AVAILABLE:
|
|||
########################################################
|
||||
|
||||
@server.list_tools()
|
||||
async def list_tools() -> List[MCPTool]:
|
||||
async def handle_list_tools() -> List[Tool]:
|
||||
"""
|
||||
List all available tools
|
||||
List all available tools.
|
||||
Also captures the active session for propagation to callbacks.
|
||||
"""
|
||||
from mcp.server.lowlevel.server import request_ctx
|
||||
|
||||
req_ctx = request_ctx.get(None)
|
||||
_session_reset_token = None
|
||||
if req_ctx:
|
||||
_session_reset_token = active_mcp_session_var.set(req_ctx.session)
|
||||
|
||||
try:
|
||||
# Get user authentication from context variable
|
||||
(
|
||||
|
|
@ -497,7 +518,7 @@ if MCP_AVAILABLE:
|
|||
oauth2_headers,
|
||||
raw_headers,
|
||||
_client_ip,
|
||||
) = get_auth_context()
|
||||
) = await get_or_extract_auth_context()
|
||||
verbose_logger.debug(
|
||||
f"MCP list_tools - User API Key Auth from context: {user_api_key_auth}"
|
||||
)
|
||||
|
|
@ -528,152 +549,178 @@ if MCP_AVAILABLE:
|
|||
# Return empty list instead of failing completely
|
||||
# This prevents the HTTP stream from failing and allows the client to get a response
|
||||
return []
|
||||
finally:
|
||||
if _session_reset_token is not None:
|
||||
active_mcp_session_var.reset(_session_reset_token)
|
||||
|
||||
@server.call_tool()
|
||||
async def mcp_server_tool_call(
|
||||
name: str, arguments: Optional[Dict[str, Any]]
|
||||
async def mcp_server_tool_call( # noqa: PLR0915
|
||||
name: str, arguments: Dict[str, Any] | None
|
||||
) -> CallToolResult:
|
||||
"""
|
||||
Call a specific tool with the provided arguments
|
||||
|
||||
Args:
|
||||
name (str): Name of the tool to call
|
||||
arguments (Dict[str, Any] | None): Arguments to pass to the tool
|
||||
|
||||
Returns:
|
||||
List[Union[MCPTextContent, MCPImageContent, MCPEmbeddedResource]]: Tool execution results
|
||||
|
||||
Raises:
|
||||
HTTPException: If tool not found or arguments missing
|
||||
"""
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.exceptions import BlockedPiiEntityError, GuardrailRaisedException
|
||||
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
|
||||
from litellm.proxy.proxy_server import proxy_config
|
||||
from mcp.types import CallToolResult
|
||||
from mcp.server.lowlevel.server import request_ctx
|
||||
|
||||
# Validate arguments
|
||||
(
|
||||
user_api_key_auth,
|
||||
mcp_auth_header,
|
||||
mcp_servers,
|
||||
mcp_server_auth_headers,
|
||||
oauth2_headers,
|
||||
raw_headers,
|
||||
_client_ip,
|
||||
) = get_auth_context()
|
||||
req_ctx = request_ctx.get(None)
|
||||
_session_reset_token = None
|
||||
if req_ctx:
|
||||
_session_reset_token = active_mcp_session_var.set(req_ctx.session)
|
||||
|
||||
verbose_logger.debug(
|
||||
f"MCP mcp_server_tool_call - User API Key Auth from context: {user_api_key_auth}"
|
||||
)
|
||||
host_progress_callback = None
|
||||
try:
|
||||
host_ctx = server.request_context
|
||||
if host_ctx and hasattr(host_ctx, "meta") and host_ctx.meta:
|
||||
host_token = getattr(host_ctx.meta, "progressToken", None)
|
||||
if host_token and hasattr(host_ctx, "session") and host_ctx.session:
|
||||
host_session = host_ctx.session
|
||||
|
||||
async def forward_progress(progress: float, total: Optional[float]):
|
||||
"""Forward progress notifications from external MCP to Host"""
|
||||
try:
|
||||
await host_session.send_progress_notification(
|
||||
progress_token=host_token,
|
||||
progress=progress,
|
||||
total=total,
|
||||
)
|
||||
verbose_logger.debug(
|
||||
f"Forwarded progress {progress}/{total} to Host"
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.error(
|
||||
f"Failed to forward progress to Host: {e}"
|
||||
)
|
||||
|
||||
host_progress_callback = forward_progress
|
||||
verbose_logger.debug(
|
||||
f"Host progressToken captured: {host_token[:8]}..."
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.warning(f"Could not capture host progress context: {e}")
|
||||
try:
|
||||
# Create a body date for logging
|
||||
body_data = {"name": name, "arguments": arguments}
|
||||
# Set trace/session id from raw_headers so spend logs and logging_obj stay consistent (same as A2A)
|
||||
chain_id = get_chain_id_from_headers(raw_headers)
|
||||
if chain_id:
|
||||
body_data["litellm_trace_id"] = chain_id
|
||||
body_data["litellm_session_id"] = chain_id
|
||||
|
||||
request = Request(
|
||||
scope={
|
||||
"type": "http",
|
||||
"method": "POST",
|
||||
"path": "/mcp/tools/call",
|
||||
"headers": [(b"content-type", b"application/json")],
|
||||
}
|
||||
# Validate arguments
|
||||
(
|
||||
user_api_key_auth,
|
||||
mcp_auth_header,
|
||||
mcp_servers,
|
||||
mcp_server_auth_headers,
|
||||
oauth2_headers,
|
||||
raw_headers,
|
||||
_client_ip,
|
||||
) = await get_or_extract_auth_context()
|
||||
verbose_logger.debug(
|
||||
f"MCP mcp_server_tool_call - user_api_key_auth={user_api_key_auth}, user_role={getattr(user_api_key_auth, 'user_role', 'N/A')}"
|
||||
)
|
||||
if user_api_key_auth is not None:
|
||||
data = await add_litellm_data_to_request(
|
||||
data=body_data,
|
||||
request=request,
|
||||
user_api_key_dict=user_api_key_auth,
|
||||
proxy_config=proxy_config,
|
||||
|
||||
verbose_logger.debug(
|
||||
f"MCP mcp_server_tool_call - User API Key Auth from context: {user_api_key_auth}"
|
||||
)
|
||||
host_progress_callback = None
|
||||
try:
|
||||
host_ctx = server.request_context
|
||||
if host_ctx and hasattr(host_ctx, "meta") and host_ctx.meta:
|
||||
host_token = getattr(host_ctx.meta, "progressToken", None)
|
||||
if host_token and hasattr(host_ctx, "session") and host_ctx.session:
|
||||
host_session = host_ctx.session
|
||||
|
||||
async def forward_progress(
|
||||
progress: float, total: Optional[float]
|
||||
):
|
||||
"""Forward progress notifications from external MCP to Host"""
|
||||
try:
|
||||
await host_session.send_progress_notification(
|
||||
progress_token=host_token,
|
||||
progress=progress,
|
||||
total=total,
|
||||
)
|
||||
verbose_logger.debug(
|
||||
f"Forwarded progress {progress}/{total} to Host"
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.error(
|
||||
f"Failed to forward progress to Host: {e}"
|
||||
)
|
||||
|
||||
host_progress_callback = forward_progress
|
||||
verbose_logger.debug(
|
||||
f"Host progressToken captured: {host_token[:8]}..."
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.warning(f"Could not capture host progress context: {e}")
|
||||
try:
|
||||
# Create a body date for logging
|
||||
body_data = {"name": name, "arguments": arguments}
|
||||
# Set trace/session id from raw_headers so spend logs and logging_obj stay consistent (same as A2A)
|
||||
chain_id = get_chain_id_from_headers(raw_headers)
|
||||
if chain_id:
|
||||
body_data["litellm_trace_id"] = chain_id
|
||||
body_data["litellm_session_id"] = chain_id
|
||||
|
||||
request = Request(
|
||||
scope={
|
||||
"type": "http",
|
||||
"method": "POST",
|
||||
"path": "/mcp/tools/call",
|
||||
"headers": [(b"content-type", b"application/json")],
|
||||
}
|
||||
)
|
||||
else:
|
||||
data = body_data
|
||||
|
||||
response = await call_mcp_tool(
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_servers=mcp_servers,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
host_progress_callback=host_progress_callback,
|
||||
**data, # for logging
|
||||
)
|
||||
except BlockedPiiEntityError as e:
|
||||
verbose_logger.error(f"BlockedPiiEntityError in MCP tool call: {str(e)}")
|
||||
return CallToolResult(
|
||||
content=[
|
||||
TextContent(
|
||||
text=f"Error: Blocked PII entity detected - {str(e)}",
|
||||
type="text",
|
||||
if user_api_key_auth is not None:
|
||||
data = await add_litellm_data_to_request(
|
||||
data=body_data,
|
||||
request=request,
|
||||
user_api_key_dict=user_api_key_auth,
|
||||
proxy_config=proxy_config,
|
||||
)
|
||||
],
|
||||
isError=True,
|
||||
)
|
||||
except GuardrailRaisedException as e:
|
||||
verbose_logger.error(f"GuardrailRaisedException in MCP tool call: {str(e)}")
|
||||
return CallToolResult(
|
||||
content=[
|
||||
TextContent(
|
||||
text=f"Error: Guardrail violation - {str(e)}", type="text"
|
||||
)
|
||||
],
|
||||
isError=True,
|
||||
)
|
||||
except HTTPException as e:
|
||||
verbose_logger.error(f"HTTPException in MCP tool call: {str(e)}")
|
||||
return CallToolResult(
|
||||
content=[TextContent(text=f"Error: {str(e.detail)}", type="text")],
|
||||
isError=True,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.exception(f"MCP mcp_server_tool_call - error: {e}")
|
||||
return CallToolResult(
|
||||
content=[TextContent(text=f"Error: {str(e)}", type="text")],
|
||||
isError=True,
|
||||
)
|
||||
else:
|
||||
data = body_data
|
||||
|
||||
return response
|
||||
response = await call_mcp_tool(
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_servers=mcp_servers,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
host_progress_callback=host_progress_callback,
|
||||
**data, # for logging
|
||||
)
|
||||
except BlockedPiiEntityError as e:
|
||||
verbose_logger.error(
|
||||
f"BlockedPiiEntityError in MCP tool call: {str(e)}"
|
||||
)
|
||||
return CallToolResult(
|
||||
content=[
|
||||
TextContent(
|
||||
text=f"Error: Blocked PII entity detected - {str(e)}",
|
||||
type="text",
|
||||
)
|
||||
],
|
||||
isError=True,
|
||||
)
|
||||
except GuardrailRaisedException as e:
|
||||
verbose_logger.error(
|
||||
f"GuardrailRaisedException in MCP tool call: {str(e)}"
|
||||
)
|
||||
return CallToolResult(
|
||||
content=[
|
||||
TextContent(
|
||||
text=f"Error: Guardrail violation - {str(e)}", type="text"
|
||||
)
|
||||
],
|
||||
isError=True,
|
||||
)
|
||||
except HTTPException as e:
|
||||
verbose_logger.error(f"HTTPException in MCP tool call: {str(e)}")
|
||||
return CallToolResult(
|
||||
content=[TextContent(text=f"Error: {str(e.detail)}", type="text")],
|
||||
isError=True,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.exception(f"MCP mcp_server_tool_call - error: {e}")
|
||||
return CallToolResult(
|
||||
content=[TextContent(text=f"Error: {str(e)}", type="text")],
|
||||
isError=True,
|
||||
)
|
||||
|
||||
return response
|
||||
finally:
|
||||
if _session_reset_token is not None:
|
||||
active_mcp_session_var.reset(_session_reset_token)
|
||||
|
||||
@server.list_prompts()
|
||||
async def list_prompts() -> List[Prompt]:
|
||||
"""
|
||||
List all available prompts
|
||||
"""
|
||||
from mcp.server.lowlevel.server import request_ctx
|
||||
|
||||
req_ctx = request_ctx.get(None)
|
||||
_session_reset_token = None
|
||||
if req_ctx:
|
||||
_session_reset_token = active_mcp_session_var.set(req_ctx.session)
|
||||
|
||||
try:
|
||||
# Get user authentication from context variable
|
||||
(
|
||||
|
|
@ -684,7 +731,7 @@ if MCP_AVAILABLE:
|
|||
oauth2_headers,
|
||||
raw_headers,
|
||||
_client_ip,
|
||||
) = get_auth_context()
|
||||
) = await get_or_extract_auth_context()
|
||||
verbose_logger.debug(
|
||||
f"MCP list_prompts - User API Key Auth from context: {user_api_key_auth}"
|
||||
)
|
||||
|
|
@ -713,6 +760,9 @@ if MCP_AVAILABLE:
|
|||
# Return empty list instead of failing completely
|
||||
# This prevents the HTTP stream from failing and allows the client to get a response
|
||||
return []
|
||||
finally:
|
||||
if _session_reset_token is not None:
|
||||
active_mcp_session_var.reset(_session_reset_token)
|
||||
|
||||
@server.get_prompt()
|
||||
async def get_prompt(
|
||||
|
|
@ -730,33 +780,13 @@ if MCP_AVAILABLE:
|
|||
"""
|
||||
|
||||
# Validate arguments
|
||||
(
|
||||
user_api_key_auth,
|
||||
mcp_auth_header,
|
||||
mcp_servers,
|
||||
mcp_server_auth_headers,
|
||||
oauth2_headers,
|
||||
raw_headers,
|
||||
_client_ip,
|
||||
) = get_auth_context()
|
||||
from mcp.server.lowlevel.server import request_ctx
|
||||
|
||||
verbose_logger.debug(
|
||||
f"MCP mcp_server_tool_call - User API Key Auth from context: {user_api_key_auth}"
|
||||
)
|
||||
return await mcp_get_prompt(
|
||||
name=name,
|
||||
arguments=arguments,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_servers=mcp_servers,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
)
|
||||
req_ctx = request_ctx.get(None)
|
||||
_session_reset_token = None
|
||||
if req_ctx:
|
||||
_session_reset_token = active_mcp_session_var.set(req_ctx.session)
|
||||
|
||||
@server.list_resources()
|
||||
async def list_resources() -> List[Resource]:
|
||||
"""List all available resources."""
|
||||
try:
|
||||
(
|
||||
user_api_key_auth,
|
||||
|
|
@ -766,7 +796,45 @@ if MCP_AVAILABLE:
|
|||
oauth2_headers,
|
||||
raw_headers,
|
||||
_client_ip,
|
||||
) = get_auth_context()
|
||||
) = await get_or_extract_auth_context()
|
||||
|
||||
verbose_logger.debug(
|
||||
f"MCP mcp_server_tool_call - User API Key Auth from context: {user_api_key_auth}"
|
||||
)
|
||||
return await mcp_get_prompt(
|
||||
name=name,
|
||||
arguments=arguments,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_servers=mcp_servers,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
)
|
||||
finally:
|
||||
if _session_reset_token is not None:
|
||||
active_mcp_session_var.reset(_session_reset_token)
|
||||
|
||||
@server.list_resources()
|
||||
async def list_resources() -> List[Resource]:
|
||||
"""List all available resources."""
|
||||
from mcp.server.lowlevel.server import request_ctx
|
||||
|
||||
req_ctx = request_ctx.get(None)
|
||||
_session_reset_token = None
|
||||
if req_ctx:
|
||||
_session_reset_token = active_mcp_session_var.set(req_ctx.session)
|
||||
|
||||
try:
|
||||
(
|
||||
user_api_key_auth,
|
||||
mcp_auth_header,
|
||||
mcp_servers,
|
||||
mcp_server_auth_headers,
|
||||
oauth2_headers,
|
||||
raw_headers,
|
||||
_client_ip,
|
||||
) = await get_or_extract_auth_context()
|
||||
verbose_logger.debug(
|
||||
f"MCP list_resources - User API Key Auth from context: {user_api_key_auth}"
|
||||
)
|
||||
|
|
@ -792,10 +860,20 @@ if MCP_AVAILABLE:
|
|||
except Exception as e:
|
||||
verbose_logger.exception(f"Error in list_resources endpoint: {str(e)}")
|
||||
return []
|
||||
finally:
|
||||
if _session_reset_token is not None:
|
||||
active_mcp_session_var.reset(_session_reset_token)
|
||||
|
||||
@server.list_resource_templates()
|
||||
async def list_resource_templates() -> List[ResourceTemplate]:
|
||||
"""List all available resource templates."""
|
||||
from mcp.server.lowlevel.server import request_ctx
|
||||
|
||||
req_ctx = request_ctx.get(None)
|
||||
_session_reset_token = None
|
||||
if req_ctx:
|
||||
_session_reset_token = active_mcp_session_var.set(req_ctx.session)
|
||||
|
||||
try:
|
||||
(
|
||||
user_api_key_auth,
|
||||
|
|
@ -805,7 +883,7 @@ if MCP_AVAILABLE:
|
|||
oauth2_headers,
|
||||
raw_headers,
|
||||
_client_ip,
|
||||
) = get_auth_context()
|
||||
) = await get_or_extract_auth_context()
|
||||
verbose_logger.debug(
|
||||
f"MCP list_resource_templates - User API Key Auth from context: {user_api_key_auth}"
|
||||
)
|
||||
|
|
@ -825,8 +903,7 @@ if MCP_AVAILABLE:
|
|||
raw_headers=raw_headers,
|
||||
)
|
||||
verbose_logger.info(
|
||||
"MCP list_resource_templates - Successfully returned "
|
||||
f"{len(resource_templates)} resource templates"
|
||||
f"MCP list_resource_templates - Successfully returned {len(resource_templates)} resource templates"
|
||||
)
|
||||
return resource_templates
|
||||
except Exception as e:
|
||||
|
|
@ -834,30 +911,44 @@ if MCP_AVAILABLE:
|
|||
f"Error in list_resource_templates endpoint: {str(e)}"
|
||||
)
|
||||
return []
|
||||
finally:
|
||||
if _session_reset_token is not None:
|
||||
active_mcp_session_var.reset(_session_reset_token)
|
||||
|
||||
@server.read_resource()
|
||||
async def read_resource(url: AnyUrl) -> list[ReadResourceContents]:
|
||||
(
|
||||
user_api_key_auth,
|
||||
mcp_auth_header,
|
||||
mcp_servers,
|
||||
mcp_server_auth_headers,
|
||||
oauth2_headers,
|
||||
raw_headers,
|
||||
_client_ip,
|
||||
) = get_auth_context()
|
||||
from mcp.server.lowlevel.server import request_ctx
|
||||
|
||||
read_resource_result = await mcp_read_resource(
|
||||
url=url,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_servers=mcp_servers,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
)
|
||||
req_ctx = request_ctx.get(None)
|
||||
_session_reset_token = None
|
||||
if req_ctx:
|
||||
_session_reset_token = active_mcp_session_var.set(req_ctx.session)
|
||||
|
||||
return _normalize_resource_contents(read_resource_result.contents)
|
||||
try:
|
||||
(
|
||||
user_api_key_auth,
|
||||
mcp_auth_header,
|
||||
mcp_servers,
|
||||
mcp_server_auth_headers,
|
||||
oauth2_headers,
|
||||
raw_headers,
|
||||
_client_ip,
|
||||
) = await get_or_extract_auth_context()
|
||||
|
||||
read_resource_result = await mcp_read_resource(
|
||||
url=url,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_servers=mcp_servers,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
)
|
||||
|
||||
return _normalize_resource_contents(read_resource_result.contents)
|
||||
finally:
|
||||
if _session_reset_token is not None:
|
||||
active_mcp_session_var.reset(_session_reset_token)
|
||||
|
||||
########################################################
|
||||
############ End of MCP Server Routes ##################
|
||||
|
|
@ -1180,8 +1271,7 @@ if MCP_AVAILABLE:
|
|||
cached_token = await mcp_per_user_token_cache.get(user_id, server_id)
|
||||
if cached_token is not None:
|
||||
verbose_logger.debug(
|
||||
"_get_user_oauth_extra_headers_from_db: Redis hit for "
|
||||
"user=%s server=%s",
|
||||
"_get_user_oauth_extra_headers_from_db: Redis hit for user=%s server=%s",
|
||||
user_id,
|
||||
server_id,
|
||||
)
|
||||
|
|
@ -1207,8 +1297,7 @@ if MCP_AVAILABLE:
|
|||
|
||||
if is_oauth_credential_expired(cred):
|
||||
verbose_logger.debug(
|
||||
"_get_user_oauth_extra_headers_from_db: token expired for "
|
||||
"user=%s server=%s — attempting refresh",
|
||||
"_get_user_oauth_extra_headers_from_db: token expired for user=%s server=%s — attempting refresh",
|
||||
user_id,
|
||||
server_id,
|
||||
)
|
||||
|
|
@ -1230,8 +1319,7 @@ if MCP_AVAILABLE:
|
|||
)
|
||||
except Exception as refresh_exc:
|
||||
verbose_logger.warning(
|
||||
"_get_user_oauth_extra_headers_from_db: refresh failed "
|
||||
"for user=%s server=%s: %s",
|
||||
"_get_user_oauth_extra_headers_from_db: refresh failed for user=%s server=%s: %s",
|
||||
user_id,
|
||||
server_id,
|
||||
refresh_exc,
|
||||
|
|
@ -1275,8 +1363,7 @@ if MCP_AVAILABLE:
|
|||
return {"Authorization": f"Bearer {access_token}"}
|
||||
except Exception as e:
|
||||
verbose_logger.warning(
|
||||
"_get_user_oauth_extra_headers_from_db: failed to retrieve credential for "
|
||||
"user=%s server=%s: %s",
|
||||
"_get_user_oauth_extra_headers_from_db: failed to retrieve credential for user=%s server=%s: %s",
|
||||
user_id,
|
||||
server_id,
|
||||
e,
|
||||
|
|
@ -2485,7 +2572,7 @@ if MCP_AVAILABLE:
|
|||
arguments=arguments or {},
|
||||
server_name=server_name or mcp_server.name,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
proxy_logging_obj=proxy_logging_obj, # type: ignore[arg-type]
|
||||
server=mcp_server,
|
||||
raw_headers=raw_headers,
|
||||
)
|
||||
|
|
@ -2744,8 +2831,7 @@ if MCP_AVAILABLE:
|
|||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=(
|
||||
"Multiple MCP servers configured; read_resource currently "
|
||||
"supports exactly one allowed server."
|
||||
"Multiple MCP servers configured; read_resource currently supports exactly one allowed server."
|
||||
),
|
||||
)
|
||||
|
||||
|
|
@ -3124,8 +3210,7 @@ if MCP_AVAILABLE:
|
|||
return False
|
||||
except Exception:
|
||||
verbose_logger.debug(
|
||||
"Unable to inspect active MCP sessions for '%s'. "
|
||||
"Deferring to session manager.",
|
||||
"Unable to inspect active MCP sessions for '%s'. Deferring to session manager.",
|
||||
_session_id,
|
||||
)
|
||||
return False
|
||||
|
|
@ -3136,8 +3221,7 @@ if MCP_AVAILABLE:
|
|||
if method == "DELETE":
|
||||
_remove_stateful_session_tracking(_session_id)
|
||||
verbose_logger.info(
|
||||
"DELETE request for non-existent MCP session '%s'. "
|
||||
"Returning success (idempotent DELETE).",
|
||||
"DELETE request for non-existent MCP session '%s'. Returning success (idempotent DELETE).",
|
||||
_session_id,
|
||||
)
|
||||
success_response = JSONResponse(
|
||||
|
|
@ -3615,6 +3699,7 @@ if MCP_AVAILABLE:
|
|||
return
|
||||
session_id = _get_session_id_from_scope(scope)
|
||||
|
||||
body = b""
|
||||
if scope.get("method") == "POST":
|
||||
consumed_messages, body = await _read_request_body_for_routing(receive)
|
||||
is_initialize = _is_initialize_request(body)
|
||||
|
|
@ -3639,8 +3724,7 @@ if MCP_AVAILABLE:
|
|||
)
|
||||
if not await _enforce_stateful_session_cap_for_owner(request_owner):
|
||||
verbose_logger.warning(
|
||||
"Rejecting MCP initialize: caller already holds the maximum "
|
||||
"number of active stateful sessions."
|
||||
"Rejecting MCP initialize: caller already holds the maximum number of active stateful sessions."
|
||||
)
|
||||
too_many_response = JSONResponse(
|
||||
status_code=429,
|
||||
|
|
@ -3672,9 +3756,54 @@ if MCP_AVAILABLE:
|
|||
# POST/DELETE are the methods that actually mutate the shared
|
||||
# auth context, so serializing those is sufficient for the
|
||||
# clobbering race between concurrent JSON-RPC calls.
|
||||
session_lock: Optional[asyncio.Lock] = None
|
||||
#
|
||||
# Also skip the lock for JSON-RPC *responses* (POSTs that carry
|
||||
# a ``result`` or ``error`` but no ``method``). These are replies
|
||||
# to server-initiated requests such as ``elicitation/create`` or
|
||||
# ``sampling/createMessage``. The in-flight tool-call POST that
|
||||
# triggered the server request already holds the session lock, so
|
||||
# trying to acquire it again for the response POST would deadlock.
|
||||
is_jsonrpc_response = False
|
||||
request_method = (scope.get("method") or "").upper()
|
||||
if use_stateful and session_id and request_method in ("POST", "DELETE"):
|
||||
if body and request_method == "POST":
|
||||
try:
|
||||
_peeked = json.loads(body)
|
||||
if (
|
||||
isinstance(_peeked, dict)
|
||||
and _peeked.get("jsonrpc") == "2.0"
|
||||
and "id" in _peeked
|
||||
and "method" not in _peeked
|
||||
and ("result" in _peeked or "error" in _peeked)
|
||||
):
|
||||
is_jsonrpc_response = True
|
||||
verbose_logger.debug(
|
||||
"MCP: detected JSON-RPC response POST (id=%s), skipping session lock to avoid deadlock",
|
||||
_peeked.get("id"),
|
||||
)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
# Body may be truncated by the peek cap. Use a substring
|
||||
# heuristic so we don't deadlock on large response payloads.
|
||||
# A false-positive (skipping the lock) is harmless; a
|
||||
# false-negative (acquiring it) would deadlock.
|
||||
_body_str = body[:512].decode("utf-8", errors="replace")
|
||||
if (
|
||||
'"jsonrpc"' in _body_str
|
||||
and '"method"' not in _body_str
|
||||
and ('"result"' in _body_str or '"error"' in _body_str)
|
||||
):
|
||||
is_jsonrpc_response = True
|
||||
verbose_logger.debug(
|
||||
"MCP: heuristic detected truncated JSON-RPC response POST, "
|
||||
"skipping session lock to avoid deadlock"
|
||||
)
|
||||
|
||||
session_lock: Optional[asyncio.Lock] = None
|
||||
if (
|
||||
use_stateful
|
||||
and session_id
|
||||
and request_method in ("POST", "DELETE")
|
||||
and not is_jsonrpc_response
|
||||
):
|
||||
session_lock = _stateful_session_locks.setdefault(
|
||||
session_id, asyncio.Lock()
|
||||
)
|
||||
|
|
@ -4099,6 +4228,119 @@ if MCP_AVAILABLE:
|
|||
)
|
||||
return None, None, None, None, None, None, None
|
||||
|
||||
def _get_current_session():
|
||||
try:
|
||||
from mcp.server.lowlevel.server import request_ctx
|
||||
|
||||
return request_ctx.get().session
|
||||
except (LookupError, ImportError):
|
||||
return None
|
||||
|
||||
def _cache_auth_context_lazily():
|
||||
session = _get_current_session()
|
||||
if session is None:
|
||||
return
|
||||
try:
|
||||
if session in _session_obj_auth_storage:
|
||||
return
|
||||
except TypeError:
|
||||
verbose_logger.debug(
|
||||
"_cache_auth_context_lazily: session object is unhashable (type=%s), cannot cache auth context",
|
||||
type(session).__name__,
|
||||
)
|
||||
return
|
||||
|
||||
auth = auth_context_var.get()
|
||||
if auth and isinstance(auth, MCPAuthenticatedUser):
|
||||
try:
|
||||
_session_obj_auth_storage[session] = auth
|
||||
except TypeError:
|
||||
verbose_logger.debug(
|
||||
"_cache_auth_context_lazily: could not store auth via "
|
||||
"session identity — session object is unhashable"
|
||||
)
|
||||
|
||||
def _recover_auth_from_session() -> Optional[MCPAuthenticatedUser]:
|
||||
session = _get_current_session()
|
||||
if session is None:
|
||||
return None
|
||||
|
||||
stored: Optional[MCPAuthenticatedUser] = None
|
||||
try:
|
||||
stored = _session_obj_auth_storage.get(session)
|
||||
except TypeError:
|
||||
verbose_logger.debug(
|
||||
"_recover_auth_from_session: session object is unhashable "
|
||||
"(type=%s), skipping _session_obj_auth_storage lookup",
|
||||
type(session).__name__,
|
||||
)
|
||||
|
||||
return stored
|
||||
|
||||
async def get_or_extract_auth_context() -> Tuple[
|
||||
Optional[UserAPIKeyAuth],
|
||||
Optional[str],
|
||||
Optional[List[str]],
|
||||
Optional[Dict[str, Dict[str, str]]],
|
||||
Optional[Dict[str, str]],
|
||||
Optional[Dict[str, str]],
|
||||
Optional[str],
|
||||
]:
|
||||
"""
|
||||
Get auth context from ContextVar first, then fall back to session
|
||||
storage (which survives cross-task boundaries in the MCP SDK).
|
||||
"""
|
||||
(
|
||||
user_api_key_auth,
|
||||
mcp_auth_header,
|
||||
mcp_servers,
|
||||
mcp_server_auth_headers,
|
||||
oauth2_headers,
|
||||
raw_headers,
|
||||
_client_ip,
|
||||
) = get_auth_context()
|
||||
|
||||
if user_api_key_auth is not None:
|
||||
_cache_auth_context_lazily()
|
||||
else:
|
||||
stored = _recover_auth_from_session()
|
||||
|
||||
if stored:
|
||||
user_api_key_auth = stored.user_api_key_auth
|
||||
mcp_auth_header = stored.mcp_auth_header
|
||||
mcp_servers = stored.mcp_servers
|
||||
mcp_server_auth_headers = stored.mcp_server_auth_headers
|
||||
oauth2_headers = stored.oauth2_headers
|
||||
raw_headers = stored.raw_headers
|
||||
_client_ip = stored.client_ip
|
||||
return (
|
||||
user_api_key_auth,
|
||||
mcp_auth_header,
|
||||
mcp_servers,
|
||||
mcp_server_auth_headers,
|
||||
oauth2_headers,
|
||||
raw_headers,
|
||||
_client_ip,
|
||||
)
|
||||
|
||||
def get_active_mcp_session() -> Optional[_McpServerSession]:
|
||||
"""Return the active MCP session captured during handler execution."""
|
||||
session = active_mcp_session_var.get()
|
||||
if session is not None:
|
||||
return session
|
||||
return _get_current_session()
|
||||
|
||||
def get_active_auth_context() -> Optional[MCPAuthenticatedUser]:
|
||||
"""Return auth context from ContextVar or session storage."""
|
||||
auth = auth_context_var.get()
|
||||
if auth and isinstance(auth, MCPAuthenticatedUser):
|
||||
return auth
|
||||
|
||||
stored = _recover_auth_from_session()
|
||||
if stored is not None:
|
||||
return stored
|
||||
return None
|
||||
|
||||
########################################################
|
||||
############ End of Auth Context Functions #############
|
||||
########################################################
|
||||
|
|
|
|||
|
|
@ -115,6 +115,8 @@ class MCPServer(BaseModel):
|
|||
# different ``server_id`` values are bumped deterministically. Left
|
||||
# ``None`` in default-prefix mode.
|
||||
short_prefix: Optional[str] = None
|
||||
allow_sampling: bool = False
|
||||
allow_elicitation: bool = False
|
||||
model_config = ConfigDict(arbitrary_types_allowed=True)
|
||||
|
||||
@property
|
||||
|
|
|
|||
|
|
@ -16,7 +16,6 @@ from typing import Any, Dict, Optional
|
|||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
|
@ -549,11 +548,7 @@ class TestHookHeaderMergePriority:
|
|||
captured_extra_headers: Dict[str, Any] = {}
|
||||
|
||||
async def fake_create_mcp_client(
|
||||
server,
|
||||
mcp_auth_header=None,
|
||||
extra_headers=None,
|
||||
stdio_env=None,
|
||||
subject_token=None,
|
||||
server, mcp_auth_header=None, extra_headers=None, stdio_env=None, **kwargs
|
||||
):
|
||||
captured_extra_headers["value"] = extra_headers
|
||||
mock_client = MagicMock()
|
||||
|
|
@ -593,11 +588,7 @@ class TestHookHeaderMergePriority:
|
|||
captured_extra_headers: Dict[str, Any] = {}
|
||||
|
||||
async def fake_create_mcp_client(
|
||||
server,
|
||||
mcp_auth_header=None,
|
||||
extra_headers=None,
|
||||
stdio_env=None,
|
||||
subject_token=None,
|
||||
server, mcp_auth_header=None, extra_headers=None, stdio_env=None, **kwargs
|
||||
):
|
||||
captured_extra_headers["value"] = extra_headers
|
||||
mock_client = MagicMock()
|
||||
|
|
@ -643,11 +634,7 @@ class TestHookHeaderMergePriority:
|
|||
captured_extra_headers: Dict[str, Any] = {}
|
||||
|
||||
async def fake_create_mcp_client(
|
||||
server,
|
||||
mcp_auth_header=None,
|
||||
extra_headers=None,
|
||||
stdio_env=None,
|
||||
subject_token=None,
|
||||
server, mcp_auth_header=None, extra_headers=None, stdio_env=None, **kwargs
|
||||
):
|
||||
captured_extra_headers["value"] = extra_headers
|
||||
mock_client = MagicMock()
|
||||
|
|
@ -703,11 +690,7 @@ class TestHookHeaderMergePriority:
|
|||
captured_extra_headers: Dict[str, Any] = {}
|
||||
|
||||
async def fake_create_mcp_client(
|
||||
server,
|
||||
mcp_auth_header=None,
|
||||
extra_headers=None,
|
||||
stdio_env=None,
|
||||
subject_token=None,
|
||||
server, mcp_auth_header=None, extra_headers=None, stdio_env=None, **kwargs
|
||||
):
|
||||
captured_extra_headers["value"] = extra_headers
|
||||
mock_client = MagicMock()
|
||||
|
|
@ -755,11 +738,7 @@ class TestHookHeaderMergePriority:
|
|||
captured_extra_headers: Dict[str, Any] = {}
|
||||
|
||||
async def fake_create_mcp_client(
|
||||
server,
|
||||
mcp_auth_header=None,
|
||||
extra_headers=None,
|
||||
stdio_env=None,
|
||||
subject_token=None,
|
||||
server, mcp_auth_header=None, extra_headers=None, stdio_env=None, **kwargs
|
||||
):
|
||||
captured_extra_headers["value"] = extra_headers
|
||||
mock_client = MagicMock()
|
||||
|
|
|
|||
|
|
@ -0,0 +1,327 @@
|
|||
"""
|
||||
Tests for MCP sampling handler model-access enforcement.
|
||||
|
||||
Verifies that handle_sampling_create_message and _check_model_access
|
||||
enforce the same model-permission checks as regular /chat/completions
|
||||
calls, preventing a malicious upstream MCP server from requesting
|
||||
inference on models the caller's API key is not authorized to use.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.sampling_handler import (
|
||||
_check_model_access,
|
||||
)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _make_user_api_key_auth(
|
||||
*,
|
||||
models=None,
|
||||
team_id=None,
|
||||
team_model_aliases=None,
|
||||
api_key="sk-test-key",
|
||||
token=None,
|
||||
user_role=None,
|
||||
):
|
||||
"""Build a minimal UserAPIKeyAuth-like object for tests."""
|
||||
auth = MagicMock()
|
||||
auth.models = models or []
|
||||
auth.team_id = team_id
|
||||
auth.team_model_aliases = team_model_aliases or {}
|
||||
auth.access_group_ids = []
|
||||
auth.api_key = api_key
|
||||
auth.token = token
|
||||
auth.user_role = user_role
|
||||
return auth
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _check_model_access
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestCheckModelAccess:
|
||||
"""Tests for the _check_model_access helper that gates sampling requests."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_should_return_none_when_no_auth_context(self):
|
||||
"""No auth context means no restriction — pass through."""
|
||||
result = await _check_model_access("gpt-4o", user_api_key_auth=None)
|
||||
assert result is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_should_allow_model_when_key_has_access(self):
|
||||
"""Key with explicit model access should be allowed."""
|
||||
auth = _make_user_api_key_auth(models=["gpt-4o", "gpt-3.5-turbo"])
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.auth.auth_checks.can_key_call_model",
|
||||
new_callable=AsyncMock,
|
||||
return_value=True,
|
||||
) as mock_check:
|
||||
result = await _check_model_access("gpt-4o", user_api_key_auth=auth)
|
||||
|
||||
assert result is None
|
||||
mock_check.assert_awaited_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_should_deny_model_when_key_lacks_access(self):
|
||||
"""Key without model access should be denied with ErrorData."""
|
||||
from litellm.proxy._types import ProxyException
|
||||
|
||||
auth = _make_user_api_key_auth(models=["gpt-3.5-turbo"])
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.auth.auth_checks.can_key_call_model",
|
||||
new_callable=AsyncMock,
|
||||
side_effect=ProxyException(
|
||||
message="key not allowed to access model",
|
||||
type="key_model_access_denied",
|
||||
param="model",
|
||||
code=401,
|
||||
),
|
||||
):
|
||||
result = await _check_model_access("gpt-4o", user_api_key_auth=auth)
|
||||
|
||||
# Should return ErrorData, not raise
|
||||
assert result is not None
|
||||
assert result.code == -1
|
||||
assert "Model access denied" in result.message
|
||||
assert "gpt-4o" in result.message
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_should_allow_wildcard_model_access(self):
|
||||
"""Key with wildcard model access should allow any model."""
|
||||
auth = _make_user_api_key_auth(models=["*"])
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.auth.auth_checks.can_key_call_model",
|
||||
new_callable=AsyncMock,
|
||||
return_value=True,
|
||||
):
|
||||
result = await _check_model_access(
|
||||
"claude-3-opus-20240229", user_api_key_auth=auth
|
||||
)
|
||||
|
||||
assert result is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_should_deny_expensive_model_requested_by_malicious_server(self):
|
||||
"""Simulates the attack: malicious MCP server hints at an expensive model
|
||||
the caller's key is restricted from using."""
|
||||
from litellm.proxy._types import ProxyException
|
||||
|
||||
# Key only has access to cheap models
|
||||
auth = _make_user_api_key_auth(models=["gpt-3.5-turbo"])
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.auth.auth_checks.can_key_call_model",
|
||||
new_callable=AsyncMock,
|
||||
side_effect=ProxyException(
|
||||
message="key not allowed to access model. This key can only access models=['gpt-3.5-turbo']. Tried to access claude-3-opus-20240229",
|
||||
type="key_model_access_denied",
|
||||
param="model",
|
||||
code=401,
|
||||
),
|
||||
):
|
||||
result = await _check_model_access(
|
||||
"claude-3-opus-20240229", user_api_key_auth=auth
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert result.code == -1
|
||||
assert "claude-3-opus-20240229" in result.message
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_should_deny_empty_oauth_passthrough_placeholder(self):
|
||||
"""Regression: process_mcp_request() returns an empty UserAPIKeyAuth()
|
||||
for OAuth2 upstream-token passthrough. The None check alone is not
|
||||
sufficient — the empty placeholder is truthy but has no api_key, no
|
||||
token, and an empty models list. can_key_call_model() would treat
|
||||
that as all-model access, letting an OAuth-only user trigger sampling
|
||||
calls on any proxy model without a LiteLLM key or budget."""
|
||||
# Simulate the empty placeholder from process_mcp_request()
|
||||
auth = _make_user_api_key_auth(
|
||||
models=[],
|
||||
api_key=None,
|
||||
token=None,
|
||||
user_role=None,
|
||||
)
|
||||
|
||||
result = await _check_model_access("gpt-4o", user_api_key_auth=auth)
|
||||
|
||||
# Must be denied — not passed through to can_key_call_model
|
||||
assert result is not None
|
||||
assert result.code == -1
|
||||
assert "sampling requires a valid LiteLLM" in result.message
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_should_allow_proxy_admin_even_without_api_key(self):
|
||||
"""Proxy admins may not have a traditional api_key but should still
|
||||
be allowed to use sampling."""
|
||||
auth = _make_user_api_key_auth(
|
||||
models=[],
|
||||
api_key=None,
|
||||
token=None,
|
||||
user_role="proxy_admin",
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.auth.auth_checks.can_key_call_model",
|
||||
new_callable=AsyncMock,
|
||||
return_value=True,
|
||||
):
|
||||
result = await _check_model_access("gpt-4o", user_api_key_auth=auth)
|
||||
|
||||
assert result is None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# handle_sampling_create_message — auth + budget gating
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestSamplingAuthAndBudgetGating:
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_should_deny_when_no_auth_context(self):
|
||||
"""Sampling must reject calls with no user_api_key_auth."""
|
||||
from litellm.proxy._experimental.mcp_server.sampling_handler import (
|
||||
handle_sampling_create_message,
|
||||
)
|
||||
|
||||
params = MagicMock()
|
||||
params.modelPreferences = None
|
||||
params.messages = []
|
||||
params.systemPrompt = None
|
||||
params.maxTokens = 100
|
||||
params.temperature = None
|
||||
params.stopSequences = None
|
||||
params.tools = None
|
||||
params.toolChoice = None
|
||||
params.metadata = None
|
||||
|
||||
result = await handle_sampling_create_message(
|
||||
context=MagicMock(),
|
||||
params=params,
|
||||
default_model="gpt-4o",
|
||||
user_api_key_auth=None,
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert result.code == -1
|
||||
assert "authenticated" in result.message.lower()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_should_run_budget_checks(self):
|
||||
"""Sampling must call _run_budget_checks after model access check."""
|
||||
from litellm.proxy._experimental.mcp_server.sampling_handler import (
|
||||
handle_sampling_create_message,
|
||||
)
|
||||
|
||||
auth = _make_user_api_key_auth(models=["gpt-4o"])
|
||||
params = MagicMock()
|
||||
params.modelPreferences = None
|
||||
params.messages = []
|
||||
params.systemPrompt = None
|
||||
params.maxTokens = 100
|
||||
params.temperature = None
|
||||
params.stopSequences = None
|
||||
params.tools = None
|
||||
params.toolChoice = None
|
||||
params.metadata = None
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.sampling_handler._check_model_access",
|
||||
new_callable=AsyncMock,
|
||||
return_value=None,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.sampling_handler._run_budget_checks",
|
||||
new_callable=AsyncMock,
|
||||
return_value=None,
|
||||
) as mock_budget,
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.sampling_handler._resolve_model_from_preferences",
|
||||
return_value="gpt-4o",
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.proxy_server.llm_router",
|
||||
new=None,
|
||||
),
|
||||
patch(
|
||||
"litellm.acompletion",
|
||||
new_callable=AsyncMock,
|
||||
return_value=MagicMock(
|
||||
choices=[
|
||||
MagicMock(
|
||||
message=MagicMock(content="hi", tool_calls=None),
|
||||
finish_reason="stop",
|
||||
)
|
||||
],
|
||||
model="gpt-4o",
|
||||
),
|
||||
),
|
||||
):
|
||||
await handle_sampling_create_message(
|
||||
context=MagicMock(),
|
||||
params=params,
|
||||
default_model="gpt-4o",
|
||||
user_api_key_auth=auth,
|
||||
)
|
||||
|
||||
mock_budget.assert_awaited_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_should_deny_over_budget_caller(self):
|
||||
"""When _run_budget_checks returns ErrorData, sampling must return it."""
|
||||
from mcp.types import ErrorData
|
||||
from litellm.proxy._experimental.mcp_server.sampling_handler import (
|
||||
handle_sampling_create_message,
|
||||
)
|
||||
|
||||
auth = _make_user_api_key_auth(models=["gpt-4o"])
|
||||
params = MagicMock()
|
||||
params.modelPreferences = None
|
||||
params.messages = []
|
||||
params.systemPrompt = None
|
||||
params.maxTokens = 100
|
||||
params.temperature = None
|
||||
params.stopSequences = None
|
||||
params.tools = None
|
||||
params.toolChoice = None
|
||||
params.metadata = None
|
||||
|
||||
budget_error = ErrorData(code=-1, message="ExceededBudget: over limit")
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.sampling_handler._check_model_access",
|
||||
new_callable=AsyncMock,
|
||||
return_value=None,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.sampling_handler._run_budget_checks",
|
||||
new_callable=AsyncMock,
|
||||
return_value=budget_error,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.sampling_handler._resolve_model_from_preferences",
|
||||
return_value="gpt-4o",
|
||||
),
|
||||
):
|
||||
result = await handle_sampling_create_message(
|
||||
context=MagicMock(),
|
||||
params=params,
|
||||
default_model="gpt-4o",
|
||||
user_api_key_auth=auth,
|
||||
)
|
||||
|
||||
assert result is budget_error
|
||||
assert "ExceededBudget" in result.message
|
||||
|
|
@ -0,0 +1,216 @@
|
|||
"""
|
||||
Tests for MCP sampling handler priority-based model selection.
|
||||
|
||||
Verifies that _resolve_model_from_preferences honours costPriority,
|
||||
speedPriority, and intelligencePriority when hints don't match,
|
||||
per the MCP spec.
|
||||
"""
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import patch
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.sampling_handler import (
|
||||
_has_priorities,
|
||||
_resolve_model_from_preferences,
|
||||
_select_model_by_priority,
|
||||
)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _prefs(*, hints=None, cost=None, speed=None, intelligence=None):
|
||||
"""Build a minimal ModelPreferences-like object."""
|
||||
return SimpleNamespace(
|
||||
hints=hints or [],
|
||||
costPriority=cost,
|
||||
speedPriority=speed,
|
||||
intelligencePriority=intelligence,
|
||||
)
|
||||
|
||||
|
||||
# Model info stubs keyed by model name
|
||||
_MODEL_INFO = {
|
||||
"gpt-3.5-turbo": {
|
||||
"input_cost_per_token": 0.0000005,
|
||||
"output_cost_per_token": 0.0000015,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"output_tokens_per_second": 50.0,
|
||||
},
|
||||
"gpt-4o": {
|
||||
"input_cost_per_token": 0.0000025,
|
||||
"output_cost_per_token": 0.0000100,
|
||||
"max_output_tokens": 16384,
|
||||
"max_tokens": 128000,
|
||||
"output_tokens_per_second": 60.0,
|
||||
},
|
||||
"claude-3-opus": {
|
||||
"input_cost_per_token": 0.0000150,
|
||||
"output_cost_per_token": 0.0000750,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 200000,
|
||||
"output_tokens_per_second": 20.0,
|
||||
},
|
||||
"gpt-4o-mini": {
|
||||
"input_cost_per_token": 0.00000015,
|
||||
"output_cost_per_token": 0.0000006,
|
||||
"max_output_tokens": 16384,
|
||||
"max_tokens": 128000,
|
||||
"output_tokens_per_second": 100.0,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _mock_get_model_info(model, **kwargs):
|
||||
"""Mock litellm.get_model_info using our test data."""
|
||||
if model in _MODEL_INFO:
|
||||
return _MODEL_INFO[model]
|
||||
raise Exception(f"Unknown model: {model}")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _has_priorities
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestHasPriorities:
|
||||
def test_should_return_false_when_no_priorities_set(self):
|
||||
prefs = _prefs()
|
||||
assert _has_priorities(prefs) is False
|
||||
|
||||
def test_should_return_false_when_all_zero(self):
|
||||
prefs = _prefs(cost=0, speed=0, intelligence=0)
|
||||
assert _has_priorities(prefs) is False
|
||||
|
||||
def test_should_return_true_when_cost_set(self):
|
||||
prefs = _prefs(cost=0.8)
|
||||
assert _has_priorities(prefs) is True
|
||||
|
||||
def test_should_return_true_when_intelligence_set(self):
|
||||
prefs = _prefs(intelligence=0.5)
|
||||
assert _has_priorities(prefs) is True
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _select_model_by_priority
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestSelectModelByPriority:
|
||||
"""Tests for the priority-based scoring logic."""
|
||||
|
||||
@patch("litellm.get_model_info", side_effect=_mock_get_model_info)
|
||||
def test_should_prefer_cheapest_when_cost_priority_high(self, _mock):
|
||||
"""High costPriority should select the cheapest model."""
|
||||
prefs = _prefs(cost=1.0, speed=0, intelligence=0)
|
||||
models = ["gpt-3.5-turbo", "gpt-4o", "claude-3-opus", "gpt-4o-mini"]
|
||||
result = _select_model_by_priority(models, prefs)
|
||||
# gpt-4o-mini has the lowest combined cost
|
||||
assert result == "gpt-4o-mini"
|
||||
|
||||
@patch("litellm.get_model_info", side_effect=_mock_get_model_info)
|
||||
def test_should_prefer_smartest_when_intelligence_priority_high(self, _mock):
|
||||
"""High intelligencePriority should select the model with highest max_output_tokens."""
|
||||
prefs = _prefs(cost=0, speed=0, intelligence=1.0)
|
||||
models = ["gpt-3.5-turbo", "gpt-4o", "claude-3-opus", "gpt-4o-mini"]
|
||||
result = _select_model_by_priority(models, prefs)
|
||||
# gpt-4o and gpt-4o-mini both have 16384 max_output_tokens (tied)
|
||||
# Either is acceptable
|
||||
assert result in ("gpt-4o", "gpt-4o-mini")
|
||||
|
||||
@patch("litellm.get_model_info", side_effect=_mock_get_model_info)
|
||||
def test_should_balance_cost_and_intelligence(self, _mock):
|
||||
"""Balanced priorities should pick a middle-ground model."""
|
||||
prefs = _prefs(cost=0.5, speed=0, intelligence=0.5)
|
||||
models = ["gpt-3.5-turbo", "gpt-4o", "claude-3-opus", "gpt-4o-mini"]
|
||||
result = _select_model_by_priority(models, prefs)
|
||||
# gpt-4o-mini is cheap AND has high max_output_tokens → best balance
|
||||
assert result == "gpt-4o-mini"
|
||||
|
||||
@patch("litellm.get_model_info", side_effect=_mock_get_model_info)
|
||||
def test_should_prefer_fastest_when_speed_priority_high(self, _mock):
|
||||
"""High speedPriority should prefer cheaper (faster proxy) models."""
|
||||
prefs = _prefs(cost=0, speed=1.0, intelligence=0)
|
||||
models = ["gpt-3.5-turbo", "gpt-4o", "claude-3-opus", "gpt-4o-mini"]
|
||||
result = _select_model_by_priority(models, prefs)
|
||||
# gpt-4o-mini has lowest cost → fastest proxy
|
||||
assert result == "gpt-4o-mini"
|
||||
|
||||
@patch(
|
||||
"litellm.get_model_info",
|
||||
side_effect=lambda m, **kw: (_ for _ in ()).throw(Exception("no info")),
|
||||
)
|
||||
def test_should_return_none_when_no_model_info(self, _mock):
|
||||
"""If get_model_info fails for all models, return None."""
|
||||
prefs = _prefs(cost=1.0)
|
||||
models = ["unknown-model-1", "unknown-model-2"]
|
||||
result = _select_model_by_priority(models, prefs)
|
||||
assert result is None
|
||||
|
||||
@patch("litellm.get_model_info", side_effect=_mock_get_model_info)
|
||||
def test_should_handle_single_model(self, _mock):
|
||||
"""Single model should always be returned regardless of priorities."""
|
||||
prefs = _prefs(cost=1.0, intelligence=1.0)
|
||||
result = _select_model_by_priority(["gpt-4o"], prefs)
|
||||
assert result == "gpt-4o"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _resolve_model_from_preferences — priority integration
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestResolveModelPriorityIntegration:
|
||||
"""End-to-end tests for priority selection within _resolve_model_from_preferences."""
|
||||
|
||||
@patch("litellm.get_model_info", side_effect=_mock_get_model_info)
|
||||
@patch("litellm.proxy.proxy_server.llm_router", None)
|
||||
@patch(
|
||||
"litellm.model_list",
|
||||
[
|
||||
{"model_name": "gpt-3.5-turbo"},
|
||||
{"model_name": "gpt-4o"},
|
||||
{"model_name": "gpt-4o-mini"},
|
||||
],
|
||||
)
|
||||
def test_should_use_priority_when_hints_empty(self, _mock_info):
|
||||
"""With no hints but priorities set, should use priority-based selection."""
|
||||
prefs = _prefs(cost=1.0, speed=0, intelligence=0)
|
||||
result = _resolve_model_from_preferences(prefs, default_model="gpt-4o")
|
||||
# Should pick cheapest, NOT fall through to default_model
|
||||
assert result == "gpt-4o-mini"
|
||||
|
||||
@patch("litellm.get_model_info", side_effect=_mock_get_model_info)
|
||||
@patch("litellm.proxy.proxy_server.llm_router", None)
|
||||
@patch(
|
||||
"litellm.model_list",
|
||||
[
|
||||
{"model_name": "gpt-3.5-turbo"},
|
||||
{"model_name": "gpt-4o"},
|
||||
{"model_name": "gpt-4o-mini"},
|
||||
],
|
||||
)
|
||||
def test_should_skip_priority_when_no_priorities_set(self, _mock_info):
|
||||
"""With no priorities set, should fall through to default_model."""
|
||||
prefs = _prefs() # no priorities
|
||||
result = _resolve_model_from_preferences(prefs, default_model="gpt-4o")
|
||||
assert result == "gpt-4o"
|
||||
|
||||
@patch("litellm.get_model_info", side_effect=_mock_get_model_info)
|
||||
@patch("litellm.proxy.proxy_server.llm_router", None)
|
||||
@patch(
|
||||
"litellm.model_list",
|
||||
[
|
||||
{"model_name": "gpt-3.5-turbo"},
|
||||
{"model_name": "gpt-4o"},
|
||||
{"model_name": "gpt-4o-mini"},
|
||||
],
|
||||
)
|
||||
def test_should_prefer_hint_over_priority(self, _mock_info):
|
||||
"""Hints should take precedence over priority-based selection."""
|
||||
hints = [SimpleNamespace(name="gpt-4o")]
|
||||
prefs = _prefs(hints=hints, cost=1.0) # cost says cheap, but hint says gpt-4o
|
||||
result = _resolve_model_from_preferences(prefs, default_model="gpt-3.5-turbo")
|
||||
assert result == "gpt-4o"
|
||||
|
|
@ -0,0 +1,147 @@
|
|||
"""
|
||||
Tests for _build_sampling_request header forwarding.
|
||||
|
||||
Verifies that the synthetic FastAPI Request built for sampling sub-calls
|
||||
correctly propagates the original MCP connection's headers and client IP
|
||||
so that header-dependent guardrails, routing hooks, and trace correlation
|
||||
function correctly.
|
||||
"""
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.sampling_handler import (
|
||||
_build_sampling_request,
|
||||
)
|
||||
|
||||
|
||||
class TestBuildSamplingRequest:
|
||||
"""Tests for the _build_sampling_request helper."""
|
||||
|
||||
def test_should_include_content_type_by_default(self):
|
||||
"""Even with no raw headers, content-type must be present."""
|
||||
req = _build_sampling_request()
|
||||
headers = dict(req.headers)
|
||||
assert headers.get("content-type") == "application/json"
|
||||
|
||||
def test_should_forward_raw_headers(self):
|
||||
"""Headers from the original MCP connection should be forwarded."""
|
||||
raw = {
|
||||
"x-litellm-tags": "tag1,tag2",
|
||||
"x-litellm-trace-id": "trace-abc-123",
|
||||
"user-agent": "MCP-Client/1.0",
|
||||
"authorization": "Bearer sk-test",
|
||||
}
|
||||
req = _build_sampling_request(raw_headers=raw)
|
||||
headers = dict(req.headers)
|
||||
|
||||
assert headers.get("x-litellm-tags") == "tag1,tag2"
|
||||
assert headers.get("x-litellm-trace-id") == "trace-abc-123"
|
||||
assert headers.get("user-agent") == "MCP-Client/1.0"
|
||||
assert headers.get("authorization") == "Bearer sk-test"
|
||||
|
||||
def test_should_skip_hop_by_hop_headers(self):
|
||||
"""content-length and transfer-encoding should not be forwarded."""
|
||||
raw = {
|
||||
"content-length": "42",
|
||||
"transfer-encoding": "chunked",
|
||||
"x-custom": "keep-me",
|
||||
}
|
||||
req = _build_sampling_request(raw_headers=raw)
|
||||
headers = dict(req.headers)
|
||||
|
||||
assert "content-length" not in headers
|
||||
assert "transfer-encoding" not in headers
|
||||
assert headers.get("x-custom") == "keep-me"
|
||||
|
||||
def test_should_not_duplicate_content_type(self):
|
||||
"""If raw_headers includes content-type, don't add it twice."""
|
||||
raw = {"content-type": "text/plain"}
|
||||
req = _build_sampling_request(raw_headers=raw)
|
||||
# Count how many content-type headers are present
|
||||
ct_count = sum(1 for k, _ in req.scope["headers"] if k == b"content-type")
|
||||
assert ct_count == 1
|
||||
|
||||
def test_should_inject_client_ip_as_x_forwarded_for(self):
|
||||
"""client_ip should be injected as x-forwarded-for."""
|
||||
req = _build_sampling_request(client_ip="10.0.0.42")
|
||||
headers = dict(req.headers)
|
||||
assert headers.get("x-forwarded-for") == "10.0.0.42"
|
||||
|
||||
def test_should_not_override_existing_x_forwarded_for(self):
|
||||
"""Caller-supplied x-forwarded-for is stripped; resolved client_ip wins."""
|
||||
raw = {"x-forwarded-for": "192.168.1.1"}
|
||||
req = _build_sampling_request(raw_headers=raw, client_ip="10.0.0.42")
|
||||
headers = dict(req.headers)
|
||||
assert headers.get("x-forwarded-for") == "10.0.0.42"
|
||||
|
||||
def test_should_set_correct_path(self):
|
||||
"""The synthetic request should have the sampling path."""
|
||||
req = _build_sampling_request()
|
||||
assert req.scope["path"] == "/mcp/sampling/createMessage"
|
||||
|
||||
def test_server_should_default_to_litellm_port(self):
|
||||
"""Server tuple should use port 4000 (LiteLLM default), not 0."""
|
||||
req = _build_sampling_request()
|
||||
_host, _port = req.scope["server"]
|
||||
assert _port == 4000, f"Expected default LiteLLM port 4000, got {_port}"
|
||||
|
||||
def test_should_populate_client_tuple_from_client_ip(self):
|
||||
"""request.client.host must return the real client IP for
|
||||
IP-based routing and guardrails."""
|
||||
req = _build_sampling_request(client_ip="10.0.0.42")
|
||||
assert req.scope.get("client") is not None
|
||||
assert req.scope["client"][0] == "10.0.0.42"
|
||||
# Verify request.client.host works (Starlette Address)
|
||||
assert req.client is not None
|
||||
assert req.client.host == "10.0.0.42"
|
||||
|
||||
def test_should_not_set_client_when_no_ip(self):
|
||||
"""If no client_ip is provided, client should not be in scope."""
|
||||
req = _build_sampling_request()
|
||||
assert "client" not in req.scope
|
||||
|
||||
def test_should_skip_all_hop_by_hop_headers(self):
|
||||
"""All hop-by-hop headers must be filtered, not just content-length
|
||||
and transfer-encoding."""
|
||||
raw = {
|
||||
"content-length": "42",
|
||||
"transfer-encoding": "chunked",
|
||||
"connection": "keep-alive",
|
||||
"keep-alive": "timeout=5",
|
||||
"upgrade": "websocket",
|
||||
"te": "trailers",
|
||||
"trailer": "Expires",
|
||||
"x-custom": "keep-me",
|
||||
}
|
||||
req = _build_sampling_request(raw_headers=raw)
|
||||
headers = dict(req.headers)
|
||||
|
||||
for hop_header in [
|
||||
"content-length",
|
||||
"transfer-encoding",
|
||||
"connection",
|
||||
"keep-alive",
|
||||
"upgrade",
|
||||
"te",
|
||||
"trailer",
|
||||
]:
|
||||
assert (
|
||||
hop_header not in headers
|
||||
), f"Hop-by-hop header '{hop_header}' should be filtered"
|
||||
assert headers.get("x-custom") == "keep-me"
|
||||
|
||||
def test_should_forward_traceparent_header(self):
|
||||
"""traceparent header must be forwarded for trace correlation."""
|
||||
raw = {
|
||||
"traceparent": "00-abcdef1234567890abcdef1234567890-1234567890abcdef-01",
|
||||
}
|
||||
req = _build_sampling_request(raw_headers=raw)
|
||||
headers = dict(req.headers)
|
||||
assert headers.get("traceparent") == (
|
||||
"00-abcdef1234567890abcdef1234567890-1234567890abcdef-01"
|
||||
)
|
||||
|
||||
def test_should_forward_x_litellm_api_key(self):
|
||||
"""x-litellm-api-key header must be forwarded for auth."""
|
||||
raw = {"x-litellm-api-key": "sk-proxy-key-123"}
|
||||
req = _build_sampling_request(raw_headers=raw)
|
||||
headers = dict(req.headers)
|
||||
assert headers.get("x-litellm-api-key") == "sk-proxy-key-123"
|
||||
|
|
@ -0,0 +1,257 @@
|
|||
"""
|
||||
Tests for MCP sampling handler tool_use / tool_result content conversion.
|
||||
|
||||
Verifies that multi-turn tool-calling conversations from upstream MCP
|
||||
servers are faithfully converted to OpenAI format instead of being
|
||||
reduced to lossy plain-text stubs.
|
||||
"""
|
||||
|
||||
import json
|
||||
from types import SimpleNamespace
|
||||
from typing import Any, Dict
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.sampling_handler import (
|
||||
_convert_mcp_messages_to_openai,
|
||||
_convert_single_content,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers — lightweight MCP type stand-ins
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _text(text: str) -> SimpleNamespace:
|
||||
return SimpleNamespace(type="text", text=text)
|
||||
|
||||
|
||||
def _tool_use(*, name: str, tool_id: str, input_data: Dict[str, Any]) -> SimpleNamespace:
|
||||
return SimpleNamespace(type="tool_use", name=name, id=tool_id, input=input_data)
|
||||
|
||||
|
||||
def _tool_result(
|
||||
*, tool_use_id: str, content: Any = None, is_error: bool = False
|
||||
) -> SimpleNamespace:
|
||||
if content is None:
|
||||
content = []
|
||||
return SimpleNamespace(
|
||||
type="tool_result", toolUseId=tool_use_id, content=content, isError=is_error
|
||||
)
|
||||
|
||||
|
||||
def _sampling_msg(role: str, content: Any) -> SimpleNamespace:
|
||||
return SimpleNamespace(role=role, content=content)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _convert_single_content — tool_use
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestConvertSingleContentToolUse:
|
||||
"""Tests for the tool_use branch of _convert_single_content."""
|
||||
|
||||
def test_should_produce_function_call_dict(self):
|
||||
"""tool_use must produce a proper function-call dict, not a text stub."""
|
||||
tu = _tool_use(name="get_weather", tool_id="call_123", input_data={"city": "NYC"})
|
||||
result = _convert_single_content(tu)
|
||||
|
||||
assert result["_marker_type"] == "tool_use"
|
||||
assert result["type"] == "function"
|
||||
assert result["id"] == "call_123"
|
||||
assert result["function"]["name"] == "get_weather"
|
||||
assert json.loads(result["function"]["arguments"]) == {"city": "NYC"}
|
||||
|
||||
def test_should_not_produce_text_stub(self):
|
||||
"""Regression: the old code produced '[Tool call: get_weather]'."""
|
||||
tu = _tool_use(name="get_weather", tool_id="call_1", input_data={})
|
||||
result = _convert_single_content(tu)
|
||||
|
||||
# Must NOT be a text content part
|
||||
assert result.get("type") != "text"
|
||||
assert "Tool call" not in str(result)
|
||||
|
||||
def test_should_handle_empty_input(self):
|
||||
tu = _tool_use(name="no_args_tool", tool_id="call_2", input_data={})
|
||||
result = _convert_single_content(tu)
|
||||
|
||||
assert json.loads(result["function"]["arguments"]) == {}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _convert_single_content — tool_result
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestConvertSingleContentToolResult:
|
||||
"""Tests for the tool_result branch of _convert_single_content."""
|
||||
|
||||
def test_should_produce_tool_role_message(self):
|
||||
"""tool_result must produce a tool-role dict, not a text content part."""
|
||||
tr = _tool_result(
|
||||
tool_use_id="call_123",
|
||||
content=[_text("Temperature: 72°F")],
|
||||
)
|
||||
result = _convert_single_content(tr)
|
||||
|
||||
assert result["_marker_type"] == "tool_result"
|
||||
assert result["role"] == "tool"
|
||||
assert result["tool_call_id"] == "call_123"
|
||||
assert "72°F" in result["content"]
|
||||
|
||||
def test_should_handle_empty_content(self):
|
||||
tr = _tool_result(tool_use_id="call_456", content=[])
|
||||
result = _convert_single_content(tr)
|
||||
|
||||
assert result["role"] == "tool"
|
||||
assert result["tool_call_id"] == "call_456"
|
||||
assert result["content"] == ""
|
||||
|
||||
def test_should_concatenate_multiple_text_parts(self):
|
||||
tr = _tool_result(
|
||||
tool_use_id="call_789",
|
||||
content=[_text("Line 1"), _text("Line 2")],
|
||||
)
|
||||
result = _convert_single_content(tr)
|
||||
assert "Line 1" in result["content"]
|
||||
assert "Line 2" in result["content"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _convert_mcp_messages_to_openai — multi-turn tool calling
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestConvertMcpMessagesMultiTurnTools:
|
||||
"""End-to-end tests for multi-turn tool-calling message sequences."""
|
||||
|
||||
def test_should_convert_assistant_tool_use_to_tool_calls_array(self):
|
||||
"""An assistant message with tool_use content should produce
|
||||
a proper tool_calls array, not a text stub."""
|
||||
messages = [
|
||||
_sampling_msg("assistant", _tool_use(
|
||||
name="search", tool_id="call_1", input_data={"query": "LiteLLM"}
|
||||
)),
|
||||
]
|
||||
result = _convert_mcp_messages_to_openai(messages)
|
||||
|
||||
assert len(result) == 1
|
||||
msg = result[0]
|
||||
assert msg["role"] == "assistant"
|
||||
assert "tool_calls" in msg
|
||||
assert len(msg["tool_calls"]) == 1
|
||||
tc = msg["tool_calls"][0]
|
||||
assert tc["function"]["name"] == "search"
|
||||
assert tc["id"] == "call_1"
|
||||
|
||||
def test_should_convert_user_tool_result_to_tool_role_message(self):
|
||||
"""A user message with tool_result content should produce
|
||||
a separate role='tool' message."""
|
||||
messages = [
|
||||
_sampling_msg("user", _tool_result(
|
||||
tool_use_id="call_1",
|
||||
content=[_text("Found 42 results")],
|
||||
)),
|
||||
]
|
||||
result = _convert_mcp_messages_to_openai(messages)
|
||||
|
||||
assert len(result) == 1
|
||||
msg = result[0]
|
||||
assert msg["role"] == "tool"
|
||||
assert msg["tool_call_id"] == "call_1"
|
||||
assert "42 results" in msg["content"]
|
||||
|
||||
def test_should_handle_full_tool_calling_round_trip(self):
|
||||
"""Simulate a complete tool-calling conversation:
|
||||
user → assistant(tool_use) → user(tool_result) → assistant(text)
|
||||
"""
|
||||
messages = [
|
||||
_sampling_msg("user", _text("What's the weather in NYC?")),
|
||||
_sampling_msg("assistant", _tool_use(
|
||||
name="get_weather", tool_id="call_w1",
|
||||
input_data={"city": "NYC"},
|
||||
)),
|
||||
_sampling_msg("user", _tool_result(
|
||||
tool_use_id="call_w1",
|
||||
content=[_text("72°F, sunny")],
|
||||
)),
|
||||
_sampling_msg("assistant", _text("It's 72°F and sunny in NYC!")),
|
||||
]
|
||||
result = _convert_mcp_messages_to_openai(messages)
|
||||
|
||||
assert len(result) == 4
|
||||
|
||||
# 1. User message
|
||||
assert result[0]["role"] == "user"
|
||||
|
||||
# 2. Assistant with tool_calls
|
||||
assert result[1]["role"] == "assistant"
|
||||
assert "tool_calls" in result[1]
|
||||
assert result[1]["tool_calls"][0]["function"]["name"] == "get_weather"
|
||||
|
||||
# 3. Tool result
|
||||
assert result[2]["role"] == "tool"
|
||||
assert result[2]["tool_call_id"] == "call_w1"
|
||||
|
||||
# 4. Final assistant text
|
||||
assert result[3]["role"] == "assistant"
|
||||
assert "72°F" in str(result[3]["content"])
|
||||
|
||||
def test_should_handle_mixed_text_and_tool_use_in_assistant(self):
|
||||
"""An assistant message with both text and tool_use content."""
|
||||
messages = [
|
||||
_sampling_msg("assistant", [
|
||||
_text("Let me check that for you."),
|
||||
_tool_use(name="lookup", tool_id="call_lu1", input_data={"id": 42}),
|
||||
]),
|
||||
]
|
||||
result = _convert_mcp_messages_to_openai(messages)
|
||||
|
||||
assert len(result) == 1
|
||||
msg = result[0]
|
||||
assert msg["role"] == "assistant"
|
||||
assert "tool_calls" in msg
|
||||
assert msg["tool_calls"][0]["function"]["name"] == "lookup"
|
||||
# Text content should also be present
|
||||
assert msg.get("content") is not None
|
||||
|
||||
def test_should_handle_multiple_tool_uses_in_single_message(self):
|
||||
"""Multiple tool_use items in a single assistant message → multiple tool_calls."""
|
||||
messages = [
|
||||
_sampling_msg("assistant", [
|
||||
_tool_use(name="tool_a", tool_id="call_a", input_data={}),
|
||||
_tool_use(name="tool_b", tool_id="call_b", input_data={"x": 1}),
|
||||
]),
|
||||
]
|
||||
result = _convert_mcp_messages_to_openai(messages)
|
||||
|
||||
assert len(result) == 1
|
||||
msg = result[0]
|
||||
assert len(msg["tool_calls"]) == 2
|
||||
names = {tc["function"]["name"] for tc in msg["tool_calls"]}
|
||||
assert names == {"tool_a", "tool_b"}
|
||||
|
||||
def test_should_handle_multiple_tool_results_in_single_message(self):
|
||||
"""Multiple tool_result items in a single user message → multiple tool messages."""
|
||||
messages = [
|
||||
_sampling_msg("user", [
|
||||
_tool_result(tool_use_id="call_a", content=[_text("Result A")]),
|
||||
_tool_result(tool_use_id="call_b", content=[_text("Result B")]),
|
||||
]),
|
||||
]
|
||||
result = _convert_mcp_messages_to_openai(messages)
|
||||
|
||||
assert len(result) == 2
|
||||
assert all(m["role"] == "tool" for m in result)
|
||||
ids = {m["tool_call_id"] for m in result}
|
||||
assert ids == {"call_a", "call_b"}
|
||||
|
||||
def test_should_preserve_system_prompt(self):
|
||||
"""System prompt should still be emitted first."""
|
||||
messages = [_sampling_msg("user", _text("Hi"))]
|
||||
result = _convert_mcp_messages_to_openai(
|
||||
messages, system_prompt="You are helpful."
|
||||
)
|
||||
|
||||
assert result[0]["role"] == "system"
|
||||
assert result[0]["content"] == "You are helpful."
|
||||
|
|
@ -1,5 +1,4 @@
|
|||
import asyncio
|
||||
import contextlib
|
||||
import contextvars
|
||||
from datetime import datetime, timedelta
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
|
@ -894,7 +893,7 @@ async def test_get_tools_from_mcp_servers_continues_when_one_server_fails():
|
|||
extra_headers=None,
|
||||
add_prefix=True,
|
||||
raw_headers=None,
|
||||
user_api_key_auth=None,
|
||||
**kwargs,
|
||||
):
|
||||
if server.name == "working_server":
|
||||
# Working server returns tools
|
||||
|
|
@ -1000,7 +999,7 @@ async def test_get_tools_from_mcp_servers_handles_all_servers_failing():
|
|||
extra_headers=None,
|
||||
add_prefix=True,
|
||||
raw_headers=None,
|
||||
user_api_key_auth=None,
|
||||
**kwargs,
|
||||
):
|
||||
# All servers fail
|
||||
raise Exception(f"Server {server.name} connection failed")
|
||||
|
|
@ -1122,8 +1121,8 @@ async def test_concurrent_initialize_session_managers():
|
|||
# Reset state before test
|
||||
original_initialized = mcp_server._SESSION_MANAGERS_INITIALIZED
|
||||
original_session_cm = mcp_server._session_manager_cm
|
||||
original_session_stateful_cm = mcp_server._session_manager_stateful_cm
|
||||
original_sse_session_cm = mcp_server._sse_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:
|
||||
|
|
@ -1131,30 +1130,38 @@ async def test_concurrent_initialize_session_managers():
|
|||
mcp_server._session_manager_cm = None
|
||||
mcp_server._session_manager_stateful_cm = None
|
||||
mcp_server._sse_session_manager_cm = None
|
||||
mcp_server._stateful_auth_context_cleanup_task = None
|
||||
|
||||
# Mock the session managers to avoid actual MCP initialization
|
||||
# Create mock context managers for all three session managers
|
||||
mock_cm_stateless = AsyncMock()
|
||||
mock_cm_stateless.__aenter__ = AsyncMock()
|
||||
mock_cm_stateless.__aexit__ = AsyncMock()
|
||||
|
||||
mock_cm_stateful = AsyncMock()
|
||||
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(
|
||||
"litellm.proxy._experimental.mcp_server.server.session_manager_stateless"
|
||||
) as mock_session_manager_stateless,
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.session_manager_stateful"
|
||||
) as mock_session_manager_stateful,
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.sse_session_manager"
|
||||
) as mock_sse_session_manager,
|
||||
patch.object(
|
||||
mcp_server.session_manager_stateless,
|
||||
"run",
|
||||
return_value=mock_cm_stateless,
|
||||
) as mock_stateless_run,
|
||||
patch.object(
|
||||
mcp_server.session_manager_stateful,
|
||||
"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.server.verbose_logger"),
|
||||
):
|
||||
# Mock the run() method to return a mock context manager
|
||||
mock_cm = AsyncMock()
|
||||
mock_cm.__aenter__ = AsyncMock()
|
||||
mock_cm.__aexit__ = AsyncMock()
|
||||
|
||||
mock_session_manager_stateless.run.return_value = mock_cm
|
||||
mock_session_manager_stateful.run.return_value = mock_cm
|
||||
mock_sse_session_manager.run.return_value = mock_cm
|
||||
|
||||
# Create multiple concurrent tasks that call initialize_session_managers
|
||||
async def init_task():
|
||||
await initialize_session_managers()
|
||||
|
|
@ -1171,19 +1178,25 @@ async def test_concurrent_initialize_session_managers():
|
|||
|
||||
# Each session manager.run() should only be called once due to the lock
|
||||
assert (
|
||||
mock_session_manager_stateless.run.call_count == 1
|
||||
), f"Expected 1 call to session_manager_stateless.run(), got {mock_session_manager_stateless.run.call_count}"
|
||||
mock_stateless_run.call_count == 1
|
||||
), f"Expected 1 call to session_manager_stateless.run(), got {mock_stateless_run.call_count}"
|
||||
assert (
|
||||
mock_session_manager_stateful.run.call_count == 1
|
||||
), f"Expected 1 call to session_manager_stateful.run(), got {mock_session_manager_stateful.run.call_count}"
|
||||
mock_stateful_run.call_count == 1
|
||||
), f"Expected 1 call to session_manager_stateful.run(), got {mock_stateful_run.call_count}"
|
||||
assert (
|
||||
mock_sse_session_manager.run.call_count == 1
|
||||
), f"Expected 1 call to sse_session_manager.run(), got {mock_sse_session_manager.run.call_count}"
|
||||
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 (3 managers)
|
||||
# The context managers should only be entered once each
|
||||
assert (
|
||||
mock_cm.__aenter__.call_count == 3
|
||||
), f"Expected 3 calls to __aenter__ (one per session manager), got {mock_cm.__aenter__.call_count}"
|
||||
mock_cm_stateless.__aenter__.call_count == 1
|
||||
), f"Expected 1 call to stateless __aenter__, got {mock_cm_stateless.__aenter__.call_count}"
|
||||
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
|
||||
|
|
@ -1195,14 +1208,12 @@ async def test_concurrent_initialize_session_managers():
|
|||
leaked_task = mcp_server._stateful_auth_context_cleanup_task
|
||||
if leaked_task is not None and leaked_task is not original_cleanup_task:
|
||||
leaked_task.cancel()
|
||||
with contextlib.suppress(asyncio.CancelledError, Exception):
|
||||
await leaked_task
|
||||
|
||||
# Restore original state
|
||||
mcp_server._SESSION_MANAGERS_INITIALIZED = original_initialized
|
||||
mcp_server._session_manager_cm = original_session_cm
|
||||
mcp_server._session_manager_stateful_cm = original_session_stateful_cm
|
||||
mcp_server._sse_session_manager_cm = original_sse_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
|
||||
|
||||
|
||||
|
|
@ -1637,10 +1648,7 @@ async def test_mcp_routing_initialize_rejected_when_owner_at_session_cap():
|
|||
active = {f"s{i}": 1 for i in range(cap)} # all in flight -> cannot evict
|
||||
contexts = {f"s{i}": MagicMock() for i in range(cap)}
|
||||
|
||||
init_body = (
|
||||
b'{"jsonrpc":"2.0","id":1,"method":"initialize",'
|
||||
b'"params":{"protocolVersion":"2024-11-05"}}'
|
||||
)
|
||||
init_body = b'{"jsonrpc":"2.0","id":1,"method":"initialize","params":{"protocolVersion":"2024-11-05"}}'
|
||||
scope = {
|
||||
"type": "http",
|
||||
"method": "POST",
|
||||
|
|
@ -2729,7 +2737,7 @@ async def test_oauth2_headers_passed_to_mcp_client():
|
|||
mcp_auth_header=None,
|
||||
extra_headers=None,
|
||||
stdio_env=None,
|
||||
subject_token=None,
|
||||
**kwargs,
|
||||
):
|
||||
# Capture the arguments for verification
|
||||
captured_client_args.update(
|
||||
|
|
@ -2738,7 +2746,7 @@ async def test_oauth2_headers_passed_to_mcp_client():
|
|||
"mcp_auth_header": mcp_auth_header,
|
||||
"extra_headers": extra_headers,
|
||||
"stdio_env": stdio_env,
|
||||
"subject_token": subject_token,
|
||||
"kwargs": kwargs,
|
||||
}
|
||||
)
|
||||
# Return a mock client that doesn't actually connect
|
||||
|
|
@ -2764,6 +2772,16 @@ async def test_oauth2_headers_passed_to_mcp_client():
|
|||
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
|
||||
AsyncMock(return_value=[oauth2_server]),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._prefetch_oauth_creds_for_user",
|
||||
new_callable=AsyncMock,
|
||||
return_value={},
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._get_user_oauth_extra_headers_from_db",
|
||||
new_callable=AsyncMock,
|
||||
return_value=None,
|
||||
),
|
||||
):
|
||||
# Call _get_tools_from_mcp_servers which should eventually call _create_mcp_client
|
||||
await _get_tools_from_mcp_servers(
|
||||
|
|
@ -2840,7 +2858,7 @@ async def test_list_tools_single_server_unprefixed_names():
|
|||
extra_headers=None,
|
||||
add_prefix=False,
|
||||
raw_headers=None,
|
||||
user_api_key_auth=None,
|
||||
**kwargs,
|
||||
):
|
||||
tool = MagicMock()
|
||||
tool.name = f"{server.alias}-toolA" if add_prefix else "toolA"
|
||||
|
|
@ -2922,7 +2940,7 @@ async def test_list_tools_multiple_servers_prefixed_names():
|
|||
extra_headers=None,
|
||||
add_prefix=True,
|
||||
raw_headers=None,
|
||||
user_api_key_auth=None,
|
||||
**kwargs,
|
||||
):
|
||||
tool = MagicMock()
|
||||
# When multiple servers, add_prefix should be True -> prefixed names
|
||||
|
|
@ -3189,7 +3207,7 @@ async def test_list_tools_filters_by_key_team_permissions():
|
|||
extra_headers=None,
|
||||
add_prefix=False,
|
||||
raw_headers=None,
|
||||
user_api_key_auth=None,
|
||||
**kwargs,
|
||||
):
|
||||
# Return 4 tools, but only 2 should be allowed
|
||||
tool1 = MagicMock()
|
||||
|
|
@ -3299,7 +3317,7 @@ async def test_list_tools_with_team_tool_permissions_inheritance():
|
|||
extra_headers=None,
|
||||
add_prefix=False,
|
||||
raw_headers=None,
|
||||
user_api_key_auth=None,
|
||||
**kwargs,
|
||||
):
|
||||
# Return 4 tools
|
||||
tool1 = MagicMock()
|
||||
|
|
@ -3395,7 +3413,7 @@ async def test_list_tools_with_no_tool_permissions_shows_all():
|
|||
extra_headers=None,
|
||||
add_prefix=False,
|
||||
raw_headers=None,
|
||||
user_api_key_auth=None,
|
||||
**kwargs,
|
||||
):
|
||||
# Return 3 tools
|
||||
tool1 = MagicMock()
|
||||
|
|
@ -3494,7 +3512,7 @@ async def test_list_tools_strips_prefix_when_matching_permissions():
|
|||
extra_headers=None,
|
||||
add_prefix=True,
|
||||
raw_headers=None,
|
||||
user_api_key_auth=None,
|
||||
**kwargs,
|
||||
):
|
||||
# Return tools WITH prefix (as they come from MCP server)
|
||||
tool1 = MagicMock()
|
||||
|
|
@ -5178,3 +5196,42 @@ def test_get_forwarded_auth_from_scope_skips_when_no_litellm_key_header():
|
|||
]
|
||||
}
|
||||
assert _get_forwarded_auth_from_scope(scope) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_mcp_client_sampling_disabled_by_default():
|
||||
"""Sampling callback must be None when allow_sampling is not set (default False)."""
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
MCPServerManager,
|
||||
)
|
||||
|
||||
manager = MCPServerManager()
|
||||
server = MCPServer(
|
||||
server_id="no-sampling",
|
||||
name="no-sampling",
|
||||
url="https://example.com/mcp",
|
||||
transport=MCPTransport.http,
|
||||
)
|
||||
|
||||
client = await manager._create_mcp_client(server=server)
|
||||
assert client._sampling_callback is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_mcp_client_sampling_enabled():
|
||||
"""Sampling callback must be set when allow_sampling=True."""
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
MCPServerManager,
|
||||
)
|
||||
|
||||
manager = MCPServerManager()
|
||||
server = MCPServer(
|
||||
server_id="with-sampling",
|
||||
name="with-sampling",
|
||||
url="https://example.com/mcp",
|
||||
transport=MCPTransport.http,
|
||||
allow_sampling=True,
|
||||
)
|
||||
|
||||
client = await manager._create_mcp_client(server=server)
|
||||
assert client._sampling_callback is not None
|
||||
|
|
|
|||
|
|
@ -320,9 +320,7 @@ class TestMCPServerManager:
|
|||
async def mock_get_tools_from_server(
|
||||
server,
|
||||
mcp_auth_header=None,
|
||||
mcp_protocol_version=None,
|
||||
raw_headers=None,
|
||||
user_api_key_auth=None,
|
||||
**kwargs,
|
||||
):
|
||||
if server.name == "github":
|
||||
tool1 = MagicMock()
|
||||
|
|
@ -375,9 +373,7 @@ class TestMCPServerManager:
|
|||
async def mock_get_tools_from_server(
|
||||
server,
|
||||
mcp_auth_header=None,
|
||||
mcp_protocol_version=None,
|
||||
raw_headers=None,
|
||||
user_api_key_auth=None,
|
||||
**kwargs,
|
||||
):
|
||||
assert mcp_auth_header == "legacy-token" # Should use legacy header
|
||||
tool = MagicMock()
|
||||
|
|
@ -414,9 +410,7 @@ class TestMCPServerManager:
|
|||
async def mock_get_tools_from_server(
|
||||
server,
|
||||
mcp_auth_header=None,
|
||||
mcp_protocol_version=None,
|
||||
raw_headers=None,
|
||||
user_api_key_auth=None,
|
||||
**kwargs,
|
||||
):
|
||||
assert (
|
||||
mcp_auth_header == "server-specific-token"
|
||||
|
|
@ -457,7 +451,7 @@ class TestMCPServerManager:
|
|||
captured_extra_headers = None
|
||||
|
||||
async def capture_create_mcp_client(
|
||||
server, mcp_auth_header, extra_headers, stdio_env, subject_token=None
|
||||
server, mcp_auth_header, extra_headers, stdio_env, **kwargs
|
||||
): # pragma: no cover - helper
|
||||
nonlocal captured_extra_headers
|
||||
captured_extra_headers = extra_headers
|
||||
|
|
@ -507,7 +501,7 @@ class TestMCPServerManager:
|
|||
captured_extra_headers = None
|
||||
|
||||
async def capture_create_mcp_client(
|
||||
server, mcp_auth_header, extra_headers, stdio_env, subject_token=None
|
||||
server, mcp_auth_header, extra_headers, stdio_env, subject_token=None, **kwargs
|
||||
): # pragma: no cover - helper
|
||||
nonlocal captured_extra_headers
|
||||
captured_extra_headers = extra_headers
|
||||
|
|
@ -560,7 +554,7 @@ class TestMCPServerManager:
|
|||
captured_extra_headers = None
|
||||
|
||||
async def capture_create_mcp_client(
|
||||
server, mcp_auth_header, extra_headers, stdio_env, subject_token=None
|
||||
server, mcp_auth_header, extra_headers, stdio_env, subject_token=None, **kwargs
|
||||
): # pragma: no cover - helper
|
||||
nonlocal captured_extra_headers
|
||||
captured_extra_headers = extra_headers
|
||||
|
|
@ -616,7 +610,7 @@ class TestMCPServerManager:
|
|||
captured_extra_headers = None
|
||||
|
||||
async def capture_create_mcp_client(
|
||||
server, mcp_auth_header, extra_headers, stdio_env, subject_token=None
|
||||
server, mcp_auth_header, extra_headers, stdio_env, subject_token=None, **kwargs
|
||||
): # pragma: no cover - helper
|
||||
nonlocal captured_extra_headers
|
||||
captured_extra_headers = extra_headers
|
||||
|
|
@ -1166,9 +1160,7 @@ class TestMCPServerManager:
|
|||
async def mock_get_tools_from_server(
|
||||
server,
|
||||
mcp_auth_header=None,
|
||||
mcp_protocol_version=None,
|
||||
raw_headers=None,
|
||||
user_api_key_auth=None,
|
||||
**kwargs,
|
||||
):
|
||||
assert (
|
||||
mcp_auth_header == "server-specific-token"
|
||||
|
|
|
|||
|
|
@ -9,8 +9,6 @@ they may send a stale `mcp-session-id` header. This test verifies that:
|
|||
|
||||
import asyncio
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
from fastapi import HTTPException
|
||||
from litellm.types.mcp import MCPAuth
|
||||
import pytest
|
||||
|
||||
|
|
@ -600,6 +598,8 @@ async def test_per_user_oauth_missing_stored_token_returns_preemptive_401():
|
|||
Per-user OAuth server with no stored token should fail fast with 401 +
|
||||
WWW-Authenticate so PKCE can start.
|
||||
"""
|
||||
from fastapi import HTTPException
|
||||
|
||||
try:
|
||||
from litellm.proxy._experimental.mcp_server.server import (
|
||||
handle_streamable_http_mcp,
|
||||
|
|
@ -612,8 +612,13 @@ async def test_per_user_oauth_missing_stored_token_returns_preemptive_401():
|
|||
"type": "http",
|
||||
"method": "POST",
|
||||
"path": "/mcp",
|
||||
"scheme": "http",
|
||||
"query_string": b"",
|
||||
"root_path": "",
|
||||
"server": ("localhost", 8000),
|
||||
"headers": [
|
||||
(b"content-type", b"application/json"),
|
||||
(b"host", b"localhost:8000"),
|
||||
],
|
||||
}
|
||||
receive = AsyncMock()
|
||||
|
|
@ -660,11 +665,12 @@ async def test_per_user_oauth_missing_stored_token_returns_preemptive_401():
|
|||
with pytest.raises(HTTPException) as exc_info:
|
||||
await handle_streamable_http_mcp(scope, receive, send)
|
||||
|
||||
exc = exc_info.value
|
||||
assert exc.status_code == 401
|
||||
assert "www-authenticate" in exc.headers
|
||||
# Verify a 401 was raised
|
||||
assert mock_get_stored_token.await_count == 1
|
||||
assert mock_handle_request.await_count == 0
|
||||
assert exc_info.value.status_code == 401
|
||||
assert "www-authenticate" in exc_info.value.headers
|
||||
assert "Bearer authorization_uri=" in exc_info.value.headers["www-authenticate"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -685,11 +691,22 @@ async def test_per_user_oauth_with_stored_token_skips_preemptive_401():
|
|||
"type": "http",
|
||||
"method": "POST",
|
||||
"path": "/mcp",
|
||||
"scheme": "http",
|
||||
"query_string": b"",
|
||||
"root_path": "",
|
||||
"server": ("localhost", 8000),
|
||||
"headers": [
|
||||
(b"content-type", b"application/json"),
|
||||
(b"host", b"localhost:8000"),
|
||||
],
|
||||
}
|
||||
receive = AsyncMock()
|
||||
receive = AsyncMock(
|
||||
return_value={
|
||||
"type": "http.request",
|
||||
"body": b'{"jsonrpc":"2.0","id":1,"method":"tools/list","params":{}}',
|
||||
"more_body": False,
|
||||
}
|
||||
)
|
||||
send = AsyncMock()
|
||||
user_auth = MagicMock()
|
||||
user_auth.user_id = "test-user-id"
|
||||
|
|
@ -729,6 +746,11 @@ async def test_per_user_oauth_with_stored_token_skips_preemptive_401():
|
|||
"handle_request",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_handle_request,
|
||||
patch.object(
|
||||
session_manager_stateless,
|
||||
"_server_instances",
|
||||
{},
|
||||
),
|
||||
):
|
||||
await handle_streamable_http_mcp(scope, receive, send)
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue