refactor(mcp): extract upstream preparation and support modern clients (#44232)

* refactor(mcp): extract upstream preparation and support modern clients

* fix(mcp): reject incompatible upstream transport before saving

* fix(mcp): serialize protocol validation with server updates

---------

Co-authored-by: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com>
This commit is contained in:
joshua-berri 2026-10-02 16:04:11 -07:00 • committed by GitHub
parent 485ad76635
commit 04bc354525
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
22 changed files with 1115 additions and 381 deletions

View file

@ -36,6 +36,7 @@ from mcp.types import (
METHOD_NOT_FOUND,
REQUEST_TIMEOUT,
ClientCapabilities,
DiscoverResult,
ElicitationCapability,
FormElicitationCapability,
GetPromptRequestParams,
@ -84,6 +85,7 @@ from litellm.types.mcp import (
MCPUpstreamProtocol,
credential_redirect_hook,
has_header,
validate_mcp_protocol_transport,
without_header,
)
@ -401,7 +403,10 @@ class MCPClient:
logging_callback: Callable | None = None,
protocol_version: MCPUpstreamProtocol = "auto",
):
self.protocol_version: MCPUpstreamProtocol = TypeAdapter(MCPUpstreamProtocol).validate_python(protocol_version)
self.protocol_version: MCPUpstreamProtocol = TypeAdapter[MCPUpstreamProtocol](
MCPUpstreamProtocol
).validate_python(protocol_version)
validate_mcp_protocol_transport(self.protocol_version, transport_type)
self.server_url: str = server_url
self.transport_type: MCPTransport = transport_type
self.auth_type: MCPAuthType = auth_type
@ -540,6 +545,17 @@ class MCPClient:
return safe_env
async def _prepare_session(self, session: ClientSession) -> InitializeResult | DiscoverResult:
if self.protocol_version != "2026-07-28":
return await self._initialize_session(session)
discovery: Final = DiscoverResult.model_validate(await session.send_discover(self.protocol_version))
if self.protocol_version not in discovery.supported_versions:
raise MCPError(code=-32022, message="Upstream did not accept the configured MCP protocol version")
session.adopt(discovery)
if session.protocol_version != self.protocol_version:
raise MCPError(code=-32022, message="Upstream selected an unsupported MCP protocol version")
return discovery
async def _initialize_session(self, session: ClientSession) -> InitializeResult:
if self.protocol_version == "auto":
automatic: Final = await session.initialize()
@ -623,7 +639,7 @@ class MCPClient:
)
session: Final = await session_ctx.__aenter__()
try:
init_result: Final = await self._initialize_session(session)
init_result: Final = await self._prepare_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

@ -3,12 +3,28 @@ from copy import deepcopy
from dataclasses import dataclass, field
from datetime import datetime
from types import MappingProxyType
from typing import Final, Protocol
from typing import TYPE_CHECKING, Final, Literal, Protocol
from litellm.proxy._experimental.mcp_server.tool_outcome import WireCompat
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.mcp_server.mcp_server_manager import MCPServer
if TYPE_CHECKING:
from litellm.proxy._experimental.mcp_server.server_resolution import ResolvedMCPServer
class TargetCatalog(Protocol):
async def resolve(
self,
server_id: str,
caller: UserAPIKeyAuth,
*,
is_admin_view: bool,
not_found_detail: Mapping[str, str],
forbidden_detail: Mapping[str, str],
non_admin_missing: Literal["not_found", "forbidden"],
) -> "ResolvedMCPServer": ...
def copy_caller(auth: UserAPIKeyAuth | None) -> UserAPIKeyAuth | None:
if auth is None:

View file

