From 04bc354525142dc4fbf4c9fde094d65b9ef2f7ed Mon Sep 17 00:00:00 2001 From: joshua-berri Date: Fri, 2 Oct 2026 16:04:11 -0700 Subject: [PATCH] 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> --- litellm/experimental_mcp_client/client.py | 20 +- .../_experimental/mcp_server/contracts.py | 18 +- litellm/proxy/_experimental/mcp_server/db.py | 43 +- .../mcp_server/mcp_server_manager.py | 368 ++---------------- .../mcp_server/server_resolution.py | 39 ++ .../_experimental/mcp_server/upstream.py | 315 +++++++++++++++ litellm/proxy/_types.py | 25 ++ .../mcp_management_endpoints.py | 55 ++- litellm/types/mcp.py | 9 +- .../types/mcp_server/mcp_server_manager.py | 6 +- tests/integration/mcp/test_mcp_management.py | 94 +++++ tests/integration/mcp/test_mcp_transports.py | 49 +++ .../test_mcp_client.py | 94 ++++- .../mcp_server/test_mcp_logging.py | 6 +- .../mcp_server/test_mcp_partial_update.py | 39 +- .../mcp_server/test_mcp_server.py | 22 +- .../mcp_server/test_mcp_server_manager.py | 76 +++- .../mcp_server/test_openapi_tool_auth.py | 25 ++ .../mcp_server/test_operations.py | 2 +- .../mcp_server/test_server_resolution.py | 93 +++++ .../test_mcp_management_endpoints.py | 86 ++++ tests/unit/proxy/test__types.py | 12 +- 22 files changed, 1115 insertions(+), 381 deletions(-) create mode 100644 litellm/proxy/_experimental/mcp_server/upstream.py diff --git a/litellm/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py index f133e4837a6..bf6780c7812 100644 --- a/litellm/experimental_mcp_client/client.py +++ b/litellm/experimental_mcp_client/client.py @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/contracts.py b/litellm/proxy/_experimental/mcp_server/contracts.py index b55034e6dc9..82b1861c1e8 100644 --- a/litellm/proxy/_experimental/mcp_server/contracts.py +++ b/litellm/proxy/_experimental/mcp_server/contracts.py @@ -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: diff --git a/litellm/proxy/_experimental/mcp_server/db.py b/litellm/proxy/_experimental/mcp_server/db.py index 481ebdb1ef2..bf1f6fa90c1 100644 --- a/litellm/proxy/_experimental/mcp_server/db.py +++ b/litellm/proxy/_experimental/mcp_server/db.py @@ -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, ) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 375c010a9a8..816ae9aea59 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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 ``).""" - 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, diff --git a/litellm/proxy/_experimental/mcp_server/server_resolution.py b/litellm/proxy/_experimental/mcp_server/server_resolution.py index 8168fea9068..54fd17fa280 100644 --- a/litellm/proxy/_experimental/mcp_server/server_resolution.py +++ b/litellm/proxy/_experimental/mcp_server/server_resolution.py @@ -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, + ) diff --git a/litellm/proxy/_experimental/mcp_server/upstream.py b/litellm/proxy/_experimental/mcp_server/upstream.py new file mode 100644 index 00000000000..89f433093e6 --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/upstream.py @@ -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 ``).""" + 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, + ), + ) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index da69b77286e..2084f6b6ee3 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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): diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 8ed8ec8752e..0d358302fa2 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -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, diff --git a/litellm/types/mcp.py b/litellm/types/mcp.py index fec5e84c8df..da7401e2a2e 100644 --- a/litellm/types/mcp.py +++ b/litellm/types/mcp.py @@ -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, diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index 91ae95eff48..2ee19b3e59a 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -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") diff --git a/tests/integration/mcp/test_mcp_management.py b/tests/integration/mcp/test_mcp_management.py index 67cdbbff5a4..bef882eae31 100644 --- a/tests/integration/mcp/test_mcp_management.py +++ b/tests/integration/mcp/test_mcp_management.py @@ -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 diff --git a/tests/integration/mcp/test_mcp_transports.py b/tests/integration/mcp/test_mcp_transports.py index 16004cdd501..b448b8b4a64 100644 --- a/tests/integration/mcp/test_mcp_transports.py +++ b/tests/integration/mcp/test_mcp_transports.py @@ -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", {}) diff --git a/tests/unit/experimental_mcp_client/test_mcp_client.py b/tests/unit/experimental_mcp_client/test_mcp_client.py index 508ab447326..015d12c3d5e 100644 --- a/tests/unit/experimental_mcp_client/test_mcp_client.py +++ b/tests/unit/experimental_mcp_client/test_mcp_client.py @@ -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) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_logging.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_logging.py index 41d0e2cb59b..44ba40afdd1 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_logging.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_logging.py @@ -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 diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_partial_update.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_partial_update.py index af4f4cbeb17..a3c52dc16b7 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_partial_update.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_partial_update.py @@ -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() diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py index f8bf72428aa..d3679506a2f 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py @@ -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 diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 46264de738a..4b460fc67ae 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -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" diff --git a/tests/unit/proxy/_experimental/mcp_server/test_openapi_tool_auth.py b/tests/unit/proxy/_experimental/mcp_server/test_openapi_tool_auth.py index 15d3b67e641..a8ec7be55f0 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_openapi_tool_auth.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_openapi_tool_auth.py @@ -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"} diff --git a/tests/unit/proxy/_experimental/mcp_server/test_operations.py b/tests/unit/proxy/_experimental/mcp_server/test_operations.py index bb900de4f98..b16b27ac919 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_operations.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_operations.py @@ -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( diff --git a/tests/unit/proxy/_experimental/mcp_server/test_server_resolution.py b/tests/unit/proxy/_experimental/mcp_server/test_server_resolution.py index f88088a4fd8..853118b8dc2 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_server_resolution.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_server_resolution.py @@ -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()) diff --git a/tests/unit/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/unit/proxy/management_endpoints/test_mcp_management_endpoints.py index 88ee36a0fef..2cbf1d578b2 100644 --- a/tests/unit/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -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 diff --git a/tests/unit/proxy/test__types.py b/tests/unit/proxy/test__types.py index 50c3eb2908c..70c5a153647 100644 --- a/tests/unit/proxy/test__types.py +++ b/tests/unit/proxy/test__types.py @@ -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"