feat(mcp): configure protocol versions and capability discovery (#43169)

* feat(mcp): configure protocol versions and capability discovery

* test(mcp): return SDK initialization result in REST pagination fixture

* test(mcp): arm cancellation deadlines after TCP calls start

* fix(mcp): avoid serialized discovery and listing spend logs

* test(mcp): isolate default protocol header policy

* fix(mcp): honor preview protocol pins and refresh migrated tests

* fix(mcp): preserve edited protocol pins in saved OAuth previews

* fix(mcp): retain saved protocol pins when previews omit versions

---------

Co-authored-by: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com>
This commit is contained in:
joshua-berri 2026-09-25 20:13:52 +00:00 • committed by GitHub
parent cf491d1df9
commit 6b7688869e
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
25 changed files with 939 additions and 42 deletions

View file

@ -10,6 +10,7 @@ import os
from collections.abc import AsyncIterator, Awaitable, Callable, Generator, Sequence
from contextlib import AbstractAsyncContextManager
from functools import partial
from importlib.metadata import version
from types import MappingProxyType
from typing import Final, TypeAlias, TypeVar, cast
@ -34,8 +35,16 @@ _TransportContext: TypeAlias = AbstractAsyncContextManager[_TransportStreams]
from mcp.types import (
METHOD_NOT_FOUND,
REQUEST_TIMEOUT,
ClientCapabilities,
ElicitationCapability,
FormElicitationCapability,
GetPromptRequestParams,
GetPromptResult,
Implementation,
InitializedNotification,
InitializeRequest,
InitializeRequestParams,
InitializeResult,
InputRequiredResult,
ListPromptsResult,
ListResourcesResult,
@ -44,12 +53,14 @@ from mcp.types import (
PaginatedResult,
Prompt,
ResourceTemplate,
SamplingCapability,
ServerNotification,
UrlElicitationCapability,
)
from mcp.types import CallToolRequestParams as MCPCallToolRequestParams
from mcp.types import CallToolResult as MCPCallToolResult
from mcp.types import Tool as MCPTool
from pydantic import AnyUrl
from pydantic import AnyUrl, TypeAdapter
from litellm._logging import verbose_logger
from litellm.constants import (
@ -64,11 +75,13 @@ from litellm.proxy._experimental.mcp_server.mcp_debug import capture_upstream_er
from litellm.proxy._experimental.mcp_server.result_conversion import error_text_result
from litellm.types.llms.custom_http import VerifyTypes
from litellm.types.mcp import (
MCP_LEGACY_VERSIONS,
MCPAuth,
MCPAuthType,
MCPStdioConfig,
MCPTransport,
MCPTransportType,
MCPUpstreamProtocol,
credential_redirect_hook,
has_header,
without_header,
@ -386,7 +399,9 @@ class MCPClient:
sampling_callback: Callable | None = None,
elicitation_callback: Callable | None = None,
logging_callback: Callable | None = None,
protocol_version: MCPUpstreamProtocol = "auto",
):
self.protocol_version: MCPUpstreamProtocol = TypeAdapter(MCPUpstreamProtocol).validate_python(protocol_version)
self.server_url: str = server_url
self.transport_type: MCPTransport = transport_type
self.auth_type: MCPAuthType = auth_type
@ -525,6 +540,35 @@ class MCPClient:
return safe_env
async def _initialize_session(self, session: ClientSession) -> InitializeResult:
if self.protocol_version == "auto":
automatic: Final = await session.initialize()
if automatic.protocol_version not in MCP_LEGACY_VERSIONS:
raise MCPError(code=-32022, message="Upstream selected an unsupported MCP protocol version")
return automatic
result: Final = await session.send_request(
InitializeRequest(
params=InitializeRequestParams(
protocol_version=self.protocol_version,
client_info=Implementation(name="litellm", version=version("litellm")),
capabilities=ClientCapabilities(
sampling=SamplingCapability() if self._sampling_callback is not None else None,
elicitation=ElicitationCapability(
form=FormElicitationCapability(), url=UrlElicitationCapability()
)
if self._elicitation_callback is not None
else None,
),
)
),
InitializeResult,
)
if result.protocol_version != self.protocol_version:
raise MCPError(code=-32022, message="Upstream did not accept the configured MCP protocol version")
session.adopt(result)
await session.send_notification(InitializedNotification())
return result
async def _execute_session_operation(
self,
transport_ctx: _TransportContext,
@ -579,7 +623,7 @@ class MCPClient:
)
session: Final = await session_ctx.__aenter__()
try:
init_result: Final = await session.initialize()
init_result: Final = await self._initialize_session(session)
instructions: Final = getattr(init_result, "instructions", None)
self._last_initialize_instructions = (
instructions.strip() or None if isinstance(instructions, str) else None

View file

@ -0,0 +1,145 @@
from collections.abc import Callable, Mapping
from dataclasses import dataclass
from itertools import product
from types import MappingProxyType
from typing import Final, Literal
from mcp.server.context import CallNext, HandlerResult, ServerRequestContext
from mcp.shared.exceptions import MCPError
from mcp.types import DiscoverResult, InitializeRequestParams, InitializeResult, ServerCapabilities
from mcp_types.methods import CLIENT_REQUESTS
from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS, LATEST_HANDSHAKE_VERSION
from pydantic import TypeAdapter
from litellm.types.mcp import MCP_LEGACY_VERSIONS, MCPAdvertisedVersions, MCPLegacyVersion, MCPSpecVersion, MCPTransport
GATEWAY_OPERATIONS: Final = frozenset(
{
"tools/list",
"tools/call",
"prompts/list",
"prompts/get",
"resources/list",
"resources/read",
"resources/templates/list",
}
)
@dataclass(frozen=True, slots=True)
class RevisionSupport:
transports: frozenset[MCPTransport]
operations: frozenset[str]
results: frozenset[Literal["complete", "input_required"]]
extensions: frozenset[str]
completed: bool
REVISION_SUPPORT: Final[Mapping[str, RevisionSupport]] = MappingProxyType(
{
version.value: RevisionSupport(
transports=frozenset(MCPTransport)
if version.value in HANDSHAKE_PROTOCOL_VERSIONS
else frozenset({MCPTransport.http, MCPTransport.stdio}),
operations=frozenset(method for method in GATEWAY_OPERATIONS if (method, version.value) in CLIENT_REQUESTS),
results=frozenset({"complete"})
if version.value in HANDSHAKE_PROTOCOL_VERSIONS
else frozenset({"complete", "input_required"}),
extensions=frozenset(),
completed=version.value in HANDSHAKE_PROTOCOL_VERSIONS,
)
for version in MCPSpecVersion
}
)
_COMPLETED_REVISIONS: Final = tuple(version for version, support in REVISION_SUPPORT.items() if support.completed)
TRANSLATION_PAIRS: Final = frozenset(product(_COMPLETED_REVISIONS, repeat=2))
_ADVERTISED_VERSIONS: Final[TypeAdapter[tuple[MCPLegacyVersion, ...]]] = TypeAdapter(MCPAdvertisedVersions)
def configured_versions() -> tuple[str, ...]:
from litellm.proxy.proxy_server import general_settings_view
configured: Final = general_settings_view().get("mcp_advertised_versions")
return _ADVERTISED_VERSIONS.validate_python(MCP_LEGACY_VERSIONS if configured is None else configured)
def build_discovery(
*,
configured: tuple[str, ...],
revision: str,
transport: MCPTransport,
authorized_operations: frozenset[str],
upstream_versions: frozenset[str],
capabilities: ServerCapabilities,
client_extensions: frozenset[str] = frozenset(),
upstream_extensions: frozenset[str] = frozenset(),
instructions: str | None = None,
) -> DiscoverResult:
supported: Final = tuple(
version
for version, support in REVISION_SUPPORT.items()
if version in configured and support.completed and transport in support.transports
)
revision_support: Final = REVISION_SUPPORT.get(revision)
operations: Final[frozenset[str]] = (
authorized_operations & revision_support.operations
if revision in supported
and revision_support is not None
and any((revision, upstream) in TRANSLATION_PAIRS for upstream in upstream_versions)
else frozenset()
)
extensions: Final[frozenset[str]] = (
revision_support.extensions & client_extensions & upstream_extensions
if operations and revision_support is not None
else frozenset()
)
caller_capabilities: Final = capabilities.model_copy(deep=True)
return DiscoverResult(
supported_versions=list(supported),
capabilities=ServerCapabilities(
tools=caller_capabilities.tools if {"tools/list", "tools/call"} <= operations else None,
prompts=caller_capabilities.prompts if {"prompts/list", "prompts/get"} <= operations else None,
resources=caller_capabilities.resources if {"resources/list", "resources/read"} <= operations else None,
extensions={
key: value for key, value in (caller_capabilities.extensions or {}).items() if key in extensions
}
or None,
),
instructions=instructions,
cache_scope="private",
ttl_ms=0,
)
class GatewayVersionPolicy:
def __init__(self, versions: Callable[[], tuple[str, ...]] = configured_versions) -> None:
self._versions = versions
async def __call__(self, ctx: ServerRequestContext[object, object], call_next: CallNext) -> HandlerResult:
versions: Final = self._versions()
requested: Final = (
InitializeRequestParams.model_validate(ctx.params or {}).protocol_version
if ctx.method == "initialize"
else ctx.protocol_version
)
negotiated: Final = (
(requested if requested in HANDSHAKE_PROTOCOL_VERSIONS else LATEST_HANDSHAKE_VERSION)
if ctx.method == "initialize"
else requested
)
if negotiated not in versions:
raise MCPError(code=-32022, message="Unsupported MCP protocol version", data={"supported": list(versions)})
result: Final = await call_next(ctx)
if ctx.method != "initialize":
return result
initialized: Final = InitializeResult.model_validate(result)
discovery: Final = build_discovery(
configured=versions,
revision=initialized.protocol_version,
transport=MCPTransport.http,
authorized_operations=GATEWAY_OPERATIONS,
upstream_versions=frozenset(HANDSHAKE_PROTOCOL_VERSIONS),
capabilities=initialized.capabilities,
instructions=initialized.instructions,
)
return initialized.model_copy(update={"capabilities": discovery.capabilities})

View file

@ -28,6 +28,7 @@ class OperationContext:
client_ip: str | None = None
mcp_proxy_mode: bool = False
wire_compat: WireCompat = WireCompat.LEGACY
protocol_version: str | None = None
def __post_init__(self) -> None:
object.__setattr__(self, "_caller", copy_caller(self._caller))

View file

@ -193,6 +193,7 @@ from litellm.types.mcp import (
MCPAuth,
MCPStdioConfig,
MCPTokenEndpointAuthMethod,
MCPUpstreamProtocol,
has_header,
without_header,
)
@ -340,6 +341,7 @@ class MCPServerConfig(TypedDict, total=False):
whatever the admin wrote, and each read applies its own default."""
server_id: ReadOnly[str]
protocol_version: ReadOnly[MCPUpstreamProtocol]
alias: str
description: str
mcp_info: MCPInfo
@ -2549,6 +2551,9 @@ class MCPServerManager:
new_server = MCPServer(
server_id=server_id,
name=name_for_prefix,
protocol_version=TypeAdapter(MCPUpstreamProtocol).validate_python(
server_config.get("protocol_version", mcp_info.get("protocol_version", "auto"))
),
alias=alias,
server_name=server_name,
spec_path=server_config.get("spec_path", None),
@ -3109,6 +3114,9 @@ class MCPServerManager:
new_server: Final = MCPServer(
server_id=mcp_server.server_id,
name=name_for_prefix,
protocol_version=TypeAdapter(MCPUpstreamProtocol).validate_python(
_mcp_info.get("protocol_version", "auto")
),
alias=getattr(mcp_server, "alias", None),
server_name=getattr(mcp_server, "server_name", None),
url=mcp_server.url,
@ -4145,6 +4153,7 @@ class MCPServerManager:
cred_provider: UpstreamCredentialProvider | None = None,
raw_headers: Mapping[str, str] | None = None,
client_ip: str | None = None,
protocol_version_override: MCPUpstreamProtocol | None = None,
) -> MCPClient:
"""
Create an MCPClient instance for the given server.
@ -4168,6 +4177,9 @@ class MCPServerManager:
"""
record_auth_resolution(server.server_id, AuthResolution.unresolved)
resolved_server: Final = await self.ensure_oauth_metadata_discovered(server)
protocol_version: Final = (
protocol_version_override if protocol_version_override is not None else resolved_server.protocol_version
)
transport: Final = resolved_server.transport or MCPTransport.sse
spec = None if transport == MCPTransport.stdio else _to_server_spec_fail_closed(resolved_server)
provider: Final = cred_provider or self._cred_provider
@ -4249,6 +4261,7 @@ class MCPServerManager:
return MCPClient(
server_url="", # Not used for stdio
transport_type=transport,
protocol_version=protocol_version,
auth_type=resolved_server.auth_type,
auth_value=auth_value,
timeout=(resolved_server.timeout if resolved_server.timeout is not None else MCP_CLIENT_TIMEOUT),
@ -4281,6 +4294,7 @@ class MCPServerManager:
MCPClient(
server_url=server_url,
transport_type=transport,
protocol_version=protocol_version,
auth_type=resolved_server.auth_type,
timeout=(
resolved_server.timeout if resolved_server.timeout is not None else MCP_CLIENT_TIMEOUT
@ -4324,6 +4338,7 @@ class MCPServerManager:
MCPClient(
server_url=server_url,
transport_type=transport,
protocol_version=protocol_version,
auth_type=resolved_server.auth_type,
auth_value=auth_value,
auth_header_name=auth_header_name,

View file

@ -14,6 +14,8 @@ from mcp.types import (
CallToolRequest,
CallToolRequestParams,
CallToolResult,
DiscoverRequest,
DiscoverResult,
GetPromptRequest,
GetPromptRequestParams,
GetPromptResult,
@ -28,10 +30,14 @@ from mcp.types import (
ListToolsResult,
PaginatedRequestParams,
Prompt,
PromptsCapability,
ReadResourceRequest,
ReadResourceRequestParams,
ResourcesCapability,
ResourceTemplate,
ServerCapabilities,
TextContent,
ToolsCapability,
)
from mcp.types import Tool as MCPTool
from pydantic import AnyUrl, ConfigDict, Field, TypeAdapter
@ -51,6 +57,11 @@ from litellm.proxy._experimental.mcp_server.byok_credential_cache import (
cache_byok_credential,
get_cached_byok_credential,
)
from litellm.proxy._experimental.mcp_server.capabilities import (
GATEWAY_OPERATIONS,
build_discovery,
configured_versions,
)
from litellm.proxy._experimental.mcp_server.contracts import (
AuthorizedToolCall,
OperationContext,
@ -122,7 +133,9 @@ from litellm.proxy.litellm_pre_call_utils import (
)
from litellm.types.mcp import (
DEFAULT_CREDENTIAL_HEADER,
MCP_LEGACY_VERSIONS,
MCPAuth,
MCPTransport,
without_header,
)
from litellm.types.mcp_server.mcp_server_manager import MCPInfo, MCPServer
@ -2657,7 +2670,11 @@ class _McpDeniedDetail(TypedDict):
async def _execute_handle_list_tools(
context: OperationContext, params: PaginatedRequestParams, host_progress_callback: ProgressCallback | None = None
context: OperationContext,
params: PaginatedRequestParams,
host_progress_callback: ProgressCallback | None = None,
*,
log_list_tools_to_spendlogs: bool = True,
) -> ListToolsResult:
try:
(
@ -2700,7 +2717,7 @@ async def _execute_handle_list_tools(
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
log_list_tools_to_spendlogs=True,
log_list_tools_to_spendlogs=log_list_tools_to_spendlogs,
list_tools_log_source="mcp_protocol",
client_ip=_client_ip,
)
@ -3065,6 +3082,7 @@ def prepare_context(
client_ip: str | None = None,
mcp_proxy_mode: bool = False,
wire_compat: WireCompat = WireCompat.LEGACY,
protocol_version: str | None = None,
) -> OperationContext:
return OperationContext(
_caller=user_api_key_auth,
@ -3076,11 +3094,13 @@ def prepare_context(
client_ip=client_ip,
mcp_proxy_mode=mcp_proxy_mode,
wire_compat=wire_compat,
protocol_version=protocol_version,
)
GatewayOperation: TypeAlias = (
AuthorizedToolCall
| DiscoverRequest
| ListToolsRequest
| CallToolRequest
| ListPromptsRequest
@ -3090,7 +3110,8 @@ GatewayOperation: TypeAlias = (
| ReadResourceRequest
)
GatewayResult: TypeAlias = (
ListToolsResult
DiscoverResult
| ListToolsResult
| CallToolResult
| InputRequiredResult
| ListPromptsResult
@ -3105,6 +3126,9 @@ class GatewayOperations:
def __init__(self, host_progress_callback: ProgressCallback | None = None) -> None:
self._host_progress_callback = host_progress_callback
@overload
async def execute(self, operation: DiscoverRequest, context: OperationContext) -> DiscoverResult: ...
@overload
async def execute(
self, operation: AuthorizedToolCall, context: OperationContext
@ -3137,6 +3161,51 @@ class GatewayOperations:
async def execute(self, operation: GatewayOperation, context: OperationContext) -> GatewayResult:
match operation:
case DiscoverRequest():
listings: Final = (
()
if context.mcp_proxy_mode
else (ListPromptsRequest(), ListResourcesRequest(), ListResourceTemplatesRequest())
)
tasks: Final = (
asyncio.create_task(
_execute_handle_list_tools(
context,
PaginatedRequestParams(),
self._host_progress_callback,
log_list_tools_to_spendlogs=False,
)
),
*(asyncio.create_task(self.execute(listing, context)) for listing in listings),
)
try:
results: Final = await asyncio.gather(*tasks)
finally:
for task in tasks:
task.cancel()
await asyncio.gather(*tasks, return_exceptions=True)
return build_discovery(
configured=configured_versions(),
revision=context.protocol_version or "2025-11-25",
transport=MCPTransport.http,
authorized_operations=GATEWAY_OPERATIONS,
upstream_versions=frozenset(MCP_LEGACY_VERSIONS),
capabilities=ServerCapabilities(
tools=ToolsCapability()
if any(isinstance(result, ListToolsResult) and result.tools for result in results)
else None,
prompts=PromptsCapability()
if any(isinstance(result, ListPromptsResult) and result.prompts for result in results)
else None,
resources=ResourcesCapability()
if any(
(isinstance(result, ListResourcesResult) and result.resources)
or (isinstance(result, ListResourceTemplatesResult) and result.resource_templates)
for result in results
)
else None,
),
)
case AuthorizedToolCall():
auth, token, _servers, server_headers, oauth_headers, headers, _client_ip = context.legacy_auth()
return await _execute_mcp_tool(

View file

@ -1375,7 +1375,16 @@ if MCP_AVAILABLE:
and headers.get(MCPRequestHandler.LITELLM_API_KEY_HEADER_NAME_PRIMARY)
else None
)
return _StagedServerTest(request=request, mcp_auth_header=mcp_auth_header, oauth2_headers=oauth2_headers)
preview_request: Final = (
request.model_copy(
update={"mcp_info": {**(request.mcp_info or {}), "protocol_version": saved_server.protocol_version}}
)
if saved_server is not None and "protocol_version" not in (request.mcp_info or {})
else request
)
return _StagedServerTest(
request=preview_request, mcp_auth_header=mcp_auth_header, oauth2_headers=oauth2_headers
)
async def _list_tools_within(client: MCPClient, deadline: float) -> list[MCPTool] | None:
with anyio.move_on_after(deadline):
@ -1512,6 +1521,7 @@ if MCP_AVAILABLE:
extra_headers=merged_headers,
stdio_env=stdio_env,
cred_provider=preview_cred_provider,
protocol_version_override=server_model.protocol_version,
)
return await operation(client)

View file

@ -125,12 +125,14 @@ def unsupported_protocol_version(scope: Scope) -> str | None:
``HANDSHAKE_PROTOCOL_VERSIONS`` to the modern single-exchange path, which
bypasses litellm's session/auth model, so the ASGI entry rejects it.
"""
from litellm.proxy._experimental.mcp_server.capabilities import configured_versions
headers: Final[Iterable[tuple[bytes, bytes]]] = scope.get("headers") or ()
values: Final = tuple(
raw.decode("latin-1").strip() for key, raw in headers if key.lower() == _MCP_PROTOCOL_VERSION_HEADER
)
for value in values:
if value and value not in HANDSHAKE_PROTOCOL_VERSIONS:
if value and value not in configured_versions():
return value
return None
@ -149,7 +151,10 @@ try:
from mcp.server.session import ServerSession as _McpServerSession
from mcp.types import (
BlobResourceContents,
DiscoverRequest,
DiscoverResult,
GetPromptResult,
RequestParams,
ResourceTemplate,
TextResourceContents,
)
@ -526,11 +531,11 @@ if MCP_AVAILABLE:
PaginatedRequestParams,
ReadResourceRequestParams,
)
from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS
from litellm.proxy._experimental.mcp_server.auth.litellm_auth_handler import (
MCPAuthenticatedUser,
)
from litellm.proxy._experimental.mcp_server.capabilities import GatewayVersionPolicy, configured_versions
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
MCPServerManager,
global_mcp_server_manager,
@ -585,6 +590,7 @@ if MCP_AVAILABLE:
name=LITELLM_MCP_SERVER_NAME,
version=LITELLM_MCP_SERVER_VERSION,
)
server.middleware.append(GatewayVersionPolicy())
server.create_initialization_options = types.MethodType(_gateway_create_initialization_options, server)
sse: Final[SseServerTransport] = SseServerTransport("/sse/messages")
@ -830,6 +836,7 @@ if MCP_AVAILABLE:
client_ip,
_mcp_proxy_mode.get(),
wire_compat_for(ctx.protocol_version),
ctx.protocol_version,
)
async def handle_list_tools(ctx: ServerRequestContext, params: PaginatedRequestParams) -> ListToolsResult:
@ -948,6 +955,11 @@ if MCP_AVAILABLE:
ReadResourceRequest(params=params), context
)
async def discover(ctx: ServerRequestContext, params: RequestParams) -> DiscoverResult:
async with _legacy_operation_context(ctx, trace=False) as context:
return await operations.GatewayOperations().execute(DiscoverRequest(params=params), context)
server.add_request_handler("server/discover", RequestParams, discover)
server.add_request_handler("tools/list", PaginatedRequestParams, handle_list_tools)
server.add_request_handler("tools/call", CallToolRequestParams, mcp_server_tool_call)
server.add_request_handler("prompts/list", PaginatedRequestParams, list_prompts)
@ -1954,7 +1966,7 @@ if MCP_AVAILABLE:
reject_disallowed_mcp_origin(StarletteRequest(scope))
bad_version: Final = unsupported_protocol_version(scope)
if bad_version is not None:
supported: Final = ", ".join(sorted(HANDSHAKE_PROTOCOL_VERSIONS))
supported: Final = ", ".join(configured_versions())
await JSONResponse(
status_code=400,
content={ # mutable-ok: JSON-RPC error payload
@ -2299,7 +2311,7 @@ if MCP_AVAILABLE:
reject_disallowed_mcp_origin(StarletteRequest(scope))
bad_version: Final = unsupported_protocol_version(scope)
if bad_version is not None:
supported: Final = ", ".join(sorted(HANDSHAKE_PROTOCOL_VERSIONS))
supported: Final = ", ".join(configured_versions())
await JSONResponse(
status_code=400,
content={ # mutable-ok: JSON-RPC error payload

View file

@ -37,6 +37,7 @@ from litellm.types.llms.openai import (
ResponsesAPIResponse,
)
from litellm.types.mcp import (
MCPAdvertisedVersions,
MCPAllowedClient,
MCPAuth,
MCPAuthType,
@ -2998,6 +2999,11 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
None,
description="Custom CIDR ranges that define internal/private networks for MCP access control. When set, only these ranges are treated as internal. Defaults to RFC 1918 private ranges (10.0.0.0/8, 172.16.0.0/12, 192.168.0.0/16, 127.0.0.0/8).",
)
mcp_advertised_versions: MCPAdvertisedVersions | None = Field(
None,
description="MCP revisions enabled by the gateway. Defaults to all completed legacy revisions. "
"Modern protocol serving and Apps/Tasks remain disabled.",
)
mcp_allowed_clients: list[MCPAllowedClient] | None = Field(
None,
description="MCP client applications admitted by the gateway, each an {alias, value} pair where alias is the name shown in the dashboard and logs and value is the identity that must match exactly. When set, every MCP request must carry a client identity equal to one of the values: a JWT caller is identified by the claim named in litellm_jwtauth.mcp_client_id_jwt_field, any other caller by the header named in mcp_client_id_header. A request with no resolvable identity, or an unlisted one, is rejected with 403. Unset means every client is admitted.",

View file

@ -6342,6 +6342,11 @@ class ProxyConfig:
if general_settings is None:
general_settings = {}
if general_settings.get("mcp_advertised_versions") is not None:
from litellm.types.mcp import MCPAdvertisedVersions
TypeAdapter(MCPAdvertisedVersions).validate_python(general_settings["mcp_advertised_versions"])
if os.getenv("NUM_WORKERS", "1") != "1" and redis_usage_cache is None:
warn_login_counters_are_per_worker(os.getenv("NUM_WORKERS", "1"))
if declared_proxy_ranges(general_settings) is None:

View file

@ -4,7 +4,7 @@ import enum
import re
from collections.abc import Awaitable, Callable, Mapping
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal
from typing import TYPE_CHECKING, Annotated, Any, Final, Literal
from urllib.parse import urlsplit
import httpx
@ -34,6 +34,8 @@ class MCPSpecVersion(str, enum.Enum):
nov_2024 = "2024-11-05"
mar_2025 = "2025-03-26"
jun_2025 = "2025-06-18"
nov_2025 = "2025-11-25"
jul_2026 = "2026-07-28"
class MCPAuth(str, enum.Enum):
@ -59,7 +61,17 @@ DEFAULT_SUBJECT_TOKEN_TYPE: Final = "urn:ietf:params:oauth:token-type:access_tok
# MCP Literals
MCPTransportType = Literal[MCPTransport.sse, MCPTransport.http, MCPTransport.stdio]
MCPSpecVersionType = Literal[MCPSpecVersion.nov_2024, MCPSpecVersion.mar_2025, MCPSpecVersion.jun_2025]
MCPLegacyVersion = Literal["2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25"]
MCP_LEGACY_VERSIONS: Final[tuple[MCPLegacyVersion, ...]] = ("2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25")
MCPUpstreamProtocol = MCPLegacyVersion | Literal["auto"]
MCPAdvertisedVersions = Annotated[tuple[MCPLegacyVersion, ...], Field(min_length=1)]
MCPSpecVersionType = Literal[
MCPSpecVersion.nov_2024,
MCPSpecVersion.mar_2025,
MCPSpecVersion.jun_2025,
MCPSpecVersion.nov_2025,
MCPSpecVersion.jul_2026,
]
MCPAuthType = (
Literal[
MCPAuth.none,

View file

@ -1,7 +1,7 @@
from datetime import datetime
from typing import Any, Final, Literal
from typing import Annotated, Any, Final, Literal
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
from pydantic import AfterValidator, BaseModel, ConfigDict, Field, TypeAdapter, field_validator, model_validator
from typing_extensions import Self
from litellm.types.mcp import (
@ -10,11 +10,19 @@ from litellm.types.mcp import (
MCPAuthType,
MCPTokenEndpointAuthMethod,
MCPTransportType,
MCPUpstreamProtocol,
normalize_upstream_header_name,
)
# MCPInfo now allows arbitrary additional fields for custom metadata
MCPInfo = dict[str, Any]
def _validate_mcp_protocol_metadata(value: dict[str, object]) -> dict[str, object]:
if "protocol_version" in value:
TypeAdapter(MCPUpstreamProtocol).validate_python(value["protocol_version"])
return value
MCPInfo = Annotated[dict[str, Any], AfterValidator(_validate_mcp_protocol_metadata)]
class MCPOAuthMetadata(BaseModel):
@ -66,6 +74,7 @@ class MCPServer(BaseModel):
server_name: str | None = None
url: str | None = None
transport: MCPTransportType
protocol_version: MCPUpstreamProtocol = "auto"
spec_path: str | None = None
auth_type: MCPAuthType | None = None
authentication_token: str | None = None
@ -246,6 +255,14 @@ class MCPServer(BaseModel):
"""
return self.oauth2_flow == "client_credentials"
@model_validator(mode="after")
def resolve_protocol_version(self) -> Self:
if "protocol_version" not in self.model_fields_set and self.mcp_info is not None:
self.protocol_version = TypeAdapter(MCPUpstreamProtocol).validate_python(
self.mcp_info.get("protocol_version", "auto")
)
return self
@model_validator(mode="after")
def validate_identity_binding_mode(self) -> Self:
binding: Final = self.oauth_identity_binding

View file

@ -84,3 +84,42 @@ def test_jsonrpc_error_and_malformed_tool_result_remain_errors(gateway: Gateway)
control: Final = call_tool(gateway, key, identity, names["add"], {"a": 3, "b": 5})
assert control.status_code == 200 and control.json()["isError"] is False, control.text
assert control.json()["content"][0]["text"] == "8"
@pytest.mark.parametrize("ingress", ("http", "sse"))
def test_configured_revision_blocks_unadvertised_handshake_and_keeps_allowed_control(
gateway: Gateway, tmp_path, ingress: str
) -> None:
import asyncio
from pathlib import Path
import yaml
from integration._support.mcp import mcp_peer
from integration._support.process import owned_proxy
from litellm.experimental_mcp_client.client import MCPClient
from litellm.types.mcp import MCPTransport
from mcp import MCPError
from mcp.types import CallToolRequestParams
with mcp_peer() as upstream, gateway.scenario() as scenario:
alias: Final = "restricted" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(scenario, upstream, alias)
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["general_settings"]["mcp_advertised_versions"] = ["2024-11-05"]
config_path: Final = tmp_path / "restricted.yaml"
config_path.write_text(yaml.safe_dump(config))
with owned_proxy(gateway, tmp_path, {"DISABLE_SCHEMA_UPDATE": "true"}, config=config_path) as restricted:
endpoint: Final = str(restricted.client.base_url).rstrip("/") + ("/mcp/sse" if ingress == "sse" else "/mcp")
headers: Final = {"Authorization": f"Bearer {key}", "x-mcp-servers": identity}
denied: Final = MCPClient(server_url=endpoint, transport_type=MCPTransport(ingress), protocol_version="2025-11-25", extra_headers=headers)
allowed: Final = MCPClient(server_url=endpoint, transport_type=MCPTransport(ingress), protocol_version="2024-11-05", extra_headers=headers)
async def exercise() -> None:
with pytest.raises(MCPError, match="Unsupported MCP protocol version"):
await denied.list_tools(raise_on_error=True)
assert f"{alias}-add" in tuple(tool.name for tool in await allowed.list_tools(raise_on_error=True))
result: Final = await allowed.call_tool(CallToolRequestParams(name=f"{alias}-add", arguments={"a": 2, "b": 5}))
assert result.is_error is False and result.content[0].text == "7"
asyncio.run(exercise())

View file

@ -155,3 +155,43 @@ def test_server_initiated_sampling_and_elicitation_surface_as_errors_not_success
assert outcome.error is not None, outcome.raw
assert outcome.text is None or not outcome.text.startswith(("sampled:", "elicited:")), outcome.raw
assert len(tool_calls(peer.drain())) == 1
@pytest.mark.parametrize("downstream", ("2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25"))
@pytest.mark.parametrize("upstream", ("2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25"))
@pytest.mark.parametrize("peer_kind", ("http", "sse", "stdio"))
@pytest.mark.parametrize("ingress", ("http", "sse"))
def test_pinned_revision_pairs_list_and_call_through_gateway(
gateway: Gateway, downstream: str, upstream: str, peer_kind: PeerKind, ingress: str
) -> None:
import asyncio
from mcp.types import CallToolRequestParams
from litellm.experimental_mcp_client.client import MCPClient
from litellm.types.mcp import MCPTransport
with peer_of(peer_kind) as peer, gateway.scenario() as scenario:
alias: Final = "versions" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(scenario, peer, alias, mcp_info={"protocol_version": upstream})
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
endpoint: Final = str(gateway.client.base_url).rstrip("/") + ("/mcp/sse" if ingress == "sse" else "/mcp")
client: Final = MCPClient(
server_url=endpoint, transport_type=MCPTransport(ingress), protocol_version=downstream,
extra_headers={"Authorization": f"Bearer {key}", "x-mcp-servers": identity}, timeout=15,
)
async def exercise() -> None:
tools: Final = await client.list_tools(raise_on_error=True)
assert f"{alias}-add" in tuple(tool.name for tool in tools)
result: Final = await client.call_tool(CallToolRequestParams(name=f"{alias}-add", arguments={"a": 3, "b": 4}))
assert result.is_error is False
assert result.content[0].text == "7"
peer.drain()
asyncio.run(exercise())
observed: Final = peer.drain()
negotiations: Final = tuple(item["body"] for item in observed if item["body"].get("method") == "initialize")
assert negotiations, "The operation must reach the upstream negotiation"
assert all(request["params"]["protocolVersion"] == upstream for request in negotiations), negotiations
assert len(tool_calls(observed)) == 1

View file

@ -0,0 +1,106 @@
from typing import Final
import pytest
from mcp import Client
from mcp.server import Server
from mcp.shared.exceptions import MCPError
from mcp.types import PromptsCapability, ResourcesCapability, ServerCapabilities, ToolsCapability
from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS
from litellm.proxy._experimental.mcp_server.capabilities import (
GATEWAY_OPERATIONS,
REVISION_SUPPORT,
TRANSLATION_PAIRS,
GatewayVersionPolicy,
build_discovery,
)
from litellm.types.mcp import MCPTransport
@pytest.mark.parametrize("revision", HANDSHAKE_PROTOCOL_VERSIONS)
@pytest.mark.parametrize("transport", tuple(MCPTransport))
def test_discovery_only_exposes_authorized_completed_support(revision, transport):
result = build_discovery(
configured=(revision, "2026-07-28", "unknown"),
revision=revision,
transport=transport,
authorized_operations=frozenset({"tools/list", "tools/call"}),
upstream_versions=frozenset(HANDSHAKE_PROTOCOL_VERSIONS),
capabilities=ServerCapabilities(
tools=ToolsCapability(), prompts=PromptsCapability(), resources=ResourcesCapability(),
extensions={"io.modelcontextprotocol/ui": {}},
),
client_extensions=frozenset({"io.modelcontextprotocol/ui"}),
upstream_extensions=frozenset({"io.modelcontextprotocol/ui"}),
)
assert result.supported_versions == [revision]
assert result.capabilities.tools is not None
assert result.capabilities.prompts is None
assert result.capabilities.resources is None
assert result.capabilities.extensions is None
assert result.capabilities.tasks is None
assert result.cache_scope == "private"
assert result.ttl_ms == 0
@pytest.mark.parametrize("upstream", [frozenset(), frozenset({"unknown"}), frozenset({"2026-07-28"})])
def test_unproven_translation_never_advertises_operations(upstream):
result = build_discovery(
configured=HANDSHAKE_PROTOCOL_VERSIONS,
revision="2025-11-25",
transport=MCPTransport.http,
authorized_operations=GATEWAY_OPERATIONS,
upstream_versions=upstream,
capabilities=ServerCapabilities(tools=ToolsCapability()),
)
assert result.capabilities.tools is None
@pytest.mark.parametrize("revision", ["2026-07-28", "unknown", "2024-11-05"])
def test_unadvertised_revision_never_gains_capabilities(revision):
result = build_discovery(
configured=("2025-11-25",), revision=revision, transport=MCPTransport.http,
authorized_operations=GATEWAY_OPERATIONS, upstream_versions=frozenset(HANDSHAKE_PROTOCOL_VERSIONS),
capabilities=ServerCapabilities(tools=ToolsCapability()),
)
assert result.capabilities.tools is None
def test_discovery_results_do_not_share_mutable_capabilities():
capabilities = ServerCapabilities(tools=ToolsCapability(), prompts=PromptsCapability(), resources=ResourcesCapability())
args = dict(
configured=HANDSHAKE_PROTOCOL_VERSIONS, revision="2025-11-25", transport=MCPTransport.http,
upstream_versions=frozenset(HANDSHAKE_PROTOCOL_VERSIONS), capabilities=capabilities,
)
allowed = build_discovery(**args, authorized_operations=GATEWAY_OPERATIONS)
denied = build_discovery(**args, authorized_operations=frozenset())
assert allowed.capabilities.prompts is not None
assert allowed.capabilities.resources is not None
assert denied.capabilities.model_dump(exclude_none=True) == {}
assert allowed.capabilities.tools is not None
allowed.capabilities.tools.list_changed = True
assert capabilities.tools.list_changed is not True
def test_modern_candidates_do_not_enable_public_serving():
modern = REVISION_SUPPORT["2026-07-28"]
assert modern.completed is False
assert "input_required" in modern.results
assert MCPTransport.sse not in modern.transports
assert not any("2026-07-28" in pair for pair in TRANSLATION_PAIRS)
@pytest.mark.asyncio
@pytest.mark.parametrize("versions,accepted", [(("2025-11-25",), True), (("2025-06-18",), False)])
async def test_version_policy_gates_the_actual_sdk_handshake(versions, accepted):
server: Final = Server("test-gateway", version="1")
server.middleware.append(GatewayVersionPolicy(lambda: versions))
if accepted:
async with Client(server, mode="legacy") as client:
assert client.protocol_version == "2025-11-25"
result = await client.session.send_ping()
assert result is not None
else:
with pytest.RaisesGroup(pytest.RaisesExc(MCPError, match="Unsupported MCP protocol version"), flatten_subgroups=True):
async with Client(server, mode="legacy"):
pytest.fail("The excluded revision must not initialize")

View file

@ -10444,11 +10444,12 @@ async def test_active_request_ctx_var_feeds_auth_resolution_recording(_mcp_reque
)
@pytest.mark.parametrize("handler", ("handle_streamable_http_mcp", "handle_sse_mcp"))
async def test_streamable_http_rejects_modern_protocol_version(
header_value: str, expected_rejected: bool, handler: str
header_value: str, expected_rejected: bool, handler: str, monkeypatch: pytest.MonkeyPatch
) -> None:
from litellm.proxy._experimental.mcp_server import server as mcp_module
from litellm.proxy._experimental.mcp_server.server import unsupported_protocol_version
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {})
scope: Scope = {
"type": "http",
"method": "POST",
@ -10659,3 +10660,32 @@ async def test_legacy_sse_mount_emits_message_endpoint(
await incoming.put({"type": "http.disconnect"})
await asyncio.wait_for(task, 2)
assert await post(initialization) == 404
@pytest.mark.parametrize("revision,rejected", [("2024-11-05", False), ("2025-11-25", True), ("2026-07-28", True)])
def test_protocol_header_respects_configured_advertisement(revision, rejected):
from litellm.proxy import proxy_server
from litellm.proxy._experimental.mcp_server.server import unsupported_protocol_version
with patch.dict(proxy_server.general_settings, {"mcp_advertised_versions": ["2024-11-05"]}):
result = unsupported_protocol_version({"headers": [(b"mcp-protocol-version", revision.encode())]})
assert result == (revision if rejected else None)
@pytest.mark.asyncio
async def test_discovery_adapter_preserves_authenticated_context(_mcp_request_ctx):
from mcp.types import DiscoverResult, RequestParams, ServerCapabilities
from litellm.proxy._experimental.mcp_server import server
expected = DiscoverResult(supported_versions=["2025-11-25"], capabilities=ServerCapabilities())
dispatched = AsyncMock(return_value=expected)
auth = UserAPIKeyAuth(user_id="discover-caller")
with (
patch.object(server, "get_or_extract_auth_context", AsyncMock(return_value=(auth, None, ["allowed"], None, None, None, None))),
patch.object(server.operations.GatewayOperations, "execute", dispatched),
):
result = await server.discover(_mcp_request_ctx(), RequestParams())
assert result is expected
context = dispatched.await_args.args[1]
assert context.user_api_key_auth.user_id == "discover-caller"
assert context.mcp_servers == ("allowed",)

View file

@ -63,7 +63,7 @@ from litellm.proxy._types import (
UserAPIKeyAuth,
)
from litellm.types.llms.custom_http import httpxSpecialProvider
from litellm.types.mcp import MCPAuth, MCPAuthType
from litellm.types.mcp import MCPAuth, MCPAuthType, MCPUpstreamProtocol
from litellm.types.mcp_server.mcp_server_manager import MCPOAuthMetadata, MCPServer
from litellm.caching.caching import DualCache
from litellm.caching.llm_caching_handler import LLMClientCache
@ -14637,3 +14637,27 @@ class TestSharedIdentifierPrefixWarning:
assert "srv-b" in shared_warnings[0]
assert "srv-c" not in shared_warnings[0]
assert "'shared'" in shared_warnings[0]
@pytest.mark.asyncio
@pytest.mark.parametrize("revision", ["auto", "2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25"])
async def test_configured_protocol_reaches_the_upstream_client(config_only_mcp_manager_factory, revision):
manager = config_only_mcp_manager_factory()
await manager.load_servers_from_config({"versions": {"url": "http://127.0.0.1:9/mcp", "transport": "http", "protocol_version": revision}})
server = next(iter(manager.config_mcp_servers.values()))
client = await manager._create_mcp_client(server)
assert server.protocol_version == revision
assert client.protocol_version == revision
@pytest.mark.parametrize("revision", ("auto", "2024-11-05", "2025-06-18"))
@pytest.mark.parametrize("explicit", (None, "auto", "2025-11-25"))
def test_runtime_protocol_metadata_preserves_explicit_precedence(
revision: MCPUpstreamProtocol, explicit: MCPUpstreamProtocol | None
) -> None:
server: Final = MCPServer.model_validate({
"server_id": "preview", "name": "preview", "transport": "http",
"mcp_info": {"protocol_version": revision},
**({"protocol_version": explicit} if explicit is not None else {}),
})
assert server.protocol_version == (explicit if explicit is not None else revision)

View file

@ -542,3 +542,126 @@ async def test_local_tool_json_array_is_converted_once_for_the_caller_revision(c
assert [block.text for block in result.content] == [body]
assert result.structured_content == (["a", "b"] if compat == "modern" else None)
@pytest.mark.asyncio
async def test_discovery_preserves_caller_scope_and_proxy_restrictions():
from mcp.types import DiscoverRequest, ListToolsResult, Tool
listed = AsyncMock(return_value=ListToolsResult(tools=[Tool(name="allowed", input_schema={"type": "object"})]))
context = prepare_context(UserAPIKeyAuth(user_id="scoped"), mcp_servers=["only-this"], mcp_proxy_mode=True, protocol_version="2025-06-18")
with patch("litellm.proxy._experimental.mcp_server.operations._execute_handle_list_tools", listed):
result = await GatewayOperations().execute(DiscoverRequest(), context)
assert result.capabilities.tools is not None
assert result.capabilities.resources is None
assert result.capabilities.prompts is None
assert listed.await_args.args[0] is context
assert listed.await_args.args[0].user_api_key_auth.user_id == "scoped"
assert listed.await_args.args[0].mcp_servers == ("only-this",)
@pytest.mark.asyncio
async def test_discovery_denial_cannot_advertise_tools():
from mcp.types import DiscoverRequest
from fastapi import HTTPException
denied = AsyncMock(side_effect=HTTPException(status_code=403, detail="Forbidden"))
with patch("litellm.proxy._experimental.mcp_server.operations._execute_handle_list_tools", denied):
with pytest.raises(HTTPException) as error:
await GatewayOperations().execute(DiscoverRequest(), prepare_context(UserAPIKeyAuth(user_id="denied")))
assert error.value.status_code == 403
@pytest.mark.asyncio
@pytest.mark.parametrize("available", ["none", "resources", "templates", "prompts"])
async def test_discovery_lists_each_capability_with_the_same_caller(available):
from mcp.types import (
DiscoverRequest, ListToolsResult, ListPromptsResult, ListResourcesResult,
ListResourceTemplatesResult, Prompt, Resource, ResourceTemplate,
)
from litellm.proxy._experimental.mcp_server import operations
context = prepare_context(UserAPIKeyAuth(user_id="scoped"), mcp_servers=["authorized"])
tools = AsyncMock(return_value=ListToolsResult(tools=[]))
prompts = AsyncMock(return_value=ListPromptsResult(prompts=[Prompt(name="allowed")] if available == "prompts" else []))
resources = AsyncMock(return_value=ListResourcesResult(resources=[Resource(name="allowed", uri="test://allowed")] if available == "resources" else []))
templates = AsyncMock(return_value=ListResourceTemplatesResult(resource_templates=[ResourceTemplate(name="allowed", uri_template="test://{id}")] if available == "templates" else []))
with (
patch.object(operations, "_execute_handle_list_tools", tools),
patch.object(operations, "_execute_list_prompts", prompts),
patch.object(operations, "_execute_list_resources", resources),
patch.object(operations, "_execute_list_resource_templates", templates),
):
result = await GatewayOperations().execute(DiscoverRequest(), context)
assert result.capabilities.tools is None
assert (result.capabilities.prompts is not None) == (available == "prompts")
assert (result.capabilities.resources is not None) == (available in {"resources", "templates"})
for listing in (tools, prompts, resources, templates):
assert listing.await_args.args[0] is context
@pytest.mark.asyncio
@pytest.mark.parametrize("outcome", ["success", "failure", "cancel"])
async def test_discovery_concurrent_listings_drain_on_failure_and_cancellation(outcome):
from mcp.types import DiscoverRequest, ListToolsResult, ListPromptsResult, ListResourcesResult, ListResourceTemplatesResult
from litellm.proxy._experimental.mcp_server import operations
ready = [asyncio.Event() for _ in range(4)]
closed = [asyncio.Event() for _ in range(4)]
release = asyncio.Event()
responses = (ListToolsResult(tools=[]), ListPromptsResult(prompts=[]), ListResourcesResult(resources=[]), ListResourceTemplatesResult(resource_templates=[]))
def listing(index):
async def run(*args, **kwargs):
ready[index].set()
try:
await release.wait()
if index == 0 and outcome == "failure":
raise ValueError("discovery failed")
if outcome != "success":
await asyncio.Event().wait()
return responses[index]
finally:
closed[index].set()
return run
with (
patch.object(operations, "_execute_handle_list_tools", side_effect=listing(0)) as tools,
patch.object(operations, "_execute_list_prompts", side_effect=listing(1)),
patch.object(operations, "_execute_list_resources", side_effect=listing(2)),
patch.object(operations, "_execute_list_resource_templates", side_effect=listing(3)),
):
task = asyncio.create_task(GatewayOperations().execute(DiscoverRequest(), prepare_context(UserAPIKeyAuth(user_id="scoped"))))
try:
await asyncio.wait_for(asyncio.gather(*(event.wait() for event in ready)), 1)
if outcome == "cancel":
task.cancel()
else:
release.set()
if outcome == "success":
result = await asyncio.wait_for(task, 1)
assert result.capabilities.model_dump(exclude_none=True) == {}
else:
with pytest.raises(asyncio.CancelledError if outcome == "cancel" else ValueError):
await asyncio.wait_for(task, 1)
assert all(event.is_set() for event in closed)
assert tools.call_args.kwargs["log_list_tools_to_spendlogs"] is False
finally:
task.cancel()
await asyncio.gather(task, return_exceptions=True)
@pytest.mark.asyncio
@pytest.mark.parametrize("log_enabled", [False, True])
async def test_tools_listing_preserves_explicit_spend_log_policy(log_enabled):
from mcp.types import PaginatedRequestParams
from litellm.proxy._experimental.mcp_server import operations
listing = AsyncMock(return_value=operations.AggregateToolListing(tools=[], outcomes={}))
with patch.object(operations, "_list_mcp_tools", listing):
result = await operations._execute_handle_list_tools(
prepare_context(UserAPIKeyAuth(user_id="caller")), PaginatedRequestParams(),
log_list_tools_to_spendlogs=log_enabled,
)
assert result.tools == []
assert listing.await_args.kwargs["log_list_tools_to_spendlogs"] is log_enabled

View file

@ -28,7 +28,7 @@ from litellm.proxy._types import (
UserAPIKeyAuth,
)
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.types.mcp import MCPAuth, MCPTransport
from litellm.types.mcp import MCPAuth, MCPTransport, MCPUpstreamProtocol
from litellm.types.mcp_server.mcp_server_manager import MCPServer
_OK_TOOL_RESULT: Final = CallToolResult(content=[TextContent(type="text", text='{"result": "ok"}')], is_error=False)
@ -1476,7 +1476,7 @@ class TestListToolsRestAPI:
monkeypatch,
):
"""The REST tools/list path should include tools beyond the upstream first page."""
from mcp.types import ListToolsResult, PaginatedRequestParams
from mcp.types import Implementation, InitializeResult, ListToolsResult, PaginatedRequestParams, ServerCapabilities
from mcp.types import Tool as MCPTool
import litellm.experimental_mcp_client.client as mcp_client_module
@ -1512,7 +1512,11 @@ class TestListToolsRestAPI:
mock_session_ctx = AsyncMock()
mock_session_instance = AsyncMock()
mock_session_instance.initialize = AsyncMock(return_value=None)
mock_session_instance.initialize = AsyncMock(return_value=InitializeResult(
protocol_version="2025-11-25",
capabilities=ServerCapabilities(),
server_info=Implementation(name="stub", version="1"),
))
mock_session_instance.list_tools.side_effect = [
ListToolsResult(
tools=[
@ -4628,3 +4632,74 @@ class TestClientAllowlistOnRestRoutes:
assert denied.value.detail["error"] == "Forbidden"
assert "'claude-code'" in denied.value.detail["details"]
acting.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize("revision", ("auto", "2024-11-05", "2025-06-18"))
async def test_preview_client_honors_protocol_metadata(revision: MCPUpstreamProtocol) -> None:
from litellm.experimental_mcp_client.client import MCPClient
payload: Final = NewMCPServerRequest(
server_name="preview", url="http://127.0.0.1:9/mcp", transport="http",
auth_type=MCPAuth.none, mcp_info={"protocol_version": revision},
)
async def inspect_client(client: MCPClient) -> dict[str, str]:
return {"protocol_version": client.protocol_version}
result: Final = await rest_endpoints._execute_with_mcp_client(payload, inspect_client)
assert result == {"protocol_version": revision}
@pytest.mark.asyncio
@pytest.mark.parametrize("auth_type", (MCPAuth.none, MCPAuth.bearer_token, MCPAuth.oauth2))
@pytest.mark.parametrize(
("metadata", "expected"),
(
(None, "2025-11-25"),
({}, "2025-11-25"),
({"description": "edited"}, "2025-11-25"),
({"protocol_version": "auto"}, "auto"),
({"protocol_version": "2024-11-05"}, "2024-11-05"),
({"protocol_version": "2025-06-18"}, "2025-06-18"),
),
)
async def test_saved_preview_protocol_omission_and_explicit_edits(
monkeypatch: pytest.MonkeyPatch, auth_type: MCPAuth,
metadata: dict[str, str] | None, expected: MCPUpstreamProtocol,
) -> None:
from starlette.datastructures import Headers
from litellm.experimental_mcp_client.client import MCPClient
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
from litellm.proxy.management_endpoints import mcp_management_endpoints
saved: Final = MCPServer(
server_id="saved-preview", name="preview", url="https://example.com/mcp",
transport="http", auth_type=auth_type, protocol_version="2025-11-25",
authentication_token="stored-token",
authorization_url="https://example.com/authorize", token_url="https://example.com/token",
)
manager: Final = MCPServerManager()
manager.registry = {saved.server_id: saved}
monkeypatch.setattr(rest_endpoints, "global_mcp_server_manager", manager)
monkeypatch.setattr(mcp_management_endpoints, "global_mcp_server_manager", manager)
payload: Final = NewMCPServerRequest(
server_id=saved.server_id, server_name=saved.name, url=saved.url, transport="http",
auth_type=auth_type, mcp_info=metadata,
authorization_url=saved.authorization_url, token_url=saved.token_url,
)
staged: Final = rest_endpoints._stage_server_test(
payload, Headers({"x-litellm-api-key": "sk-admin", "authorization": "Bearer preview-token"})
)
async def inspect_client(client: MCPClient) -> dict[str, str]:
return {"protocol_version": client.protocol_version}
result: Final = await rest_endpoints._execute_with_mcp_client(
staged.request, inspect_client,
mcp_auth_header=staged.mcp_auth_header, oauth2_headers=staged.oauth2_headers,
)
assert result == {"protocol_version": expected}
assert saved.protocol_version == "2025-11-25"
assert payload.mcp_info == metadata

View file

@ -4903,3 +4903,19 @@ async def test_model_refresh_updates_availability_catalog_and_retains_it_on_db_f
assert await pc._get_models_from_db(client) == []
assert pc.auto_router_db_catalog == ()
assert find_many.await_count == 3
@pytest.mark.asyncio
@pytest.mark.parametrize("versions", [None, ["2024-11-05"], [], ["2026-07-28"], ["unknown"]])
async def test_proxy_config_validates_advertised_mcp_versions_at_load(tmp_path, monkeypatch, versions):
config = tmp_path / "mcp-versions.yaml"
config.write_text(json.dumps({"model_list": [], "general_settings": {"mcp_advertised_versions": versions}}))
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", False)
monkeypatch.delenv("LITELLM_CONFIG_BUCKET_NAME", raising=False)
if versions is None or versions == ["2024-11-05"]:
_, _, settings = await ProxyConfig().load_config(router=None, config_file_path=str(config))
assert settings["mcp_advertised_versions"] == versions
return
with pytest.raises(ValidationError):
await ProxyConfig().load_config(router=None, config_file_path=str(config))

View file

@ -377,3 +377,21 @@ def test_change_password_request_passwords_hidden_from_repr():
for rendered in (repr(request), str(request)):
assert "hunter2hunter2" not in rendered
assert "NewP@ssw0rd-2026" not in rendered
@pytest.mark.parametrize("versions", [[], ["2099-01-01"], ["2026-07-28"]])
def test_mcp_advertised_versions_reject_unavailable_revisions(versions):
from pydantic import ValidationError
from litellm.proxy._types import ConfigGeneralSettings
with pytest.raises(ValidationError):
ConfigGeneralSettings(mcp_advertised_versions=versions)
@pytest.mark.parametrize("revision", ["2026-07-28", "unknown", None])
def test_mcp_metadata_rejects_unavailable_upstream_protocol(revision):
from litellm.proxy._types import NewMCPServerRequest, UpdateMCPServerRequest
payload = {"server_id": "test", "transport": "http", "url": "https://example.com/mcp", "mcp_info": {"protocol_version": revision}}
for model in (NewMCPServerRequest, UpdateMCPServerRequest):
with pytest.raises(ValidationError):
model.model_validate(payload)

View file

@ -58,6 +58,15 @@ from litellm.types.mcp_server.mcp_server_manager import MCPServer
_JSONRPC_MESSAGE_ADAPTER: Final = TypeAdapter(JSONRPCMessage)
def _initialized(instructions: str | None = None) -> InitializeResult:
return InitializeResult(
protocol_version=LATEST_HANDSHAKE_VERSION,
capabilities=ServerCapabilities(),
server_info=Implementation(name="test", version="1"),
instructions=instructions,
)
class _MockTransportClient(MCPClient):
"""An MCPClient whose streamable-HTTP transport runs on an httpx2 MockTransport."""
@ -125,7 +134,7 @@ class TestMCPClient:
mock_stdio_client.return_value = mock_stdio_ctx
mock_session_instance = AsyncMock()
mock_session_instance.initialize = AsyncMock()
mock_session_instance.initialize = AsyncMock(return_value=_initialized())
mock_session_ctx = AsyncMock()
mock_session_ctx.__aenter__.return_value = mock_session_instance
mock_session_ctx.__aexit__.return_value = None
@ -168,7 +177,7 @@ class TestMCPClient:
# Mock the session
with patch("litellm.experimental_mcp_client.client.ClientSession") as mock_session:
mock_session_instance = AsyncMock()
mock_session_instance.initialize = AsyncMock()
mock_session_instance.initialize = AsyncMock(return_value=_initialized())
mock_session_ctx = AsyncMock()
mock_session_ctx.__aenter__.return_value = mock_session_instance
mock_session_ctx.__aexit__.return_value = None
@ -214,7 +223,7 @@ class TestMCPClient:
# Mock the session
with patch("litellm.experimental_mcp_client.client.ClientSession") as mock_session:
mock_session_instance = AsyncMock()
mock_session_instance.initialize = AsyncMock()
mock_session_instance.initialize = AsyncMock(return_value=_initialized())
mock_session_ctx = AsyncMock()
mock_session_ctx.__aenter__.return_value = mock_session_instance
mock_session_ctx.__aexit__.return_value = None
@ -266,7 +275,7 @@ class TestMCPClient:
# Mock the session
with patch("litellm.experimental_mcp_client.client.ClientSession") as mock_session:
mock_session_instance = AsyncMock()
mock_session_instance.initialize = AsyncMock()
mock_session_instance.initialize = AsyncMock(return_value=_initialized())
mock_session_ctx = AsyncMock()
mock_session_ctx.__aenter__.return_value = mock_session_instance
mock_session_ctx.__aexit__.return_value = None
@ -413,8 +422,7 @@ class TestMCPClientInstructionsCapture:
)
mock_session = AsyncMock()
init_result = MagicMock()
init_result.instructions = " upstream says hello "
init_result = _initialized(" upstream says hello ")
mock_session.initialize = AsyncMock(return_value=init_result)
session_ctx = MagicMock()
@ -442,8 +450,7 @@ class TestMCPClientInstructionsCapture:
)
mock_session = AsyncMock()
init_result = MagicMock()
init_result.instructions = None
init_result = _initialized()
mock_session.initialize = AsyncMock(return_value=init_result)
session_ctx = MagicMock()
@ -600,8 +607,7 @@ class TestExecuteSessionOperationSurfacesTransportError:
@patch("litellm.experimental_mcp_client.client.ClientSession")
async def test_cleanup_error_after_success_is_swallowed(self, mock_session_cls):
client = MCPClient(server_url="http://example.com/mcp", transport_type="http")
init_result = MagicMock()
init_result.instructions = None
init_result = _initialized()
self._make_session(mock_session_cls, AsyncMock(return_value=init_result))
transport_ctx = self._make_transport(_FakeExceptionGroup("late", [httpx2.ConnectError("late cleanup error")]))
@ -634,7 +640,7 @@ class TestExecuteSessionOperationSurfacesTransportError:
@pytest.mark.parametrize("original_error", (False, True))
@patch("litellm.experimental_mcp_client.client.ClientSession")
async def test_session_exit_cancellation_preserves_original_failure(self, session_class, original_error):
self._make_session(session_class, AsyncMock(return_value=None))
self._make_session(session_class, AsyncMock(return_value=_initialized()))
cancelled: Final = asyncio.CancelledError("cancelled while closing session")
session_class.return_value.__aexit__ = AsyncMock(side_effect=cancelled)
original: Final = RuntimeError("operation failed")
@ -656,7 +662,7 @@ class TestExecuteSessionOperationSurfacesTransportError:
@pytest.mark.parametrize("phase", ("session", "transport"))
@patch("litellm.experimental_mcp_client.client.ClientSession")
async def test_cleanup_preserves_process_exit(self, session_class, phase, signal_type):
self._make_session(session_class, AsyncMock(return_value=None))
self._make_session(session_class, AsyncMock(return_value=_initialized()))
signal: Final = signal_type("process stopping")
if phase == "session":
session_class.return_value.__aexit__ = AsyncMock(side_effect=signal)
@ -670,7 +676,7 @@ class TestExecuteSessionOperationSurfacesTransportError:
@pytest.mark.asyncio
@patch("litellm.experimental_mcp_client.client.ClientSession")
async def test_session_and_termination_share_one_cleanup_deadline(self, session_class):
self._make_session(session_class, AsyncMock(return_value=None))
self._make_session(session_class, AsyncMock(return_value=_initialized()))
deleting: Final = asyncio.Event()
async def close_session(*args):
@ -1883,16 +1889,17 @@ async def test_sse_read_failure_is_preserved() -> None:
@pytest.mark.asyncio
@pytest.mark.parametrize("protocol_version", ["auto", "2025-06-18"])
@pytest.mark.parametrize("transport", [MCPTransport.sse, MCPTransport.stdio])
@pytest.mark.parametrize("mode", ["ok", "closed", "silent"])
async def test_transport_completion_and_normal_messages(transport: MCPTransport, mode: str) -> None:
async def test_transport_completion_and_normal_messages(transport: MCPTransport, mode: str, protocol_version: str) -> None:
from mcp import ClientSession
from litellm.proxy._experimental.mcp_server.rest_endpoints import _connection_error_message
logging_callback: Final = AsyncMock()
read_timeout: Final = 0.2 if mode == "silent" else 30
client: Final = MCPClient(
server_url="https://example.com/sse", transport_type=transport, timeout=read_timeout, logging_callback=logging_callback
server_url="https://example.com/sse", transport_type=transport, timeout=read_timeout, logging_callback=logging_callback, protocol_version=protocol_version
)
async def operation(session: ClientSession) -> CallToolResult:
@ -2754,12 +2761,13 @@ async def test_http_close_cancellation_cannot_turn_into_success(original_error:
@pytest.mark.asyncio
@pytest.mark.parametrize("protocol_version", ("auto", "2025-06-18"))
@pytest.mark.parametrize("cancel_mode", ("scope", "task", "wait_for", "read_timeout"))
@pytest.mark.parametrize("concurrency", (1, 5))
@pytest.mark.parametrize("termination", ("ok", "hang", "hang_body"))
@pytest.mark.parametrize("raise_on_error", (False, True))
async def test_cancellation_delivers_termination_over_tcp(
cancel_mode: str, concurrency: int, termination: str, raise_on_error: bool
cancel_mode: str, concurrency: int, termination: str, raise_on_error: bool, protocol_version: str
) -> None:
started: Final = asyncio.Event()
scope_ready: Final[asyncio.Future[anyio.CancelScope]] = asyncio.get_running_loop().create_future()
@ -2813,6 +2821,8 @@ async def test_cancellation_delivers_termination_over_tcp(
await stop.wait()
return
if payload["method"] == "initialize":
if cancel_mode != "read_timeout":
await asyncio.sleep(0.75)
response: Final = json.dumps(
{
"jsonrpc": "2.0",
@ -2839,7 +2849,7 @@ async def test_cancellation_delivers_termination_over_tcp(
listener: Final = await asyncio.start_server(handle_connection, "127.0.0.1", 0)
port: Final = listener.sockets[0].getsockname()[1]
client: Final = MCPClient(
server_url=f"http://127.0.0.1:{port}/mcp", timeout=2 if cancel_mode == "read_timeout" else 0.5 if termination != "ok" else 30
server_url=f"http://127.0.0.1:{port}/mcp", protocol_version=protocol_version, timeout=2 if cancel_mode == "read_timeout" else 30
)
async def calls():
@ -2866,7 +2876,7 @@ async def test_cancellation_delivers_termination_over_tcp(
try:
task: Final = asyncio.create_task(invoke())
await asyncio.wait_for(started.wait(), 3)
await asyncio.wait_for(started.wait(), 30)
if cancel_mode == "scope":
(await scope_ready).deadline = anyio.current_time() + 0.2
if cancel_mode == "task":
@ -2901,3 +2911,52 @@ async def test_cancellation_delivers_termination_over_tcp(
closed: Final = await asyncio.wait_for(asyncio.gather(*connections, return_exceptions=True), 2)
assert all(result is None or isinstance(result, asyncio.CancelledError) for result in closed), closed
await asyncio.wait_for(listener.wait_closed(), 2)
@pytest.mark.asyncio
@pytest.mark.parametrize("revision", ["2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25", "auto"])
@pytest.mark.parametrize("accepted", [True, False])
@pytest.mark.parametrize("callbacks", [False, True])
async def test_configured_upstream_revision_is_offered_and_checked(revision, accepted, callbacks):
from mcp.types import JSONRPCRequest
from mcp_types.version import LATEST_HANDSHAKE_VERSION
offered = LATEST_HANDSHAKE_VERSION if revision == "auto" else revision
def respond(request):
if request.method == "DELETE":
return httpx2.Response(200)
payload = _JSONRPC_MESSAGE_ADAPTER.validate_json(request.content)
if not isinstance(payload, JSONRPCRequest):
return httpx2.Response(202)
if payload.method == "initialize":
assert payload.params["protocolVersion"] == offered
assert ("sampling" in payload.params["capabilities"]) == callbacks
assert ("elicitation" in payload.params["capabilities"]) == callbacks
return httpx2.Response(200, json={
"jsonrpc": "2.0", "id": payload.id,
"result": {"protocolVersion": offered if accepted else "unsupported",
"capabilities": {"tools": {}}, "serverInfo": {"name": "upstream", "version": "1"}},
})
assert accepted, "No operation may execute after failed version negotiation"
return httpx2.Response(200, json={"jsonrpc": "2.0", "id": payload.id, "result": {"tools": [{"name": "echo", "inputSchema": {"type": "object"}}]}})
client = _MockTransportClient(
respond, server_url="https://example.com/mcp", protocol_version=revision,
sampling_callback=AsyncMock() if callbacks else None,
elicitation_callback=AsyncMock() if callbacks else None,
)
if accepted:
result = await client.list_tools(raise_on_error=True)
assert [tool.name for tool in result] == ["echo"]
else:
with pytest.raises((MCPError, RuntimeError), match="protocol version"):
await client.list_tools(raise_on_error=True)
@pytest.mark.parametrize("revision", ["2026-07-28", "unknown", "", None])
def test_upstream_protocol_configuration_rejects_unavailable_modes(revision):
from pydantic import ValidationError
with pytest.raises(ValidationError):
MCPClient(protocol_version=revision)

View file

@ -12,7 +12,7 @@ import litellm.experimental_mcp_client.client as mcp_client_module
from litellm.experimental_mcp_client.client import MCPClient
from litellm.types.mcp import MCPAuth, MCPTransport
from mcp.types import CallToolResult as MCPCallToolResult
from mcp.types import ListToolsResult, PaginatedRequestParams
from mcp.types import Implementation, InitializeResult, ListToolsResult, PaginatedRequestParams, ServerCapabilities
from mcp.types import Tool as MCPTool
@ -128,6 +128,11 @@ class TestMCPClientUnitTests:
mock_session_ctx = AsyncMock()
mock_session_class.return_value = mock_session_ctx
mock_session_instance = AsyncMock()
mock_session_instance.initialize.return_value = InitializeResult(
protocol_version="2025-11-25",
capabilities=ServerCapabilities(),
server_info=Implementation(name="test-peer", version="1"),
)
mock_session_ctx.__aenter__ = AsyncMock(return_value=mock_session_instance)
client = MCPClient(
@ -163,6 +168,11 @@ class TestMCPClientUnitTests:
mock_session_ctx = AsyncMock()
mock_session_class.return_value = mock_session_ctx
mock_session_instance = AsyncMock()
mock_session_instance.initialize.return_value = InitializeResult(
protocol_version="2025-11-25",
capabilities=ServerCapabilities(),
server_info=Implementation(name="test-peer", version="1"),
)
mock_session_ctx.__aenter__ = AsyncMock(return_value=mock_session_instance)
mock_tools = [
@ -204,6 +214,11 @@ class TestMCPClientUnitTests:
mock_session_ctx = AsyncMock()
mock_session_class.return_value = mock_session_ctx
mock_session_instance = AsyncMock()
mock_session_instance.initialize.return_value = InitializeResult(
protocol_version="2025-11-25",
capabilities=ServerCapabilities(),
server_info=Implementation(name="test-peer", version="1"),
)
mock_session_ctx.__aenter__ = AsyncMock(return_value=mock_session_instance)
first_page_tools = [
@ -245,6 +260,11 @@ class TestMCPClientUnitTests:
mock_session_ctx = AsyncMock()
mock_session_class.return_value = mock_session_ctx
mock_session_instance = AsyncMock()
mock_session_instance.initialize.return_value = InitializeResult(
protocol_version="2025-11-25",
capabilities=ServerCapabilities(),
server_info=Implementation(name="test-peer", version="1"),
)
mock_session_ctx.__aenter__ = AsyncMock(return_value=mock_session_instance)
mock_session_instance.list_tools.side_effect = [
@ -277,6 +297,11 @@ class TestMCPClientUnitTests:
mock_session_ctx = AsyncMock()
mock_session_class.return_value = mock_session_ctx
mock_session_instance = AsyncMock()
mock_session_instance.initialize.return_value = InitializeResult(
protocol_version="2025-11-25",
capabilities=ServerCapabilities(),
server_info=Implementation(name="test-peer", version="1"),
)
mock_session_ctx.__aenter__ = AsyncMock(return_value=mock_session_instance)
mock_result = MCPCallToolResult(content=[])
@ -289,7 +314,7 @@ class TestMCPClientUnitTests:
assert result == mock_result
mock_session_instance.initialize.assert_called_once()
mock_session_instance.call_tool.assert_called_once_with(
name="test_tool", arguments={"arg1": "value1"}, progress_callback=ANY
name="test_tool", arguments={"arg1": "value1"}, progress_callback=ANY, allow_input_required=False
)

View file

@ -489,7 +489,7 @@ async def test_sse_mcp_handler_mock():
)
with (
patch("litellm.proxy._experimental.mcp_server.server.server.run", run),
patch("litellm.proxy._experimental.mcp_server.server.serve_loop", run),
patch(
"litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED",
True,
@ -511,7 +511,7 @@ async def test_sse_mcp_handler_mock():
# Call the handler
await handle_sse_mcp(mock_scope, mock_receive, mock_send)
assert run.await_args.args[:2] == (read_stream, write_stream)
assert run.await_args.args[1:3] == (read_stream, write_stream)
assert mock_sse.connect_sse.call_args.args[0]["path"] == "/mcp/sse"

View file

@ -38,6 +38,7 @@ def _run_shard(tmp_path: Path, test_path: str, workers: str) -> subprocess.Compl
"PATH": f"{shim_dir}{os.pathsep}{os.environ['PATH']}",
"GITHUB_OUTPUT": str(tmp_path / "github_output"),
"TEST_PATH": test_path,
"UNIT_FLAG": "",
"WORKERS": workers,
"UNIT_FLAG": "",
},

View file

@ -28350,6 +28350,11 @@ export interface components {
* @description Maximum retention period for spend logs (e.g., '7d' for 7 days). Logs older than this will be deleted.
*/
maximum_spend_logs_retention_period?: string | null;
/**
* Mcp Advertised Versions
* @description MCP revisions enabled by the gateway. Defaults to all completed legacy revisions. Modern protocol serving and Apps/Tasks remain disabled.
*/
mcp_advertised_versions?: ("2024-11-05" | "2025-03-26" | "2025-06-18" | "2025-11-25")[] | null;
/**
* Mcp Allowed Clients
* @description MCP client applications admitted by the gateway, each an {alias, value} pair where alias is the name shown in the dashboard and logs and value is the identity that must match exactly. When set, every MCP request must carry a client identity equal to one of the values: a JWT caller is identified by the claim named in litellm_jwtauth.mcp_client_id_jwt_field, any other caller by the header named in mcp_client_id_header. A request with no resolvable identity, or an unlisted one, is rejected with 403. Unset means every client is admitted.