@ -8,6 +8,7 @@ from datetime import datetime, timedelta, timezone
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, TypedDict, cast
from fastapi import HTTPException
from pydantic import TypeAdapter
from typing_extensions import ReadOnly
from litellm._logging import verbose_proxy_logger
@ -52,8 +53,8 @@ from litellm.repositories.verification_token_repository import (
VerificationTokenRepository,
)
from litellm.types.llms.custom_http import httpxSpecialProvider
from litellm.types.mcp import MCPCredentials
from litellm.types.mcp_server.mcp_server_manager import PinnedMCPTool
from litellm.types.mcp import MCPCredentials, MCPTransportType, MCPUpstreamProtocol, validate_mcp_protocol_transport
from litellm.types.mcp_server.mcp_server_manager import MCPInfo, PinnedMCPTool
if TYPE_CHECKING:
from prisma import models as prisma_db_models
@ -600,6 +601,7 @@ async def _mcp_server_write_if_identifier_free(
alias: str | None,
exclude_server_id: str | None,
write: "Callable[[TableActions[prisma_db_models.LiteLLM_MCPServerTable]], Awaitable[prisma_db_models.LiteLLM_MCPServerTable | None]]",
lock_server_id: str | None = None,
) -> "prisma_db_models.LiteLLM_MCPServerTable | McpIdentifierConflict | None":
"""Run ``write`` only when no other live row owns ``server_name``/``alias``.
@ -619,6 +621,10 @@ async def _mcp_server_write_if_identifier_free(
)
if conflict is not None:
return conflict
if lock_server_id is not None:
await tx.execute_raw(
'SELECT server_id FROM "LiteLLM_MCPServerTable" WHERE server_id=$1 FOR UPDATE', lock_server_id
)
return await write(tx.litellm_mcpservertable)
@ -1100,6 +1106,27 @@ async def get_draft_mcp_server(
return table
def _validate_mcp_protocol_write(
stored: "prisma_db_models.LiteLLM_MCPServerTable", data_dict: Mapping[str, object]
) -> None:
raw_info: Final = data_dict.get("mcp_info", stored.mcp_info)
adapter: Final = TypeAdapter[MCPInfo | None](MCPInfo | None)
try:
info: Final = (
adapter.validate_json(raw_info) if isinstance(raw_info, str) else adapter.validate_python(raw_info)
)
validate_mcp_protocol_transport(
TypeAdapter[MCPUpstreamProtocol](MCPUpstreamProtocol).validate_python(
(info or {}).get("protocol_version", "auto")
),
TypeAdapter[MCPTransportType](MCPTransportType).validate_python(
data_dict.get("transport", stored.transport)
),
)
except ValueError as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
async def _update_mcp_server_row(
prisma_client: PrismaClient,
*,
@ -1107,29 +1134,36 @@ async def _update_mcp_server_row(
data_dict: Mapping[str, object],
) -> "prisma_db_models.LiteLLM_MCPServerTable | McpIdentifierConflict | None":
identifier_write: Final = any(field in data_dict for field in ("server_name", "alias"))
protocol_write: Final = bool({"transport", "mcp_info"}.intersection(data_dict))
async def _update(
table: "TableActions[prisma_db_models.LiteLLM_MCPServerTable]",
) -> "prisma_db_models.LiteLLM_MCPServerTable | None":
if protocol_write:
stored: Final = await table.find_unique(where={"server_id": server_id})
if stored is None:
return None
_validate_mcp_protocol_write(stored, data_dict)
return await table.update(
where={"server_id": server_id},
data=data_dict,
)
if not identifier_write:
if not identifier_write and not protocol_write:
return await _update(_mcp_server_table_actions(prisma_client))
if "alias" in data_dict and not data_dict["alias"] and "server_name" not in data_dict:
# Clearing the alias drops the prefix to the stored server_name, which
# may already belong to another row, so that name needs the check too.
existing: Final = await _db_find_mcp_server_row(prisma_client, server_id)
if existing is None:
return await _update(_mcp_server_table_actions(prisma_client))
return None
return await _mcp_server_write_if_identifier_free(
prisma_client,
server_name=existing.server_name,
alias=None,
exclude_server_id=server_id,
write=_update,
lock_server_id=server_id if protocol_write else None,
)
return await _mcp_server_write_if_identifier_free(
prisma_client,
@ -1137,6 +1171,7 @@ async def _update_mcp_server_row(
alias=_identifier_field(data_dict, "alias"),
exclude_server_id=server_id,
write=_update,
lock_server_id=server_id if protocol_write else None,
)

View file

@ -58,12 +58,10 @@ from litellm.constants import (
MCP_CLIENT_TIMEOUT,
MCP_HEALTH_CHECK_TIMEOUT,
MCP_METADATA_TIMEOUT,
MCP_NPM_CACHE_DIR,
MCP_STDIO_ALLOWED_COMMANDS,
MCP_TOOL_LISTING_TIMEOUT,
)
from litellm.exceptions import BlockedPiiEntityError, GuardrailRaisedException
from litellm.experimental_mcp_client.client import MCPClient, MCPSigV4Auth, strip_auth_scheme, to_basic_credentials
from litellm.experimental_mcp_client.client import MCPClient, strip_auth_scheme, to_basic_credentials
from litellm.integrations.custom_guardrail import (
_sync_guardrail_info_to_logging_obj, # pyright: ignore[reportPrivateUsage] - the same bridge @log_guardrail_information uses; reimplementing it here would fork the metadata-key logic
)
@ -91,8 +89,6 @@ from litellm.proxy._experimental.mcp_server.mcp_debug import describe_upstream_h
from litellm.proxy._experimental.mcp_server.oauth2_token_cache import (
MCPPerUserTokenCache,
mcp_per_user_token_cache,
resolve_mcp_auth,
resolved_token_header,
)
from litellm.proxy._experimental.mcp_server.oauth_utils import (
_redact_mcp_resource_url,
@ -105,10 +101,8 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials import (
UpstreamCredentialProvider,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import (
prepare_mcp_client,
raise_public,
raise_token_exchange_challenge,
raise_user_oauth_challenge,
to_server_spec,
to_subject,
)
@ -118,19 +112,14 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_sto
from litellm.proxy._experimental.mcp_server.outbound_credentials.per_user_oauth_store import (
LazyPerUserOAuthTokenStore,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.resolver import resolve_credentials_with_source
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchange_provider import (
build_token_exchanger,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
DEFAULT_CREDENTIAL_HEADER,
AuthorizationCodeConfig,
AuthResolution,
ClientCredentialsConfig,
CredError,
IdJagConfig,
PassthroughConfig,
ServerSpec,
TokenExchangeConfig,
)
from litellm.proxy._experimental.mcp_server.result_conversion import (
@ -145,7 +134,6 @@ from litellm.proxy._experimental.mcp_server.sampling_handler import (
from litellm.proxy._experimental.mcp_server.stdio_gate import (
MCP_STDIO_DISABLED_MESSAGE,
is_mcp_stdio_blocked,
is_mcp_stdio_enabled,
warn_if_mcp_stdio_blocked,
)
from litellm.proxy._experimental.mcp_server.tool_catalog_guard import (
@ -154,6 +142,13 @@ from litellm.proxy._experimental.mcp_server.tool_catalog_guard import (
pin_tool_catalog,
scan_tool_descriptions,
)
from litellm.proxy._experimental.mcp_server.upstream import (
passthrough_token_from_mcp_auth_header,
prepare_upstream_client,
resolve_upstream_auth,
take_forwarded_authorization,
to_server_spec_fail_closed,
)
from litellm.proxy._experimental.mcp_server.utils import (
MCP_TOOL_PREFIX_SEPARATOR,
MCPMissingUserEnvVarsError,
@ -204,10 +199,8 @@ from litellm.types.llms.custom_http import httpxSpecialProvider
from litellm.types.mcp import (
DEFAULT_SUBJECT_TOKEN_TYPE,
MCPAuth,
MCPStdioConfig,
MCPTokenEndpointAuthMethod,
MCPUpstreamProtocol,
has_header,
without_header,
)
from litellm.types.mcp_server.mcp_server_manager import (
@ -1234,7 +1227,7 @@ def _resolve_openapi_tool_auth(
Returns the ``Authorization`` value to inject, the extra headers to forward, and the credential to
hand ``resolve_openapi_upstream_auth``, whose passthrough arm reads it via
``_passthrough_token_from_mcp_auth_header``. The per-server Authorization travels only in the
``passthrough_token_from_mcp_auth_header``. The per-server Authorization travels only in the
credential, never also in the forwarded headers, because the resolver pops Authorization out of
those and would otherwise have two sources to reconcile.
"""
@ -1325,35 +1318,6 @@ def _client_forwarded_authorization_headers(
return extra_headers
def _take_forwarded_authorization(
headers: dict[str, str] | None,
) -> tuple[str | None, dict[str, str] | None]:
"""Pop the ``Authorization`` value out of ``headers`` (case-insensitive), returning it with the
remaining headers, so the passthrough resolver arm is the single Authorization source rather than
the header also riding in ``extra_headers`` (which the resolved auth would then defer to)."""
if not headers:
return None, headers
value: Final = next((v for k, v in headers.items() if k.lower() == "authorization"), None)
return value, without_header(headers, DEFAULT_CREDENTIAL_HEADER)
def _passthrough_token_from_mcp_auth_header(
mcp_auth_header: str | dict[str, str] | None,
) -> str | None:
"""The caller's per-server upstream credential for a passthrough-mode server, or None.
Sourced from ``x-mcp-{alias}-authorization`` (string or per-header dict form) or the deprecated
global ``x-mcp-auth`` fallback. Per-server headers are the multi-server shape: they bind one
token to one server, so an aggregate scope with several passthrough-mode servers never replays
a single credential across upstreams. The value is forwarded verbatim, so it must be the full
header value (e.g. ``Bearer <upstream-token>``)."""
if isinstance(mcp_auth_header, str):
return mcp_auth_header or None
if isinstance(mcp_auth_header, dict):
return next((v for k, v in mcp_auth_header.items() if k.lower() == "authorization"), None)
return None
async def _materialize_auth_headers(auth: httpx2.Auth | None) -> dict[str, str] | None:
"""Extract the header a resolved ``httpx2.Auth`` would set, as a plain dict, or None.
@ -1418,25 +1382,6 @@ def _redacted_registry_dump(servers: dict[str, MCPServer]) -> dict[str, dict[str
}
def _to_server_spec_fail_closed(server: MCPServer) -> ServerSpec | None:
"""`to_server_spec`, except a half-configured `oauth2_id_jag` server refuses instead of deferring.
ID-JAG has no v1 arm, so deferring to v1 would let `resolve_mcp_auth` honor a caller x-mcp-*
override or fall through to the static `authentication_token`, both of which bypass the per-user
identity assertion the mode promises. That is an operator misconfiguration, not a fallback.
"""
spec: Final = to_server_spec(server)
if spec is None and server.auth_type == MCPAuth.oauth2_id_jag:
raise_public(
CredError.of_misconfigured(
"oauth2_id_jag requires token_exchange_endpoint, id_jag_resource_token_endpoint, "
"client_id, and a client_secret or client_private_key; refusing to fall back to "
"a static credential."
)
)
return spec
def _caller_authorization_fans_out(
server: MCPServer,
scope_servers: list[MCPServer] | None,
@ -4088,75 +4033,6 @@ class MCPServerManager:
_write_user_env_vars_cache(user_id, server.server_id, values)
return values
async def _resolve_v2_auth(
self,
*,
server: MCPServer,
spec: ServerSpec,
provider: UpstreamCredentialProvider,
subject_token: str | None,
user_api_key_auth: UserAPIKeyAuth | None,
extra_headers: dict[str, str] | None,
) -> tuple[httpx2.Auth | None, dict[str, str] | None]:
"""Resolve a v2-owned server's upstream credential into ``(resolved_auth, extra_headers)``.
On a missing/rejected per-user credential this raises the mode's discovery challenge
(authorization_code's browser-OAuth 401, token_exchange's RFC 9728 challenge) or maps any
other ``CredError`` onto its public HTTP status; it never returns an error as a value.
"""
match await resolve_credentials_with_source(provider, to_subject(user_api_key_auth, subject_token), spec):
case Ok(credential):
auth: Final = credential.auth
# NoOpAuth has no header_name and so never conflicts.
header_name: Final[str | None] = getattr(auth, "header_name", None)
if header_name is None or not extra_headers:
source: Final = (
AuthResolution.extra_headers
if credential.source == AuthResolution.no_auth and extra_headers
else credential.source
)
record_auth_resolution(server.server_id, source)
return auth, extra_headers
if not has_header(extra_headers, header_name):
record_auth_resolution(server.server_id, credential.source)
return auth, extra_headers
if isinstance(
spec.config,
(TokenExchangeConfig, AuthorizationCodeConfig, IdJagConfig, ClientCredentialsConfig),
):
# The resolver owns the credential here (token_exchange's exchanged token,
# authorization_code's stored token, id_jag's minted assertion,
# client_credentials' gateway-minted M2M token). It is authoritative: a
# guardrail such as MCPJWTSigner, static_headers, or any other injected
# Authorization must NOT shadow it (otherwise the upstream gets e.g. the
# signer's JWT instead of the minted token and rejects it, and for M2M the
# one-shot 401 refetch is lost with it). Drop only the header the resolved
# credential is about to occupy, so a static credential the operator aimed at a
# DIFFERENT header still reaches upstream.
record_auth_resolution(server.server_id, credential.source)
return auth, without_header(extra_headers, header_name)
# Other modes: an Authorization already supplied via extra_headers (a forwarded caller
# header or static_headers) is intentional and wins; v1 applies those last.
record_auth_resolution(server.server_id, AuthResolution.extra_headers)
return None, extra_headers
case Error(err):
record_auth_resolution(server.server_id, AuthResolution.failed)
if err.tag == "unauthorized" and isinstance(spec.config, AuthorizationCodeConfig):
# authorization_code's missing per-user token -> the per-server browser-OAuth
# challenge, built here where the full MCPServer is in hand.
raise_user_oauth_challenge(server, root_path=get_request_root_path())
if err.tag == "unauthorized" and isinstance(spec.config, TokenExchangeConfig):
# token_exchange (OBO): a missing/rejected subject token -> the RFC 9728 challenge
# pointing at the IdP the client must SSO with to obtain one, rather than an opaque
# 401. No gateway-side browser flow. An IdP step-up rejection (Entra Conditional
# Access) threads its claims blob into the challenge for the client to satisfy.
raise_token_exchange_challenge(
server,
root_path=get_request_root_path(),
claims=err.unauthorized.claims,
)
raise_public(err)
async def preflight_token_exchange(
self,
server: MCPServer,
@ -4196,7 +4072,7 @@ class MCPServerManager:
case _:
return
resolved_server: Final = await self.ensure_oauth_metadata_discovered(server)
spec: Final = _to_server_spec_fail_closed(resolved_server)
spec: Final = to_server_spec_fail_closed(resolved_server)
if spec is None or not isinstance(spec.config, (TokenExchangeConfig, IdJagConfig)):
return
if subject_token is None and isinstance(spec.config, TokenExchangeConfig):
@ -4226,202 +4102,33 @@ class MCPServerManager:
client_ip: str | None = None,
protocol_version_override: MCPUpstreamProtocol | None = None,
) -> MCPClient:
"""
Create an MCPClient instance for the given server.
Auth resolution (single place for all auth logic):
1. ``mcp_auth_header`` — per-request/per-user override
2. OAuth2 Token Exchange (OBO) — exchange user token for scoped token
3. OAuth2 client_credentials token — auto-fetched and cached
4. ``server.authentication_token`` — static token from config/DB
Args:
server: The server configuration.
mcp_auth_header: Optional per-request auth override.
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.
"""
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
# A caller-supplied per-request override (mcp_auth_header / x-mcp-*) defers to the v1 path
# so it wins - except for the modes the v2 resolver owns per-caller (authorization_code's
# stored token, token_exchange's RFC 8693 minted token, id_jag's minted assertion, and the
# passthrough modes' forwarded caller token). A caller must not be able to substitute another
# user's stored credential, nor silently disable the OBO / ID-JAG exchange and forward an
# arbitrary bearer upstream, so we keep the v2 spec and ignore the override for these; the
# REST tools preview supplies its not-yet-persisted token through the resolver
# (cred_provider), never this path.
if (
spec is not None
and mcp_auth_header
and not isinstance(
spec.config,
(AuthorizationCodeConfig, IdJagConfig, PassthroughConfig, TokenExchangeConfig),
)
):
spec = None
auth_value: Final = await resolve_mcp_auth(resolved_server, mcp_auth_header) if spec is None else None
auth_header_name: Final = resolved_token_header(resolved_server, mcp_auth_header) if spec is None else None
# Create sampling and elicitation callbacks for this client
sampling_cb = (
_create_sampling_callback(
operation_context=OperationContext(
_caller=user_api_key_auth, raw_headers=raw_headers, client_ip=client_ip
)
)
if resolved_server.allow_sampling
else None
)
elicitation_cb: Final = _create_elicitation_callback() if resolved_server.allow_elicitation else None
# Handle stdio transport
if transport == MCPTransport.stdio:
if not is_mcp_stdio_enabled():
raise HTTPException(status_code=403, detail=MCP_STDIO_DISABLED_MESSAGE)
resolved_env: Final = (
stdio_env
if stdio_env is not None
else (dict(resolved_server.env) if resolved_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 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 resolved_server.command:
base_command: Final = os.path.basename(resolved_server.command)
# 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 '{resolved_server.command}' is not in the allowlist ({sorted(MCP_STDIO_ALLOWED_COMMANDS)}). "
f"Add it to LITELLM_MCP_STDIO_EXTRA_COMMANDS to allow this command.",
return await prepare_upstream_client(
resolved_server,
provider=cred_provider or self._cred_provider,
root_path=get_request_root_path(),
mcp_auth_header=mcp_auth_header,
extra_headers=extra_headers,
stdio_env=stdio_env,
subject_token=subject_token,
user_api_key_auth=user_api_key_auth,
protocol_version=(
protocol_version_override if protocol_version_override is not None else resolved_server.protocol_version
),
sampling_callback=(
_create_sampling_callback(
operation_context=OperationContext(
_caller=user_api_key_auth,
raw_headers=raw_headers,
client_ip=client_ip,
)
stdio_config: MCPStdioConfig | None = None
if resolved_server.command and resolved_server.args is not None:
stdio_config = MCPStdioConfig(
command=resolved_server.command,
args=resolved_server.args,
env=resolved_env,
)
record_auth_resolution(server.server_id, AuthResolution.not_applicable)
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),
stdio_config=stdio_config,
extra_headers=extra_headers,
sampling_callback=sampling_cb,
elicitation_callback=elicitation_cb,
)
else:
# For HTTP/SSE transports
server_url: Final = resolved_server.url or ""
if spec is not None:
inbound_token = subject_token
if isinstance(spec.config, PassthroughConfig):
inbound_token, extra_headers = _take_forwarded_authorization(extra_headers)
per_server_token: Final = _passthrough_token_from_mcp_auth_header(mcp_auth_header)
if per_server_token is not None:
inbound_token = per_server_token
resolved_auth, extra_headers = await self._resolve_v2_auth(
server=resolved_server,
spec=spec,
provider=provider,
subject_token=inbound_token,
user_api_key_auth=user_api_key_auth,
extra_headers=extra_headers,
)
return await prepare_mcp_client(
resolved_server,
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
),
extra_headers=extra_headers,
resolved_auth=resolved_auth,
sampling_callback=sampling_cb,
elicitation_callback=elicitation_cb,
),
)
# Create SigV4 auth if configured
aws_auth = None
if resolved_server.auth_type == MCPAuth.aws_sigv4:
aws_auth = MCPSigV4Auth(
aws_access_key_id=resolved_server.aws_access_key_id,
aws_secret_access_key=resolved_server.aws_secret_access_key,
aws_session_token=resolved_server.aws_session_token,
aws_region_name=resolved_server.aws_region_name,
aws_service_name=resolved_server.aws_service_name,
aws_role_name=resolved_server.aws_role_name,
aws_session_name=resolved_server.aws_session_name,
)
legacy_source: Final = (
AuthResolution.aws_sigv4
if aws_auth is not None
else AuthResolution.extra_headers
if extra_headers and has_header(extra_headers, auth_header_name or "Authorization")
else AuthResolution.per_request_header
if mcp_auth_header
else AuthResolution.static_token
if auth_value
else AuthResolution.extra_headers
if extra_headers
else AuthResolution.no_auth
)
record_auth_resolution(server.server_id, legacy_source)
return await prepare_mcp_client(
resolved_server,
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,
timeout=(resolved_server.timeout if resolved_server.timeout is not None else MCP_CLIENT_TIMEOUT),
extra_headers=extra_headers,
aws_auth=aws_auth,
sampling_callback=sampling_cb,
elicitation_callback=elicitation_cb,
),
)
if resolved_server.allow_sampling
else None
),
elicitation_callback=(_create_elicitation_callback() if resolved_server.allow_elicitation else None),
)
async def _get_tools_from_server(
self,
@ -6457,7 +6164,7 @@ class MCPServerManager:
token, token_exchange's exchanged token, passthrough's forwarded caller token) must be
materialized into headers here. Returns ``(resolved_auth_headers, forwarded_headers)``:
the resolved headers are authoritative over every other Authorization source (the same
rule ``_resolve_v2_auth`` applies on the MCPClient path) and ``forwarded_headers`` comes
rule ``resolve_upstream_auth`` applies on the MCPClient path) and ``forwarded_headers`` comes
back with any header the resolver claimed already dropped. Unmigrated (v1) servers resolve
through the stored-token lookup instead, and a missing per-user credential raises the same
discovery challenge the MCPClient path serves, rather than egressing unauthenticated.
@ -6480,11 +6187,12 @@ class MCPServerManager:
if isinstance(spec.config, (TokenExchangeConfig, IdJagConfig)):
subject_token = self._extract_subject_token(oauth2_headers, raw_headers, user_api_key_auth)
elif isinstance(spec.config, PassthroughConfig):
inbound_token, forwarded_headers = _take_forwarded_authorization(forwarded_headers)
per_server_token: Final = _passthrough_token_from_mcp_auth_header(mcp_auth_header)
inbound_token, forwarded_headers = take_forwarded_authorization(forwarded_headers)
per_server_token: Final = passthrough_token_from_mcp_auth_header(mcp_auth_header)
subject_token = per_server_token if per_server_token is not None else inbound_token
resolved_auth, forwarded_headers = await self._resolve_v2_auth(
resolved_auth, forwarded_headers = await resolve_upstream_auth(
root_path=get_request_root_path(),
server=mcp_server,
spec=spec,
provider=self._cred_provider,

View file

@ -120,3 +120,42 @@ async def authorize_mcp_server(
)
return resolved
@dataclass(frozen=True, slots=True)
class MCPServerTargetCatalog:
manager: MCPServerRegistry
db_lookup: Callable[[str], Awaitable[LiteLLM_MCPServerTable | None]] | None = None
temp_lookup: Callable[[str], Awaitable[MCPServer | None]] | None = None
id_client_ip: str | None = None
name_client_ip: str | None = None
match_name: bool = False
async def resolve(
self,
server_id: str,
caller: UserAPIKeyAuth,
*,
is_admin_view: bool,
not_found_detail: Mapping[str, str],
forbidden_detail: Mapping[str, str],
non_admin_missing: Literal["not_found", "forbidden"],
) -> ResolvedMCPServer:
resolved: Final = await resolve_mcp_server(
server_id,
manager=self.manager,
db_lookup=self.db_lookup,
temp_lookup=self.temp_lookup,
id_client_ip=self.id_client_ip,
name_client_ip=self.name_client_ip,
match_name=self.match_name,
)
return await authorize_mcp_server(
resolved,
caller,
manager=self.manager,
is_admin_view=is_admin_view,
not_found_detail=not_found_detail,
forbidden_detail=forbidden_detail,
non_admin_missing=non_admin_missing,
)

View file

@ -0,0 +1,315 @@
from __future__ import annotations
import os
from typing import Final
import httpx2
from fastapi import HTTPException
from litellm.constants import MCP_CLIENT_TIMEOUT, MCP_NPM_CACHE_DIR, MCP_STDIO_ALLOWED_COMMANDS
from litellm.experimental_mcp_client.client import MCPClient, MCPSigV4Auth
from litellm.proxy._experimental.mcp_server.legacy_callbacks import ElicitationCallback, SamplingCallback
from litellm.proxy._experimental.mcp_server.mcp_debug import record_auth_resolution
from litellm.proxy._experimental.mcp_server.oauth2_token_cache import resolve_mcp_auth, resolved_token_header
from litellm.proxy._experimental.mcp_server.outbound_credentials import Error, Ok, UpstreamCredentialProvider
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import (
prepare_mcp_client,
raise_public,
raise_token_exchange_challenge,
raise_user_oauth_challenge,
to_server_spec,
to_subject,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.resolver import resolve_credentials_with_source
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
DEFAULT_CREDENTIAL_HEADER,
AuthorizationCodeConfig,
AuthResolution,
ClientCredentialsConfig,
CredError,
IdJagConfig,
PassthroughConfig,
ServerSpec,
TokenExchangeConfig,
)
from litellm.proxy._experimental.mcp_server.stdio_gate import MCP_STDIO_DISABLED_MESSAGE, is_mcp_stdio_enabled
from litellm.proxy._types import MCPTransport, UserAPIKeyAuth
from litellm.types.mcp import (
MCPAuth,
MCPStdioConfig,
MCPTransportType,
MCPUpstreamProtocol,
has_header,
without_header,
)
from litellm.types.mcp_server.mcp_server_manager import MCPServer
def take_forwarded_authorization(
headers: dict[str, str] | None,
) -> tuple[str | None, dict[str, str] | None]:
"""Pop the ``Authorization`` value out of ``headers`` (case-insensitive), returning it with the
remaining headers, so the passthrough resolver arm is the single Authorization source rather than
the header also riding in ``extra_headers`` (which the resolved auth would then defer to)."""
if not headers:
return None, headers
value: Final = next((v for k, v in headers.items() if k.lower() == "authorization"), None)
return value, without_header(headers, DEFAULT_CREDENTIAL_HEADER)
def passthrough_token_from_mcp_auth_header(
mcp_auth_header: str | dict[str, str] | None,
) -> str | None:
"""The caller's per-server upstream credential for a passthrough-mode server, or None.
Sourced from ``x-mcp-{alias}-authorization`` (string or per-header dict form) or the deprecated
global ``x-mcp-auth`` fallback. Per-server headers are the multi-server shape: they bind one
token to one server, so an aggregate scope with several passthrough-mode servers never replays
a single credential across upstreams. The value is forwarded verbatim, so it must be the full
header value (e.g. ``Bearer <upstream-token>``)."""
if isinstance(mcp_auth_header, str):
return mcp_auth_header or None
if isinstance(mcp_auth_header, dict):
return next((v for k, v in mcp_auth_header.items() if k.lower() == "authorization"), None)
return None
def to_server_spec_fail_closed(server: MCPServer) -> ServerSpec | None:
"""`to_server_spec`, except a half-configured `oauth2_id_jag` server refuses instead of deferring.
ID-JAG has no v1 arm, so deferring to v1 would let `resolve_mcp_auth` honor a caller x-mcp-*
override or fall through to the static `authentication_token`, both of which bypass the per-user
identity assertion the mode promises. That is an operator misconfiguration, not a fallback.
"""
spec: Final = to_server_spec(server)
if spec is None and server.auth_type == MCPAuth.oauth2_id_jag:
raise_public(
CredError.of_misconfigured(
"oauth2_id_jag requires token_exchange_endpoint, id_jag_resource_token_endpoint, "
"client_id, and a client_secret or client_private_key; refusing to fall back to "
"a static credential."
)
)
return spec
async def resolve_upstream_auth(
*,
server: MCPServer,
spec: ServerSpec,
root_path: str,
provider: UpstreamCredentialProvider,
subject_token: str | None,
user_api_key_auth: UserAPIKeyAuth | None,
extra_headers: dict[str, str] | None,
) -> tuple[httpx2.Auth | None, dict[str, str] | None]:
"""Resolve a v2-owned server's upstream credential into ``(resolved_auth, extra_headers)``.
On a missing/rejected per-user credential this raises the mode's discovery challenge
(authorization_code's browser-OAuth 401, token_exchange's RFC 9728 challenge) or maps any
other ``CredError`` onto its public HTTP status; it never returns an error as a value.
"""
match await resolve_credentials_with_source(provider, to_subject(user_api_key_auth, subject_token), spec):
case Ok(credential):
auth: Final = credential.auth
# NoOpAuth has no header_name and so never conflicts.
header_name: Final[str | None] = getattr(auth, "header_name", None)
if header_name is None or not extra_headers:
source: Final = (
AuthResolution.extra_headers
if credential.source == AuthResolution.no_auth and extra_headers
else credential.source
)
record_auth_resolution(server.server_id, source)
return auth, extra_headers
if not has_header(extra_headers, header_name):
record_auth_resolution(server.server_id, credential.source)
return auth, extra_headers
if isinstance(
spec.config,
(TokenExchangeConfig, AuthorizationCodeConfig, IdJagConfig, ClientCredentialsConfig),
):
# The resolver owns the credential here (token_exchange's exchanged token,
# authorization_code's stored token, id_jag's minted assertion,
# client_credentials' gateway-minted M2M token). It is authoritative: a
# guardrail such as MCPJWTSigner, static_headers, or any other injected
# Authorization must NOT shadow it (otherwise the upstream gets e.g. the
# signer's JWT instead of the minted token and rejects it, and for M2M the
# one-shot 401 refetch is lost with it). Drop only the header the resolved
# credential is about to occupy, so a static credential the operator aimed at a
# DIFFERENT header still reaches upstream.
record_auth_resolution(server.server_id, credential.source)
return auth, without_header(extra_headers, header_name)
# Other modes: an Authorization already supplied via extra_headers (a forwarded caller
# header or static_headers) is intentional and wins; v1 applies those last.
record_auth_resolution(server.server_id, AuthResolution.extra_headers)
return None, extra_headers
case Error(err):
record_auth_resolution(server.server_id, AuthResolution.failed)
if err.tag == "unauthorized" and isinstance(spec.config, AuthorizationCodeConfig):
# authorization_code's missing per-user token -> the per-server browser-OAuth
# challenge, built here where the full MCPServer is in hand.
raise_user_oauth_challenge(server, root_path=root_path)
if err.tag == "unauthorized" and isinstance(spec.config, TokenExchangeConfig):
# token_exchange (OBO): a missing/rejected subject token -> the RFC 9728 challenge
# pointing at the IdP the client must SSO with to obtain one, rather than an opaque
# 401. No gateway-side browser flow. An IdP step-up rejection (Entra Conditional
# Access) threads its claims blob into the challenge for the client to satisfy.
raise_token_exchange_challenge(
server,
root_path=root_path,
claims=err.unauthorized.claims,
)
raise_public(err)
def _stdio_config(server: MCPServer, stdio_env: dict[str, str] | None) -> MCPStdioConfig | None:
if not is_mcp_stdio_enabled():
raise HTTPException(status_code=403, detail=MCP_STDIO_DISABLED_MESSAGE)
if server.command:
command: Final = os.path.basename(server.command)
lowercase: Final = command.lower()
normalized: Final = next(
(
lowercase.removesuffix(suffix)
for suffix in (".exe", ".cmd", ".bat", ".com")
if lowercase.endswith(suffix)
),
lowercase,
)
if command not in MCP_STDIO_ALLOWED_COMMANDS and normalized 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)}). "
"Add it to LITELLM_MCP_STDIO_EXTRA_COMMANDS to allow this command.",
)
if not server.command or server.args is None:
return None
environment: Final = stdio_env if stdio_env is not None else server.env
return MCPStdioConfig(
command=server.command,
args=server.args,
env={"NPM_CONFIG_CACHE": MCP_NPM_CACHE_DIR, **environment} if environment is not None else None,
)
async def prepare_upstream_client(
server: MCPServer,
*,
provider: UpstreamCredentialProvider,
root_path: str,
mcp_auth_header: str | dict[str, str] | None = None,
extra_headers: dict[str, str] | None = None,
stdio_env: dict[str, str] | None = None,
subject_token: str | None = None,
user_api_key_auth: UserAPIKeyAuth | None = None,
protocol_version: MCPUpstreamProtocol,
sampling_callback: SamplingCallback | None = None,
elicitation_callback: ElicitationCallback | None = None,
) -> MCPClient:
transport: Final[MCPTransportType] = server.transport or MCPTransport.sse
server_spec: Final = None if transport == MCPTransport.stdio else to_server_spec_fail_closed(server)
spec: Final = (
None
if server_spec is not None
and mcp_auth_header
and not isinstance(
server_spec.config, (AuthorizationCodeConfig, IdJagConfig, PassthroughConfig, TokenExchangeConfig)
)
else server_spec
)
auth_value: Final = await resolve_mcp_auth(server, mcp_auth_header) if spec is None else None
auth_header_name: Final = resolved_token_header(server, mcp_auth_header) if spec is None else None
timeout: Final = server.timeout if server.timeout is not None else MCP_CLIENT_TIMEOUT
if transport == MCPTransport.stdio:
config: Final = _stdio_config(server, stdio_env)
record_auth_resolution(server.server_id, AuthResolution.not_applicable)
return MCPClient(
server_url="",
transport_type=transport,
protocol_version=protocol_version,
auth_type=server.auth_type,
auth_value=auth_value,
timeout=timeout,
stdio_config=config,
extra_headers=extra_headers,
sampling_callback=sampling_callback,
elicitation_callback=elicitation_callback,
)
if spec is not None:
inbound_token, forwarded_headers = (
take_forwarded_authorization(extra_headers)
if isinstance(spec.config, PassthroughConfig)
else (subject_token, extra_headers)
)
per_server_token: Final = (
passthrough_token_from_mcp_auth_header(mcp_auth_header)
if isinstance(spec.config, PassthroughConfig)
else None
)
resolved_auth, resolved_headers = await resolve_upstream_auth(
server=server,
spec=spec,
provider=provider,
root_path=root_path,
subject_token=per_server_token if per_server_token is not None else inbound_token,
user_api_key_auth=user_api_key_auth,
extra_headers=forwarded_headers,
)
return await prepare_mcp_client(
server,
MCPClient(
server_url=server.url or "",
transport_type=transport,
protocol_version=protocol_version,
auth_type=server.auth_type,
timeout=timeout,
extra_headers=resolved_headers,
resolved_auth=resolved_auth,
sampling_callback=sampling_callback,
elicitation_callback=elicitation_callback,
),
)
aws_auth: Final = (
MCPSigV4Auth(
aws_access_key_id=server.aws_access_key_id,
aws_secret_access_key=server.aws_secret_access_key,
aws_session_token=server.aws_session_token,
aws_region_name=server.aws_region_name,
aws_service_name=server.aws_service_name,
aws_role_name=server.aws_role_name,
aws_session_name=server.aws_session_name,
)
if server.auth_type == MCPAuth.aws_sigv4
else None
)
source: Final = (
AuthResolution.aws_sigv4
if aws_auth is not None
else AuthResolution.extra_headers
if extra_headers and has_header(extra_headers, auth_header_name or "Authorization")
else AuthResolution.per_request_header
if mcp_auth_header
else AuthResolution.static_token
if auth_value
else AuthResolution.extra_headers
if extra_headers
else AuthResolution.no_auth
)
record_auth_resolution(server.server_id, source)
return await prepare_mcp_client(
server,
MCPClient(
server_url=server.url or "",
transport_type=transport,
protocol_version=protocol_version,
auth_type=server.auth_type,
auth_value=auth_value,
auth_header_name=auth_header_name,
timeout=timeout,
extra_headers=extra_headers,
aws_auth=aws_auth,
sampling_callback=sampling_callback,
elicitation_callback=elicitation_callback,
),
)

View file

@ -16,6 +16,7 @@ from pydantic import (
JsonValue,
PositiveInt,
PrivateAttr,
TypeAdapter,
field_validator,
model_validator,
)
@ -46,6 +47,8 @@ from litellm.types.mcp import (
MCPCredentials,
MCPTransport,
MCPTransportType,
MCPUpstreamProtocol,
validate_mcp_protocol_transport,
)
from litellm.types.mcp_server.mcp_server_manager import MCPInfo
from litellm.types.proxy.agent_identity import ManagedAgentContext
@ -1673,6 +1676,16 @@ class NewMCPServerRequest(LiteLLMPydanticObjectBase):
description="Server-managed: set by the endpoint; caller values are overridden.",
)
@model_validator(mode="after")
def validate_protocol_transport(self) -> "NewMCPServerRequest":
validate_mcp_protocol_transport(
TypeAdapter[MCPUpstreamProtocol](MCPUpstreamProtocol).validate_python(
(self.mcp_info or {}).get("protocol_version", "auto")
),
self.transport,
)
return self
@model_validator(mode="before")
@classmethod
def validate_transport_fields(cls, values):
@ -1756,6 +1769,18 @@ class UpdateMCPServerRequest(LiteLLMPydanticObjectBase):
timeout: float | None = None
max_concurrent_requests: int | None = None
@model_validator(mode="after")
def validate_protocol_transport(self) -> "UpdateMCPServerRequest":
if not {"transport", "mcp_info"}.issubset(self.model_fields_set):
return self
validate_mcp_protocol_transport(
TypeAdapter[MCPUpstreamProtocol](MCPUpstreamProtocol).validate_python(
(self.mcp_info or {}).get("protocol_version", "auto")
),
self.transport,
)
return self
@model_validator(mode="before")
@classmethod
def validate_transport_fields(cls, values):

View file

@ -44,6 +44,7 @@ from fastapi import (
status,
)
from fastapi.responses import JSONResponse
from pydantic import TypeAdapter
from typing_extensions import ReadOnly, TypedDict
try:
@ -137,6 +138,7 @@ if MCP_AVAILABLE:
def validate_tool_name(name: str) -> _ToolNameValidationResult:
return _ToolNameValidationResult()
from litellm.proxy._experimental.mcp_server.contracts import TargetCatalog
from litellm.proxy._experimental.mcp_server.db import (
McpIdentifierConflict,
approve_mcp_server,
@ -178,6 +180,7 @@ if MCP_AVAILABLE:
global_mcp_server_manager,
)
from litellm.proxy._experimental.mcp_server.server_resolution import (
MCPServerTargetCatalog,
authorize_mcp_server,
resolve_mcp_server,
)
@ -237,7 +240,9 @@ if MCP_AVAILABLE:
MCPCredentials,
MCPGatewaySessionsResponse,
MCPGatewaySessionsTerminateResponse,
MCPUpstreamProtocol,
normalize_upstream_header_name,
validate_mcp_protocol_transport,
)
from litellm.types.mcp_server.mcp_server_manager import MCPServer, PinnedMCPTool
@ -2185,18 +2190,16 @@ if MCP_AVAILABLE:
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
client_ip: Final = IPAddressUtils.get_mcp_client_ip(request) if request is not None else None
resolved: Final = await resolve_mcp_server(
server_id,
catalog: Final[TargetCatalog] = MCPServerTargetCatalog(
manager=global_mcp_server_manager,
temp_lookup=get_cached_temporary_mcp_server,
id_client_ip=None,
name_client_ip=client_ip,
match_name=True,
)
authorized: Final = await authorize_mcp_server(
resolved,
authorized: Final = await catalog.resolve(
server_id,
user_api_key_dict,
manager=global_mcp_server_manager,
is_admin_view=_user_has_admin_view(user_api_key_dict),
not_found_detail={"error": f"MCP server {server_id} not found"},
forbidden_detail={"error": f"Access denied to MCP server {server_id}"},
@ -2739,15 +2742,13 @@ if MCP_AVAILABLE:
404, so server ids can't be enumerated), using the same allowed-server
resolution the MCP gateway enforces on tool calls.
"""
resolved: Final = await resolve_mcp_server(
server_id,
catalog: Final[TargetCatalog] = MCPServerTargetCatalog(
manager=global_mcp_server_manager,
db_lookup=lambda sid: get_mcp_server(prisma_client, sid),
)
authorized: Final = await authorize_mcp_server(
resolved,
authorized: Final = await catalog.resolve(
server_id,
user_api_key_dict,
manager=global_mcp_server_manager,
is_admin_view=_user_has_admin_view(user_api_key_dict),
not_found_detail={"error": f"MCP Server {server_id} not found"},
forbidden_detail={
@ -2938,6 +2939,32 @@ if MCP_AVAILABLE:
statuses.append(status_obj)
return statuses
def _validate_mcp_protocol_update(
payload: UpdateMCPServerRequest,
fields_set: set[str],
stored: LiteLLM_MCPServerTable | None,
read_failed: bool,
) -> None:
if not {"transport", "mcp_info"}.intersection(fields_set):
return
if read_failed:
raise HTTPException(
status_code=503, detail="Cannot validate MCP configuration while stored state is unavailable"
)
if stored is None:
return
effective_transport: Final = payload.transport if "transport" in fields_set else stored.transport
effective_info: Final = payload.mcp_info if "mcp_info" in fields_set else stored.mcp_info
try:
validate_mcp_protocol_transport(
TypeAdapter[MCPUpstreamProtocol](MCPUpstreamProtocol).validate_python(
(effective_info or {}).get("protocol_version", "auto")
),
effective_transport,
)
except ValueError as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
@router.put(
"/server",
description="Allows deleting mcp serves in the db",
@ -2986,9 +3013,9 @@ if MCP_AVAILABLE:
},
)
# Snapshot the pre-update identity so we can detect a mint-relevant change below. The read is
# advisory (it only feeds the stale-token purge decision), so a failure skips the purge with a
# warning instead of failing the edit, whose primary job is the update itself.
# Snapshot stored configuration for protocol validation and mint-relevant changes below.
# Protocol or transport edits require this read; other edits may continue on read failure
# while skipping the best-effort stale-token purge.
try:
old_server_record = await get_mcp_server(prisma_client, payload.server_id)
old_server_record_read_failed = False
@ -3001,6 +3028,8 @@ if MCP_AVAILABLE:
old_server_record = None
old_server_record_read_failed = True
_validate_mcp_protocol_update(payload, payload_fields_set, old_server_record, old_server_record_read_failed)
if payload.per_server_oauth_discovery and (old_server_record is not None or old_server_record_read_failed):
relay_eligible: Final = old_server_record is not None and is_per_server_oauth_discovery_eligible(
payload.auth_type if "auth_type" in payload_fields_set else old_server_record.auth_type,

View file

@ -63,7 +63,14 @@ DEFAULT_SUBJECT_TOKEN_TYPE: Final = "urn:ietf:params:oauth:token-type:access_tok
MCPTransportType = Literal[MCPTransport.sse, MCPTransport.http, MCPTransport.stdio]
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"]
MCPUpstreamProtocol = MCPLegacyVersion | Literal["auto", "2026-07-28"]
def validate_mcp_protocol_transport(protocol_version: MCPUpstreamProtocol, transport: MCPTransportType) -> None:
if protocol_version == "2026-07-28" and transport == MCPTransport.sse:
raise ValueError("Modern MCP requires HTTP or stdio transport")
MCPAdvertisedVersions = Annotated[tuple[MCPLegacyVersion, ...], Field(min_length=1)]
MCPSpecVersionType = Literal[
MCPSpecVersion.nov_2024,

View file

@ -13,13 +13,14 @@ from litellm.types.mcp import (
MCPTransportType,
MCPUpstreamProtocol,
normalize_upstream_header_name,
validate_mcp_protocol_transport,
)
# MCPInfo now allows arbitrary additional fields for custom metadata
def _validate_mcp_protocol_metadata(value: dict[str, object]) -> dict[str, object]:
if "protocol_version" in value:
TypeAdapter(MCPUpstreamProtocol).validate_python(value["protocol_version"])
TypeAdapter[MCPUpstreamProtocol](MCPUpstreamProtocol).validate_python(value["protocol_version"])
return value
@ -277,9 +278,10 @@ class MCPServer(BaseModel):
@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.protocol_version = TypeAdapter[MCPUpstreamProtocol](MCPUpstreamProtocol).validate_python(
self.mcp_info.get("protocol_version", "auto")
)
validate_mcp_protocol_transport(self.protocol_version, self.transport)
return self
@model_validator(mode="after")

View file

@ -390,3 +390,97 @@ def test_ui_session_lists_and_fetches_team_granted_config_server(
assert detail.status_code == 200, f"Team-granted server detail access should succeed: {detail.text}"
assert detail.json()["server_id"] == server_id, detail.text
assert detail.json()["alias"] == alias, detail.text
@pytest.mark.parametrize("explicit_transport", [False, True])
def test_modern_sse_registration_rejected_without_saving(gateway: Gateway, explicit_transport: bool) -> None:
identity: Final = str(uuid.uuid4())
with mcp_peer() as peer:
response: Final = gateway.request("POST", "/v1/mcp/server", {
"server_id": identity, "server_name": "invalid" + uuid.uuid4().hex[:8],
"url": peer.url, "mcp_info": {"protocol_version": "2026-07-28"},
**({"transport": "sse"} if explicit_transport else {}),
})
try:
assert response.status_code == 422, response.text
assert "Modern MCP requires HTTP or stdio" in response.text
assert identity not in _servers(gateway)
assert peer.drain() == (), "Rejected configuration reached upstream"
finally:
if response.status_code == 201:
delete_mcp(gateway, identity)
def test_protocol_transport_updates_validate_effective_configuration(gateway: Gateway) -> None:
from integration._support.database import read_rows
with mcp_peer() as peer, gateway.scenario() as scenario:
identity: Final = register_mcp(scenario, peer, "protocol" + uuid.uuid4().hex[:8])
modern: Final = gateway.request("PUT", "/v1/mcp/server", {
"server_id": identity, "mcp_info": {"protocol_version": "2026-07-28"},
})
assert modern.status_code == 202, modern.text
assert modern.json()["transport"] == "http"
changed: Final = gateway.request("PUT", "/v1/mcp/server", {"server_id": identity, "description": "renamed"})
assert changed.status_code == 202, changed.text
assert changed.json()["mcp_info"]["protocol_version"] == "2026-07-28"
snapshot: Final = read_rows('SELECT transport, mcp_info, updated_at::text FROM "LiteLLM_MCPServerTable" WHERE server_id=%s', (identity,))
peer.drain()
rejected: Final = gateway.request("PUT", "/v1/mcp/server", {"server_id": identity, "transport": "sse", "url": peer.url})
assert rejected.status_code == 400, rejected.text
assert "Modern MCP requires HTTP or stdio" in rejected.text
assert read_rows('SELECT transport, mcp_info, updated_at::text FROM "LiteLLM_MCPServerTable" WHERE server_id=%s', (identity,)) == snapshot
assert tool_calls(peer.drain()) == ()
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
assert call_tool(gateway, key, identity, "add", ADD).status_code == 200
legacy: Final = gateway.request("PUT", "/v1/mcp/server", {"server_id": identity, "transport": "sse", "url": peer.url, "mcp_info": {}})
assert legacy.status_code == 202, legacy.text
legacy_snapshot: Final = read_rows('SELECT transport, mcp_info, updated_at::text FROM "LiteLLM_MCPServerTable" WHERE server_id=%s', (identity,))
rejected_protocol: Final = gateway.request("PUT", "/v1/mcp/server", {"server_id": identity, "mcp_info": {"protocol_version": "2026-07-28"}})
assert rejected_protocol.status_code == 400, rejected_protocol.text
assert read_rows('SELECT transport, mcp_info, updated_at::text FROM "LiteLLM_MCPServerTable" WHERE server_id=%s', (identity,)) == legacy_snapshot
repaired: Final = gateway.request("PUT", "/v1/mcp/server", {"server_id": identity, "transport": "http", "url": peer.url, "mcp_info": {"protocol_version": "2026-07-28"}})
assert repaired.status_code == 202, repaired.text
assert call_tool(gateway, key, identity, "add", ADD).status_code == 200
@pytest.mark.parametrize("rename", [False, True])
def test_concurrent_protocol_transport_edits_cannot_save_incompatible_configuration(
gateway: Gateway, peer: Gateway, rename: bool
) -> None:
import os
from concurrent.futures import ThreadPoolExecutor
import psycopg
from integration._support.database import read_rows
with mcp_peer() as upstream, gateway.scenario() as scenario:
identity: Final = register_mcp(scenario, upstream, "race" + uuid.uuid4().hex[:8])
with ThreadPoolExecutor(max_workers=2) as pool:
# Hold the row so both workers read the old configuration before either can write.
with psycopg.connect(os.environ["DATABASE_URL"]) as blocker:
blocker.execute('SELECT server_id FROM "LiteLLM_MCPServerTable" WHERE server_id=%s FOR UPDATE', (identity,))
protocol = pool.submit(gateway.request, "PUT", "/v1/mcp/server", {
"server_id": identity, "mcp_info": {"protocol_version": "2026-07-28"},
**({"alias": "renamed" + uuid.uuid4().hex[:8]} if rename else {}),
})
transport = pool.submit(peer.request, "PUT", "/v1/mcp/server", {
"server_id": identity, "transport": "sse", "url": upstream.url,
})
eventually(
lambda: read_rows("SELECT pid FROM pg_stat_activity WHERE datname=current_database() AND wait_event_type='Lock' AND query LIKE %s", ('%LiteLLM_MCPServerTable%',)),
lambda rows: len(rows) >= 2,
seconds=3,
)
responses: Final = [protocol.result(), transport.result()]
assert sorted(response.status_code for response in responses) == [202, 400], [r.text for r in responses]
saved: Final = read_rows('SELECT transport, mcp_info FROM "LiteLLM_MCPServerTable" WHERE server_id=%s', (identity,))[0]
assert saved["transport"] == "http" or saved["mcp_info"] != {"protocol_version": "2026-07-28"}
repaired: Final = gateway.request("PUT", "/v1/mcp/server", {
"server_id": identity, "transport": "http", "url": upstream.url,
"mcp_info": {"protocol_version": "2026-07-28"},
})
assert repaired.status_code == 202, repaired.text
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
assert call_tool(gateway, key, identity, "add", ADD).status_code == 200

View file

@ -195,3 +195,52 @@ def test_pinned_revision_pairs_list_and_call_through_gateway(
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
@pytest.mark.parametrize("downstream", ("2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25"))
@pytest.mark.parametrize("peer_kind", ("http", "stdio"))
def test_legacy_gateway_calls_modern_upstream_without_initialize(
gateway: Gateway,
downstream: str,
peer_kind: PeerKind,
) -> 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 = "modern" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(scenario, peer, alias, mcp_info={"protocol_version": "2026-07-28"})
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
client: Final = MCPClient(
server_url=str(gateway.client.base_url).rstrip("/") + "/mcp",
transport_type=MCPTransport.http,
protocol_version=downstream,
extra_headers={"Authorization": f"Bearer {key}"},
)
async def exercise() -> None:
listed: Final = await client.list_tools(raise_on_error=True)
assert f"{alias}-add" in tuple(tool.name for tool in listed)
called: Final = await client.call_tool(
CallToolRequestParams(name=f"{alias}-add", arguments={"a": 2, "b": 3}),
raise_on_error=True,
)
assert called.is_error is False
assert called.content[0].text == "5"
peer.drain()
asyncio.run(exercise())
observed: Final = peer.drain()
assert len(tool_calls(observed)) == 1
assert all(row["body"].get("method") not in ("initialize", "notifications/initialized") for row in observed)
requests: Final = tuple(row for row in observed if "id" in row["body"])
assert requests
for row in requests:
metadata: Final = row["body"]["params"]["_meta"]
assert metadata["io.modelcontextprotocol/protocolVersion"] == "2026-07-28"
assert "io.modelcontextprotocol/clientCapabilities" in metadata
assert b"mcp-session-id" not in row.get("headers", {})

View file

@ -3028,9 +3028,101 @@ async def test_configured_upstream_revision_is_offered_and_checked(revision, acc
await client.list_tools(raise_on_error=True)
@pytest.mark.parametrize("revision", ["2026-07-28", "unknown", "", None])
@pytest.mark.parametrize("revision", ["unknown", "", None])
def test_upstream_protocol_configuration_rejects_unavailable_modes(revision):
from pydantic import ValidationError
with pytest.raises(ValidationError):
MCPClient(protocol_version=revision)
@pytest.mark.asyncio
@pytest.mark.parametrize("accepted", [True, False])
async def test_modern_upstream_requests_are_self_contained_without_initialization(accepted: bool) -> None:
from queue import SimpleQueue
from mcp.types import DiscoverResult, ToolsCapability
methods: Final[SimpleQueue[str]] = SimpleQueue()
def respond(request: httpx2.Request) -> httpx2.Response:
if request.method != "POST":
return httpx2.Response(405)
payload: Final = _JSONRPC_MESSAGE_ADAPTER.validate_json(request.content)
assert isinstance(payload, JSONRPCRequest)
methods.put(payload.method)
assert payload.method != "initialize", "Modern operations must not establish a legacy session"
assert "mcp-session-id" not in request.headers
assert request.headers["mcp-protocol-version"] == "2026-07-28"
assert request.headers["authorization"] == "Bearer upstream-credential"
assert payload.params is not None
metadata: Final = payload.params["_meta"]
assert metadata["io.modelcontextprotocol/protocolVersion"] == "2026-07-28"
assert "io.modelcontextprotocol/clientCapabilities" in metadata
if payload.method == "server/discover":
discovery: Final = DiscoverResult(
supported_versions=["2026-07-28"] if accepted else ["2025-11-25"],
capabilities=ServerCapabilities(tools=ToolsCapability()),
instructions="modern instructions",
)
return httpx2.Response(
200,
json={
"jsonrpc": "2.0",
"id": payload.id,
"result": discovery.model_dump(by_alias=True, exclude_none=True),
},
)
assert accepted, "Rejected negotiation must prevent upstream execution"
if payload.method == "tools/list":
return httpx2.Response(
200,
json={
"jsonrpc": "2.0",
"id": payload.id,
"result": {
"resultType": "complete",
"cacheScope": "private",
"ttlMs": 0,
"tools": [{"name": "add", "inputSchema": {"type": "object"}}],
},
},
)
assert payload.method == "tools/call"
assert payload.params["arguments"] == {"a": 2, "b": 3}
return httpx2.Response(
200,
json={
"jsonrpc": "2.0",
"id": payload.id,
"result": {"resultType": "complete", "content": [{"type": "text", "text": "5"}], "isError": False},
},
)
client: Final = _MockTransportClient(
respond,
server_url="https://example.com/mcp",
protocol_version="2026-07-28",
auth_type=MCPAuth.bearer_token,
auth_value="upstream-credential",
)
params: Final = CallToolRequestParams(name="add", arguments={"a": 2, "b": 3})
if accepted:
result: Final = await client.call_tool(params, raise_on_error=True)
assert result.content[0].text == "5"
assert not result.is_error
assert client._last_initialize_instructions == "modern instructions"
assert tuple(methods.get_nowait() for _ in range(methods.qsize())) == (
"server/discover",
"tools/call",
"tools/list",
)
else:
with pytest.raises((MCPError, RuntimeError), match="protocol version"):
await client.call_tool(params, raise_on_error=True)
assert tuple(methods.get_nowait() for _ in range(methods.qsize())) == ("server/discover",)
def test_modern_upstream_rejects_legacy_sse_transport() -> None:
with pytest.raises(ValueError, match="transport"):
MCPClient(protocol_version="2026-07-28", transport_type=MCPTransport.sse)

View file

@ -111,7 +111,7 @@ async def test_mcp_cost_tracking():
local_mcp_server_manager = MCPServerManager()
with patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient",
"litellm.proxy._experimental.mcp_server.upstream.MCPClient",
mock_client_constructor,
):
# Load the server config
@ -244,7 +244,7 @@ async def test_mcp_cost_tracking_per_tool():
local_mcp_server_manager = MCPServerManager()
with patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient",
"litellm.proxy._experimental.mcp_server.upstream.MCPClient",
mock_client_constructor,
):
# Load the server config with per-tool costs
@ -417,7 +417,7 @@ async def test_mcp_tool_call_hook():
local_mcp_server_manager = MCPServerManager()
with patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient",
"litellm.proxy._experimental.mcp_server.upstream.MCPClient",
mock_client_constructor,
):
# Load the server config

View file

@ -13,6 +13,7 @@ from unittest.mock import AsyncMock, MagicMock
import pytest
from prisma import Json, models
from fastapi import HTTPException
from litellm.proxy._experimental.mcp_server.db import (
create_mcp_server,
@ -35,7 +36,7 @@ def _mock_prisma():
mock_prisma.db.litellm_mcpservertable.update = AsyncMock(return_value=row)
mock_prisma.db.litellm_mcpservertable.create = AsyncMock(return_value=row)
mock_prisma.db.litellm_mcpservertable.find_first = AsyncMock(return_value=None)
mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=None)
mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=row)
tx_client = MagicMock()
tx_client.execute_raw = AsyncMock()
tx_client.litellm_mcpservertable = mock_prisma.db.litellm_mcpservertable
@ -1141,6 +1142,42 @@ async def test_set_mcp_server_pinned_tools_writes_the_snapshot_and_null_clears_i
@pytest.mark.asyncio
async def test_set_mcp_server_pinned_tools_on_a_missing_server_writes_nothing():
mock_prisma = _mock_prisma()
mock_prisma.db.litellm_mcpservertable.find_unique.return_value = None
assert await set_mcp_server_pinned_tools(mock_prisma, "ghost", None, "admin") is None
mock_prisma.db.litellm_mcpservertable.update.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize("protocol_only", [False, True])
async def test_protocol_update_revalidates_current_stored_configuration_before_writing(protocol_only: bool):
prisma = _mock_prisma()
table = prisma.db.litellm_mcpservertable
table.find_unique.return_value = models.LiteLLM_MCPServerTable.model_construct(
server_id="test-server", transport="sse" if protocol_only else "http",
mcp_info={} if protocol_only else {"protocol_version": "2026-07-28"}, env={}, env_vars=[],
)
payload = UpdateMCPServerRequest.model_validate({
"server_id": "test-server",
**({"mcp_info": {"protocol_version": "2026-07-28"}} if protocol_only else {"transport": "sse", "url": "https://upstream.example/sse"}),
})
with pytest.raises(HTTPException) as error:
await update_mcp_server(prisma, payload, "admin")
assert error.value.status_code == 400
assert "Modern MCP requires HTTP or stdio" in str(error.value.detail)
table.update.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize("clear_alias", [False, True])
async def test_protocol_update_preserves_missing_server_without_writing(clear_alias: bool):
prisma = _mock_prisma()
table = prisma.db.litellm_mcpservertable
table.find_unique.return_value = None
payload = UpdateMCPServerRequest.model_validate({
"server_id": "missing", "mcp_info": {"protocol_version": "2026-07-28"},
**({"alias": None} if clear_alias else {}),
})
result = await update_mcp_server(prisma, payload, "admin")
assert result is None
table.update.assert_not_awaited()

View file

@ -71,7 +71,7 @@ async def test_mcp_server_manager_https_server():
return mock_client
with patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient",
"litellm.proxy._experimental.mcp_server.upstream.MCPClient",
mock_client_constructor,
):
await mcp_server_manager.load_servers_from_config(
@ -179,7 +179,7 @@ async def test_mcp_http_transport_list_tools_mock():
return mock_client
with patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient",
"litellm.proxy._experimental.mcp_server.upstream.MCPClient",
mock_client_constructor,
):
# Load server config with HTTP transport
@ -256,7 +256,7 @@ async def test_mcp_http_transport_call_tool_mock():
return mock_client
with patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient",
"litellm.proxy._experimental.mcp_server.upstream.MCPClient",
mock_client_constructor,
):
# Load server config with HTTP transport
@ -322,7 +322,7 @@ async def test_mcp_http_transport_call_tool_error_mock():
return mock_client
with patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient",
"litellm.proxy._experimental.mcp_server.upstream.MCPClient",
mock_client_constructor,
):
# Load server config with HTTP transport
@ -1093,7 +1093,7 @@ async def test_list_tools_only_returns_allowed_servers(monkeypatch):
return mock_client
with patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient",
"litellm.proxy._experimental.mcp_server.upstream.MCPClient",
mock_client_constructor,
):
# Call list_tools
@ -1390,7 +1390,7 @@ async def test_mcp_server_manager_alias_tool_prefixing():
return mock_client
with patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient",
"litellm.proxy._experimental.mcp_server.upstream.MCPClient",
mock_client_constructor,
):
# Get tools from server
@ -1450,7 +1450,7 @@ async def test_mcp_server_manager_server_name_tool_prefixing():
return mock_client
with patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient",
"litellm.proxy._experimental.mcp_server.upstream.MCPClient",
mock_client_constructor,
):
# Get tools from server
@ -1510,7 +1510,7 @@ async def test_mcp_server_manager_server_id_tool_prefixing():
return mock_client
with patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient",
"litellm.proxy._experimental.mcp_server.upstream.MCPClient",
mock_client_constructor,
):
# Get tools from server
@ -2506,7 +2506,7 @@ async def test_filter_tools_by_allowed_tools_integration():
# Mock the MCPClient constructor
with patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient",
"litellm.proxy._experimental.mcp_server.upstream.MCPClient",
mock_client_constructor,
):
# Call _get_tools_from_mcp_servers which should apply the filtering
@ -2620,7 +2620,7 @@ async def test_filter_tools_by_disallowed_tools_integration():
# Mock the MCPClient constructor
with patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient",
"litellm.proxy._experimental.mcp_server.upstream.MCPClient",
mock_client_constructor,
):
# Call _get_tools_from_mcp_servers which should apply the filtering
@ -2722,7 +2722,7 @@ async def test_filter_tools_no_restrictions_integration():
# Mock the MCPClient constructor
with patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient",
"litellm.proxy._experimental.mcp_server.upstream.MCPClient",
mock_client_constructor,
):
# Call _get_tools_from_mcp_servers which should apply the filtering

View file

@ -1,3 +1,4 @@
from litellm.proxy._experimental.mcp_server.upstream import resolve_upstream_auth
import importlib
import asyncio
import functools
@ -94,7 +95,7 @@ async def test_manager_sampling_preserves_explicit_headers_without_ambient_conte
client.call_tool = AsyncMock(return_value=CallToolResult(content=[]))
assert legacy_server.get_active_auth_context() is None
with (
patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient", return_value=client) as factory,
patch("litellm.proxy._experimental.mcp_server.upstream.MCPClient", return_value=client) as factory,
patch("litellm.proxy._experimental.mcp_server.sampling_handler.handle_sampling_create_message", sampling),
):
await MCPServerManager()._call_regular_mcp_tool(
@ -1343,7 +1344,7 @@ class TestMCPServerManager:
"ensure_oauth_metadata_discovered",
new=ensure_oauth_metadata_discovered,
),
patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient"),
patch("litellm.proxy._experimental.mcp_server.upstream.MCPClient"),
):
await manager._create_mcp_client(server)
@ -3900,10 +3901,10 @@ class TestMCPServerManager:
)
with (
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.resolve_mcp_auth",
"litellm.proxy._experimental.mcp_server.upstream.resolve_mcp_auth",
new_callable=AsyncMock,
) as mock_resolve,
patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient") as mock_client_cls,
patch("litellm.proxy._experimental.mcp_server.upstream.MCPClient") as mock_client_cls,
):
await manager._create_mcp_client(server=server, extra_headers={"Authorization": "Bearer upstream-token"})
mock_resolve.assert_not_awaited()
@ -3953,10 +3954,10 @@ class TestMCPServerManager:
)
with (
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.resolve_mcp_auth",
"litellm.proxy._experimental.mcp_server.upstream.resolve_mcp_auth",
new_callable=AsyncMock,
) as mock_resolve,
patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient") as mock_client_cls,
patch("litellm.proxy._experimental.mcp_server.upstream.MCPClient") as mock_client_cls,
):
await manager._create_mcp_client(
server=server,
@ -3996,10 +3997,10 @@ class TestMCPServerManager:
)
with (
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.resolve_mcp_auth",
"litellm.proxy._experimental.mcp_server.upstream.resolve_mcp_auth",
new_callable=AsyncMock,
) as mock_resolve,
patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient") as mock_client_cls,
patch("litellm.proxy._experimental.mcp_server.upstream.MCPClient") as mock_client_cls,
):
await manager._create_mcp_client(
server=server,
@ -13707,7 +13708,8 @@ async def test_debug_resolution_matches_final_header_conflict_winner(_mcp_reques
"none": NoneConfig(),
}[config]
try:
auth, remaining = await MCPServerManager()._resolve_v2_auth(
auth, remaining = await resolve_upstream_auth(
root_path="",
server=MCPServer(
server_id="s",
name="s",
@ -15085,7 +15087,7 @@ async def test_client_sampling_does_not_fill_explicit_context_from_another_ambie
try:
legacy_server.set_auth_context(UserAPIKeyAuth(user_id="unrelated"), raw_headers={"authorization": "unrelated-credential"}, client_ip="192.0.2.99")
with (
patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient") as factory,
patch("litellm.proxy._experimental.mcp_server.upstream.MCPClient") as factory,
patch("litellm.proxy._experimental.mcp_server.sampling_handler.handle_sampling_create_message", sampling),
):
if legacy_factory:
@ -15800,3 +15802,57 @@ class TestToolCatalogGuard:
proxy_logging_obj=proxy_logging_obj,
server=server,
)
@pytest.mark.asyncio
@pytest.mark.parametrize("command,args", [(None, []), ("python", None), ("blocked-executable", [])])
async def test_upstream_preparation_rejects_blocked_or_preserves_incomplete_stdio_config(
monkeypatch: pytest.MonkeyPatch,
command: str | None,
args: list[str] | None,
) -> None:
monkeypatch.setenv("LITELLM_ENABLE_MCP_STDIO", "true")
server: Final = MCPServer(server_id="stdio", name="stdio", transport=MCPTransport.stdio, command=command, args=args)
if command == "blocked-executable":
with pytest.raises(HTTPException) as error:
await MCPServerManager()._create_mcp_client(server)
assert error.value.status_code == 403
assert "not in the allowlist" in error.value.detail
else:
client: Final = await MCPServerManager()._create_mcp_client(server)
assert client.stdio_config is None
@pytest.mark.asyncio
async def test_upstream_preparation_preserves_windows_command_and_caller_environment(
monkeypatch: pytest.MonkeyPatch,
) -> None:
from litellm.constants import MCP_NPM_CACHE_DIR
monkeypatch.setenv("LITELLM_ENABLE_MCP_STDIO", "true")
environment: Final = {"PEER_USER": "alice"}
server: Final = MCPServer(
server_id="stdio", name="stdio", transport=MCPTransport.stdio, command="python.exe", args=[]
)
client: Final = await MCPServerManager()._create_mcp_client(server, stdio_env=environment)
assert client.stdio_config == {
"command": "python.exe",
"args": [],
"env": {"PEER_USER": "alice", "NPM_CONFIG_CACHE": MCP_NPM_CACHE_DIR},
}
assert environment == {"PEER_USER": "alice"}
@pytest.mark.asyncio
async def test_upstream_preparation_honors_case_sensitive_extra_command(monkeypatch: pytest.MonkeyPatch) -> None:
from litellm.proxy._experimental.mcp_server import upstream
monkeypatch.setenv("LITELLM_ENABLE_MCP_STDIO", "true")
monkeypatch.setattr(upstream, "MCP_STDIO_ALLOWED_COMMANDS", frozenset({"CustomRunner"}))
server: Final = MCPServer(
server_id="custom-stdio", name="custom-stdio", transport=MCPTransport.stdio,
command="/opt/tools/CustomRunner", args=[],
)
client: Final = await MCPServerManager()._create_mcp_client(server)
assert client.stdio_config is not None
assert client.stdio_config["command"] == "/opt/tools/CustomRunner"

View file

@ -851,3 +851,28 @@ def test_the_openapi_arm_keeps_the_shared_client_when_no_guard_is_needed(resolve
assert not client.client.event_hooks.get("request")
finally:
_request_resolved_auth_headers.reset(token)
@pytest.mark.asyncio
@pytest.mark.parametrize("mode", [MCPAuth.true_passthrough, MCPAuth.oauth_delegate])
@pytest.mark.parametrize("per_server", [None, "Bearer per-server"])
async def test_openapi_passthrough_preparation_preserves_credential_precedence(
mode: MCPAuth, per_server: str | None,
) -> None:
from typing import Final
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
server: Final = MCPServer(
server_id="openapi-passthrough", name="openapi-passthrough", transport=MCPTransport.http,
url="https://upstream.example", spec_path="https://upstream.example/openapi.json", auth_type=mode,
)
headers: Final = {"authorization": "Bearer forwarded", "X-Trace": "trace"}
resolved, remaining = await MCPServerManager().resolve_openapi_upstream_auth(
mcp_server=server, oauth2_headers=None, raw_headers=None,
mcp_auth_header=per_server, user_api_key_auth=UserAPIKeyAuth(user_id="alice"),
forwarded_headers=headers,
)
assert resolved == {"Authorization": per_server or "Bearer forwarded"}
assert remaining == {"X-Trace": "trace"}
assert headers == {"authorization": "Bearer forwarded", "X-Trace": "trace"}

View file

@ -136,7 +136,7 @@ async def test_prompt_sampling_receives_explicit_operation_caller_headers_and_ip
sampling = AsyncMock()
with (
patch.object(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[upstream])),
patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient", return_value=client) as factory,
patch("litellm.proxy._experimental.mcp_server.upstream.MCPClient", return_value=client) as factory,
patch("litellm.proxy._experimental.mcp_server.sampling_handler.handle_sampling_create_message", sampling),
):
result = await GatewayOperations().execute(

View file

@ -10,7 +10,9 @@ from unittest.mock import Mock
import pytest
from fastapi import HTTPException
from litellm.proxy._experimental.mcp_server.contracts import TargetCatalog
from litellm.proxy._experimental.mcp_server.server_resolution import (
MCPServerTargetCatalog,
ResolutionSource,
ResolvedMCPServer,
authorize_mcp_server,
@ -460,3 +462,94 @@ async def test_missing_alias_does_not_produce_a_resolution() -> None:
manager: Final = _manager()
assert await resolve_mcp_server("missing", manager=manager, match_name=True) is None
manager.name_lookup_spy.assert_called_once_with("missing", None)
@pytest.mark.asyncio
@pytest.mark.parametrize("source", ["db", "registry", "temp"])
@pytest.mark.parametrize("allowed", [False, True])
async def test_target_catalog_authorizes_canonical_identity(source: ResolutionSource, allowed: bool) -> None:
server: Final = _runtime_server()
manager: Final = _manager(
servers_by_name={"requested-alias": server},
allowed_server_ids=(server.server_id,) if allowed else ("requested-alias",),
)
async def database_lookup(server_id: str) -> LiteLLM_MCPServerTable | None:
return _table_server(server.server_id)
async def temporary_lookup(server_id: str) -> MCPServer | None:
return server
catalog: Final[TargetCatalog] = MCPServerTargetCatalog(
manager=manager,
db_lookup=database_lookup if source == "db" else None,
temp_lookup=temporary_lookup if source == "temp" else None,
match_name=True,
)
operation: Final = catalog.resolve(
"requested-alias",
_auth(),
is_admin_view=False,
not_found_detail={"error": "missing"},
forbidden_detail={"error": "denied"},
non_admin_missing="forbidden",
)
if allowed and source != "temp":
result: Final = await operation
assert result.table.server_id == server.server_id
assert result.source == source
assert result.runtime is (None if source == "db" else server)
else:
with pytest.raises(HTTPException) as error:
await operation
assert (error.value.status_code, error.value.detail) == (403, {"error": "denied"})
@pytest.mark.asyncio
@pytest.mark.parametrize(
"admin,missing,status_code", [(True, "forbidden", 404), (False, "forbidden", 403), (False, "not_found", 404)]
)
async def test_target_catalog_preserves_missing_target_policy(
admin: bool,
missing: Literal["forbidden", "not_found"],
status_code: int,
) -> None:
catalog: Final[TargetCatalog] = MCPServerTargetCatalog(manager=_manager())
with pytest.raises(HTTPException) as error:
await catalog.resolve(
"missing",
_auth(),
is_admin_view=admin,
not_found_detail={"error": "missing"},
forbidden_detail={"error": "denied"},
non_admin_missing=missing,
)
assert error.value.status_code == status_code
assert error.value.detail == {"error": "missing" if status_code == 404 else "denied"}
@pytest.mark.asyncio
async def test_target_catalog_does_not_reuse_admin_authorization_for_another_caller() -> None:
server: Final = _runtime_server()
manager: Final = _manager(servers_by_id={server.server_id: server})
catalog: Final[TargetCatalog] = MCPServerTargetCatalog(manager=manager)
admin: Final = await catalog.resolve(
server.server_id,
UserAPIKeyAuth(user_id="admin"),
is_admin_view=True,
not_found_detail={"error": "missing"},
forbidden_detail={"error": "denied"},
non_admin_missing="forbidden",
)
assert admin.runtime is server
with pytest.raises(HTTPException) as error:
await catalog.resolve(
server.server_id,
_auth(),
is_admin_view=False,
not_found_detail={"error": "missing"},
forbidden_detail={"error": "denied"},
non_admin_missing="forbidden",
)
assert (error.value.status_code, error.value.detail) == (403, {"error": "denied"})
manager.allowed_servers_spy.assert_called_once_with(_auth())

View file

@ -11059,3 +11059,89 @@ class TestMCPServerResolutionCharacterization:
health_check.assert_not_awaited()
effects.assert_no_writes()
assert httpx_mock.calls.call_count == 0
@pytest.mark.parametrize("explicit_transport", [False, True])
def test_modern_sse_create_is_rejected_before_persistence(explicit_transport: bool) -> None:
with pytest.raises(ValidationError, match="Modern MCP requires HTTP or stdio"):
NewMCPServerRequest.model_validate({
"url": "https://upstream.example/sse",
"mcp_info": {"protocol_version": "2026-07-28"},
**({"transport": "sse"} if explicit_transport else {}),
})
@pytest.mark.parametrize("metadata", [False, True])
def test_modern_sse_runtime_configuration_is_rejected(metadata: bool) -> None:
with pytest.raises(ValidationError, match="Modern MCP requires HTTP or stdio"):
MCPServer.model_validate({
"server_id": "modern", "name": "modern", "transport": "sse",
**({"mcp_info": {"protocol_version": "2026-07-28"}} if metadata else {"protocol_version": "2026-07-28"}),
})
@pytest.mark.parametrize("transport,version", [("http", "2026-07-28"), ("stdio", "2026-07-28"), ("sse", "2025-11-25"), ("sse", "auto")])
def test_supported_protocol_transport_configurations_remain_valid(transport: str, version: str, monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("LITELLM_ENABLE_MCP_STDIO", "true")
payload: Final = NewMCPServerRequest.model_validate({
"transport": transport, "url": "https://upstream.example/mcp", "command": "python", "args": ["peer.py"],
"mcp_info": {"protocol_version": version},
})
assert payload.transport == transport
assert payload.mcp_info == {"protocol_version": version}
@pytest.mark.asyncio
@pytest.mark.parametrize("protocol_only", [False, True])
async def test_modern_sse_partial_update_rejected_without_writes(protocol_only: bool) -> None:
old_record: Final = LiteLLM_MCPServerTable(
server_id="srv-1", transport="sse" if protocol_only else "http",
mcp_info={"protocol_version": "auto" if protocol_only else "2026-07-28"},
)
payload: Final = UpdateMCPServerRequest.model_validate({
"server_id": "srv-1",
**({"mcp_info": {"protocol_version": "2026-07-28"}} if protocol_only else {"transport": "sse", "url": "https://upstream.example/sse"}),
})
update_mock: Final = AsyncMock(side_effect=HTTPException(status_code=418, detail="Unexpected persistence"))
p1, p2, p3, p4, p5 = _edit_endpoint_patches(old_record, update_mock)
with p1, p2, p3, p4, p5:
with pytest.raises(HTTPException) as error:
await mgmt_endpoints.edit_mcp_server(payload=payload, user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN))
assert error.value.status_code == 400
assert "Modern MCP requires HTTP or stdio" in str(error.value.detail)
update_mock.assert_not_awaited()
@pytest.mark.asyncio
async def test_protocol_partial_update_fails_closed_when_stored_configuration_is_unreadable() -> None:
update_mock: Final = AsyncMock(side_effect=HTTPException(status_code=418, detail="Unexpected persistence"))
p1, p2, p3, p4, p5 = _edit_endpoint_patches(RuntimeError("db unavailable"), update_mock)
with p1, p2, p3, p4, p5:
with pytest.raises(HTTPException) as error:
await mgmt_endpoints.edit_mcp_server(
payload=UpdateMCPServerRequest(server_id="srv-1", mcp_info={"protocol_version": "2026-07-28"}),
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN),
)
assert error.value.status_code == 503
update_mock.assert_not_awaited()
def test_modern_sse_complete_update_is_rejected() -> None:
with pytest.raises(ValidationError, match="Modern MCP requires HTTP or stdio"):
UpdateMCPServerRequest(
server_id="server", transport=MCPTransport.sse, url="https://upstream.example/sse",
mcp_info={"protocol_version": "2026-07-28"},
)
@pytest.mark.asyncio
async def test_protocol_update_on_missing_server_preserves_not_found() -> None:
update_mock: Final = AsyncMock(return_value=None)
p1, p2, p3, p4, p5 = _edit_endpoint_patches(None, update_mock)
with p1, p2, p3, p4, p5:
with pytest.raises(HTTPException) as error:
await mgmt_endpoints.edit_mcp_server(
payload=UpdateMCPServerRequest(server_id="missing", mcp_info={"protocol_version": "2026-07-28"}),
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN),
)
assert error.value.status_code == 404

View file

@ -397,7 +397,7 @@ def test_mcp_advertised_versions_reject_unavailable_revisions(versions):
ConfigGeneralSettings(mcp_advertised_versions=versions)
@pytest.mark.parametrize("revision", ["2026-07-28", "unknown", None])
@pytest.mark.parametrize("revision", ["unknown", None])
def test_mcp_metadata_rejects_unavailable_upstream_protocol(revision):
from litellm.proxy._types import NewMCPServerRequest, UpdateMCPServerRequest
@ -466,3 +466,13 @@ def test_an_http_mcp_server_is_unaffected_by_the_stdio_flag(monkeypatch, request
def test_a_non_mapping_mcp_server_payload_gets_a_validation_error(request_model):
with pytest.raises(ValidationError, match="valid dictionary"):
request_model.model_validate("not-a-server")
@pytest.mark.parametrize("request_model", MCP_SERVER_REQUESTS)
def test_modern_http_upstream_protocol_is_available(request_model):
parsed = request_model.model_validate({
"server_id": "modern", "transport": "http", "url": "https://example.com/mcp",
"mcp_info": {"protocol_version": "2026-07-28"},
})
assert parsed.mcp_info["protocol_version"] == "2026-07-28"
assert parsed.transport == "http"