From a4fd58501a9b45fdb509b7b8ee21ec38949ba84d Mon Sep 17 00:00:00 2001 From: joshua-berri Date: Fri, 9 Oct 2026 03:21:07 -0700 Subject: [PATCH] fix(mcp): never forward the caller's LiteLLM key to MCP servers (#45401) * fix(mcp): never forward the caller's LiteLLM key to MCP servers Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(mcp): tidy caller-key scrub Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(mcp): avoid master-key scanner false positive Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): match admission when scrubbing the caller key Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(mcp): drop header copy churn Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): normalize bearer variants in caller-key match Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): inject admission header name and cover provider keys Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(mcp): exercise real REST header extraction Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(mcp): cover caller-key scrub on the real REST path Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): exclude gateway admission header from upstream forwarding * fix(mcp): skip empty admission headers when scrubbing keys * fix(mcp): use scrubbed server credentials for auth probes * fix(mcp): match probe credential precedence to upstream calls * test(mcp): simplify probe header assertion --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> --- .../mcp_server/auth/user_api_key_auth_mcp.py | 93 +++- .../mcp_server/mcp_server_manager.py | 4 +- .../_experimental/mcp_server/operations.py | 2 +- .../mcp_server/rest_endpoints.py | 84 ++-- .../proxy/_experimental/mcp_server/server.py | 85 +++- tests/integration/_support/mcp.py | 12 +- tests/integration/mcp/test_mcp_credentials.py | 155 ++++++ .../auth/test_user_api_key_auth_mcp.py | 470 +++++++++++++++++- .../mcp_server/test_mcp_server.py | 14 +- .../mcp_server/test_mcp_server_manager.py | 14 + .../test_mcp_server_tool_calls_and_headers.py | 162 +++++- .../mcp_server/test_operations.py | 17 + .../mcp_server/test_rest_endpoints.py | 60 +++ 13 files changed, 1093 insertions(+), 79 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index ab14a317a66..87dab4a8d67 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -1,4 +1,5 @@ import re +import secrets from collections.abc import Mapping, Sequence from collections.abc import Set as AbstractSet from dataclasses import dataclass @@ -7,6 +8,7 @@ from types import MappingProxyType from typing import TYPE_CHECKING, Final, Literal, cast from fastapi import HTTPException +from pydantic import TypeAdapter from starlette.datastructures import Headers from starlette.requests import Request from starlette.types import Scope @@ -16,6 +18,7 @@ import litellm from litellm._internal_context import with_service_target from litellm._logging import verbose_logger from litellm.constants import MCP_ALL_TOOLS_WILDCARD +from litellm.experimental_mcp_client.client import strip_auth_scheme from litellm.proxy._experimental.mcp_server.catalog import global_manager from litellm.proxy._experimental.mcp_server.oauth_utils import ( get_passthrough_resource_metadata_url, @@ -84,6 +87,11 @@ if TYPE_CHECKING: _EMPTY_TOOLSET_GRANTS: Final[Mapping[str, Sequence[str]]] = MappingProxyType({}) +OPTIONAL_STRING_ADAPTER: Final[TypeAdapter[str | None]] = TypeAdapter(str | None) + + +def _normalize_caller_admission_credential(value: str) -> str: + return _get_bearer_token_or_received_api_key(strip_auth_scheme(value, "Bearer")).strip() def _as_list(values: Sequence[str] | None) -> list[str] | None: # mutable-ok: resolver returns a list @@ -623,22 +631,30 @@ class MCPRequestHandler: bearer_presented=False, ) - # Leak-defense (single chokepoint): a gateway admission credential (session bearer or bridge - # envelope) is NEVER a valid upstream token. Scrub it from EVERY egress context so no - # client-forwarded, OBO, or passthrough path can send it upstream for replay. Anchored to the - # credential SHAPE, so a legitimate upstream/passthrough token is forwarded unchanged. + # Scrub gateway-shaped credentials and the exact LiteLLM credential that admitted this request. raw_headers = dict(headers) + from litellm.proxy.proxy_server import general_settings + + custom_key_header_name: Final = OPTIONAL_STRING_ADAPTER.validate_python( + general_settings.get("litellm_key_header_name") + ) + admitted_credential: Final = MCPRequestHandler.caller_admission_credential( + headers, + validated_user_api_key_auth, + custom_key_header_name=custom_key_header_name, + ) ( oauth2_headers, raw_headers, mcp_auth_header, mcp_server_auth_headers, - ) = MCPRequestHandler._scrub_gateway_admission_credentials( + ) = MCPRequestHandler.scrub_gateway_admission_credentials( admitted=is_mcp_admitted_user_subject(validated_user_api_key_auth), oauth2_headers=oauth2_headers, raw_headers=raw_headers, mcp_auth_header=mcp_auth_header, mcp_server_auth_headers=mcp_server_auth_headers, + admitted_credential=admitted_credential, ) return ( @@ -658,19 +674,36 @@ class MCPRequestHandler: return value is not None and (is_session_bearer_shaped(value) or is_bridge_envelope_shaped(value)) @staticmethod - def _scrub_gateway_admission_credentials( + def _is_caller_admission_key(value: str | None, admitted_credential: str | None) -> bool: + if value is None or admitted_credential is None: + return False + normalized_value: Final = _normalize_caller_admission_credential(value) + normalized_admitted_credential: Final = _normalize_caller_admission_credential(admitted_credential) + return bool( + normalized_value + and normalized_admitted_credential + and secrets.compare_digest(normalized_value.encode(), normalized_admitted_credential.encode()) + ) + + @staticmethod + def scrub_gateway_admission_credentials( admitted: bool, oauth2_headers: dict[str, str] | None, raw_headers: dict[str, str], mcp_auth_header: str | None, mcp_server_auth_headers: dict[str, dict[str, str]] | None, + *, + admitted_credential: str | None, ) -> tuple[dict[str, str] | None, dict[str, str], str | None, dict[str, dict[str, str]] | None]: - """Remove any gateway admission credential from EVERY egress header context, keyed on the credential - SHAPE: top-level ``Authorization`` (oauth2 + raw), the deprecated ``x-mcp-auth``, and per-server - ``x-mcp-{alias}-authorization``. A legitimate upstream/passthrough token is never gateway-shaped so - it survives (including the real upstream token the bridge arm injects per-server); an admitted - subject's top-level Authorization is dropped unconditionally as defense-in-depth.""" - cred: Final = MCPRequestHandler._is_gateway_admission_credential + """Remove gateway-shaped credentials and the caller's admitted credential from EVERY egress + header context: top-level ``Authorization`` (oauth2 + raw), the deprecated ``x-mcp-auth``, and + per-server ``x-mcp-{alias}-authorization``. Preserve the raw ``x-litellm-api-key`` admission + header; an admitted subject's top-level Authorization is otherwise dropped unconditionally.""" + + def cred(value: str | None) -> bool: + return MCPRequestHandler._is_gateway_admission_credential( + value + ) or MCPRequestHandler._is_caller_admission_key(value, admitted_credential) # 1. Top-level Authorization → oauth2_headers. authz: Final = oauth2_headers.get("Authorization") if oauth2_headers else None @@ -680,7 +713,9 @@ class MCPRequestHandler: # 2. raw_headers: drop the admitted subject's Authorization, and ANY header whose value is a # gateway credential (covers x-mcp-auth and x-mcp-{alias}-authorization in their raw form). raw_headers = { - k: v for k, v in raw_headers.items() if not ((admitted and k.lower() == "authorization") or cred(v)) + k: v + for k, v in raw_headers.items() + if k.lower() == "x-litellm-api-key" or not ((admitted and k.lower() == "authorization") or cred(v)) } # 3. Deprecated x-mcp-auth value. @@ -697,6 +732,8 @@ class MCPRequestHandler: return oauth2_headers, raw_headers, mcp_auth_header, mcp_server_auth_headers + _scrub_gateway_admission_credentials = scrub_gateway_admission_credentials + @staticmethod def extract_target_server_names_from_path(path: str) -> list[str]: """ @@ -1553,6 +1590,36 @@ class MCPRequestHandler: return None + @staticmethod + def caller_admission_credential( + headers: Headers, + user_api_key_auth: UserAPIKeyAuth, + *, + custom_key_header_name: str | None, + ) -> str | None: + if user_api_key_auth.api_key is None: + return None + admission_header_names: Final = ( + MCPRequestHandler.LITELLM_API_KEY_HEADER_NAME_PRIMARY, + MCPRequestHandler.LITELLM_API_KEY_HEADER_NAME_SECONDARY, + SpecialHeaders.azure_authorization.value, + SpecialHeaders.anthropic_authorization.value, + SpecialHeaders.google_ai_studio_authorization.value, + SpecialHeaders.azure_apim_authorization.value, + ) + litellm_api_key: Final = ( + headers.get(custom_key_header_name) + if custom_key_header_name is not None + else next( + (header_value for header_name in admission_header_names if (header_value := headers.get(header_name))), + None, + ) + ) + if litellm_api_key is None: + return None + admitted_credential: Final = _normalize_caller_admission_credential(litellm_api_key) + return admitted_credential or None + @staticmethod def safe_get_headers_from_scope(scope: Scope) -> Headers: """ diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index aa31799f1aa..7c375969a44 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -1296,7 +1296,7 @@ def _openapi_forwarded_extra_headers( ) forwarded: Final[dict[str, str]] = {} for header_name in mcp_server.extra_headers: - if not isinstance(header_name, str): + if not isinstance(header_name, str) or header_name.lower() == "x-litellm-api-key": continue if skip_caller_authorization and header_name.lower() == "authorization": continue @@ -6272,7 +6272,7 @@ class MCPServerManager: ) for header in mcp_server.extra_headers: - if not isinstance(header, str): + if not isinstance(header, str) or header.lower() == "x-litellm-api-key": continue if header.lower() == "authorization" and strip_caller_authorization: continue diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index 039fbfdb916..d98748612bc 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -949,7 +949,7 @@ def _prepare_mcp_server_headers( ) for header in server.extra_headers: - if not isinstance(header, str): + if not isinstance(header, str) or header.lower() == "x-litellm-api-key": continue if header.lower() == "authorization" and (strip_caller_authorization or withhold_forwarded_authorization): continue diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index ac543e1f041..496c92f88ee 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -332,9 +332,6 @@ if MCP_AVAILABLE: ``mcp_tool_search_enabled``). Kept out of ``call_tool_rest_api`` so that endpoint stays a single dispatch. An upstream 401 raised by the virtual ``mcp_tool_call`` propagates unhandled to the caller's ``except MCPUpstreamAuthError`` relay, the same as the direct call path.""" - from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( - MCPRequestHandler, - ) from litellm.proxy._experimental.mcp_server.tool_search import ( AGENT_SEARCH_TOOL_NAME, DEFAULT_AGENT_SEARCH_TOP_K, @@ -376,8 +373,8 @@ if MCP_AVAILABLE: virtual_mcp_auth_header, virtual_mcp_server_auth_headers, virtual_raw_headers, - ) = _extract_mcp_headers_from_request(request, MCPRequestHandler) - virtual_oauth2_headers: Final = MCPRequestHandler.get_oauth2_headers_from_headers(request.headers) + virtual_oauth2_headers, + ) = _extract_mcp_headers_from_request(request, user_api_key_dict) if tool_name == MCP_TOOL_SEARCH_TOOL_NAME: return await handle_mcp_tool_search( query=tool_arguments.get("query", ""), @@ -587,19 +584,49 @@ if MCP_AVAILABLE: def _extract_mcp_headers_from_request( request: Request, - mcp_request_handler_cls, - ) -> tuple[str | None, dict[str, dict[str, str]], dict[str, str]]: + user_api_key_dict: UserAPIKeyAuth, + ) -> tuple[ # mutable-ok: preserve the existing plain-dict header contract for MCP helpers + str | None, + dict[str, dict[str, str]] | None, + dict[str, str], + dict[str, str] | None, + ]: """ Extract MCP auth headers from HTTP request. Returns: - Tuple of (mcp_auth_header, mcp_server_auth_headers, raw_headers) + Tuple of (mcp_auth_header, mcp_server_auth_headers, raw_headers, oauth2_headers) """ + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( + OPTIONAL_STRING_ADAPTER, + MCPRequestHandler, + ) + from litellm.proxy.proxy_server import general_settings + headers: Final = request.headers - raw_headers: Final = dict(headers) - mcp_auth_header: Final = mcp_request_handler_cls._get_mcp_auth_header_from_headers(headers) - mcp_server_auth_headers: Final = mcp_request_handler_cls._get_mcp_server_auth_headers_from_headers(headers) - return mcp_auth_header, mcp_server_auth_headers, raw_headers + custom_key_header_name: Final = OPTIONAL_STRING_ADAPTER.validate_python( + general_settings.get("litellm_key_header_name") + ) + admitted_credential: Final = MCPRequestHandler.caller_admission_credential( + headers, + user_api_key_dict, + custom_key_header_name=custom_key_header_name, + ) + oauth2_headers_from_request: Final = MCPRequestHandler.get_oauth2_headers_from_headers(headers) + ( + scrubbed_oauth2_headers, + raw_headers, + mcp_auth_header, + mcp_server_auth_headers, + ) = MCPRequestHandler.scrub_gateway_admission_credentials( + admitted=False, + admitted_credential=admitted_credential, + oauth2_headers=oauth2_headers_from_request, + raw_headers=dict(headers), + mcp_auth_header=MCPRequestHandler.get_mcp_auth_header_from_headers(headers), + mcp_server_auth_headers=MCPRequestHandler.get_mcp_server_auth_headers_from_headers(headers), + ) + return mcp_auth_header, mcp_server_auth_headers, raw_headers, scrubbed_oauth2_headers def _resolve_mcp_server_id_for_rest( server_id: str, @@ -790,11 +817,10 @@ if MCP_AVAILABLE: async def fetch_pinnable_tool_catalog( server: MCPServer, request: Request, user_api_key_dict: UserAPIKeyAuth ) -> dict[str, PinnedMCPTool]: - from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler from litellm.proxy.proxy_server import proxy_logging_obj - mcp_auth_header, mcp_server_auth_headers, raw_headers = _extract_mcp_headers_from_request( - request, MCPRequestHandler + mcp_auth_header, mcp_server_auth_headers, raw_headers, _ = _extract_mcp_headers_from_request( + request, user_api_key_dict ) upstream: Final = await _list_server_tools( server.model_copy(update={"pinned_tools": None, "tool_name_to_description": None}), @@ -807,7 +833,11 @@ if MCP_AVAILABLE: record_listing=False, ) scan: Final = await scan_tool_descriptions( - apply_description_overrides(upstream, server), server, proxy_logging_obj, user_api_key_dict, raw_headers + apply_description_overrides(upstream, server), + server, + proxy_logging_obj, + user_api_key_dict, + raw_headers, ) pinnable: Final = frozenset(tool.name for tool in scan.served) return { @@ -1002,10 +1032,6 @@ if MCP_AVAILABLE: "message": "Successfully retrieved tools" } """ - from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( - MCPRequestHandler, - ) - reject_disallowed_mcp_client(request.headers, user_api_key_dict) try: mcp_server_name = _as_query_str(mcp_server_name) @@ -1039,11 +1065,9 @@ if MCP_AVAILABLE: "message": "Successfully retrieved tools", } - # Extract auth headers from request - headers: Final = request.headers - raw_headers_from_request: Final = dict(headers) - mcp_auth_header: Final = MCPRequestHandler.get_mcp_auth_header_from_headers(headers) - mcp_server_auth_headers: Final = MCPRequestHandler.get_mcp_server_auth_headers_from_headers(headers) + mcp_auth_header, mcp_server_auth_headers, raw_headers_from_request, _ = _extract_mcp_headers_from_request( + request, user_api_key_dict + ) auth_contexts: Final = await build_effective_auth_contexts(user_api_key_dict) @@ -1069,7 +1093,7 @@ if MCP_AVAILABLE: server_id=server_id, allowed_server_ids=allowed_server_ids, rest_client_ip=_rest_client_ip, - mcp_server_auth_headers=mcp_server_auth_headers, + mcp_server_auth_headers=mcp_server_auth_headers or {}, mcp_auth_header=mcp_auth_header, raw_headers_from_request=raw_headers_from_request, user_api_key_dict=user_api_key_dict, @@ -1200,9 +1224,6 @@ if MCP_AVAILABLE: from fastapi import HTTPException from litellm.exceptions import BlockedPiiEntityError, GuardrailRaisedException - from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( - MCPRequestHandler, - ) from litellm.proxy.common_request_processing import ( ProxyBaseLLMRequestProcessing, ) @@ -1272,7 +1293,8 @@ if MCP_AVAILABLE: mcp_auth_header, mcp_server_auth_headers, raw_headers_from_request, - ) = _extract_mcp_headers_from_request(request, MCPRequestHandler) + caller_oauth2_headers_from_request, + ) = _extract_mcp_headers_from_request(request, user_api_key_dict) if mcp_auth_header: data["mcp_auth_header"] = mcp_auth_header if mcp_server_auth_headers: @@ -1301,7 +1323,7 @@ if MCP_AVAILABLE: if target_server is not None: user_oauth_extra_headers = await _get_user_oauth_extra_headers(target_server, user_api_key_dict) caller_oauth2_headers: Final = ( - MCPRequestHandler.get_oauth2_headers_from_headers(request.headers) + caller_oauth2_headers_from_request if target_server is not None and target_server.auth_type in _CLIENT_FORWARDED_TOKEN_AUTH_TYPES else None ) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index dce792cfee1..bf813eaff2b 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -81,6 +81,7 @@ from litellm.proxy._experimental.mcp_server.utils import ( LITELLM_MCP_SERVER_DESCRIPTION, LITELLM_MCP_SERVER_NAME, LITELLM_MCP_SERVER_VERSION, + merge_mcp_headers, ) from litellm.proxy._types import ( ProxyException, @@ -1941,6 +1942,10 @@ if MCP_AVAILABLE: user_api_key_auth: UserAPIKeyAuth | None, mcp_servers: list[str] | None, client_ip: str | None, + *, + oauth2_headers: Mapping[str, str] | None, + mcp_server_auth_headers: dict[str, dict[str, str]] | None, + raw_headers: Mapping[str, str] | None, ) -> None: """Probe pass-through upstream servers in parallel before the MCP session starts. @@ -1955,8 +1960,7 @@ if MCP_AVAILABLE: Fails-open: network errors are logged and the request is allowed through. """ - forwarded_auth: Final = _get_forwarded_auth_from_scope(scope) - if not forwarded_auth: + if not oauth2_headers and not mcp_server_auth_headers: return # Use the authorized server set, not the raw user-supplied names, so that @@ -1966,27 +1970,48 @@ if MCP_AVAILABLE: mcp_servers=mcp_servers, client_ip=client_ip, ) - passthrough_targets: Final[tuple[tuple[MCPServer, str, str], ...]] = tuple( - (srv, forwarded_auth, srv.name) + prepared_headers: Final = tuple( + ( + srv, + operations._prepare_mcp_server_headers( + server=srv, + mcp_server_auth_headers=mcp_server_auth_headers, + mcp_auth_header=None, + oauth2_headers=dict(oauth2_headers) if oauth2_headers is not None else None, + raw_headers=dict(raw_headers) if raw_headers is not None else None, + user_api_key_auth=user_api_key_auth, + scope_servers=allowed_servers, + ), + ) for srv in allowed_servers - # Restrict to genuine OAuth pass-through servers (auth_type none + - # Authorization in extra_headers). Gateway-managed OAuth2 servers - # must not receive the ``resource_metadata=`` challenge emitted - # below — they require ``authorization_uri=`` pointing at the - # gateway AS metadata. ``is_oauth_passthrough`` already requires - # ``auth_type in (None, MCPAuth.none)``, which is mutually - # exclusive with ``has_client_credentials`` (oauth2 + M2M flow), - # so M2M servers are implicitly excluded here. if srv.is_oauth_passthrough ) - probe_targets: Final = passthrough_targets + probe_headers: Final = tuple( + ( + srv, + merge_mcp_headers( + extra_headers=auth if isinstance(auth, dict) else None, + static_headers=extra, + ), + ) + for srv, (auth, extra) in prepared_headers + ) + probe_targets: Final[tuple[tuple[MCPServer, str, str], ...]] = tuple( + (srv, auth_header, srv.name) + for srv, headers in probe_headers + if ( + auth_header := next( + (value for name, value in (headers or {}).items() if name.lower() == "authorization"), None + ) + ) + ) if not probe_targets: return probe_results: Final = await asyncio.gather( *[_probe_upstream_auth(srv.url or "", auth_header) for srv, auth_header, _ in probe_targets] ) - for (srv, _, challenge_server_name), (probe_status, _) in zip(probe_targets, probe_results): + for (_srv, _, challenge_server_name), (probe_status, _) in zip(probe_targets, probe_results): if probe_status == 401: # Token is missing or expired: keep pass-through clients on the # protected-resource discovery flow so they re-authorize against @@ -2060,7 +2085,17 @@ if MCP_AVAILABLE: allowed_server_ids={target.server_id for target in allowed}, raw_headers=context.raw_headers, ) - await _check_passthrough_upstream_auth(scope, context.user_api_key_auth, authorized_names, context.client_ip) + await _check_passthrough_upstream_auth( + scope, + context.user_api_key_auth, + authorized_names, + context.client_ip, + oauth2_headers=context.oauth2_headers, + mcp_server_auth_headers={key: dict(value) for key, value in context.mcp_server_auth_headers.items()} + if context.mcp_server_auth_headers is not None + else None, + raw_headers=context.raw_headers, + ) return None async def handle_streamable_http_mcp(scope: Scope, receive: Receive, send: Send) -> None: @@ -2140,7 +2175,15 @@ if MCP_AVAILABLE: # Pre-flight auth check for pass-through servers. Must run after # toolset scoping so the probe list is derived from the fully-authorized # server set, not the raw user-supplied names. - await _check_passthrough_upstream_auth(scope, user_api_key_auth, mcp_servers, _client_ip) + await _check_passthrough_upstream_auth( + scope, + user_api_key_auth, + mcp_servers, + _client_ip, + oauth2_headers=oauth2_headers, + mcp_server_auth_headers=mcp_server_auth_headers, + raw_headers=raw_headers, + ) # Inject masked debug headers when client sends x-litellm-mcp-debug: true _debug_headers: Final = MCPDebug.maybe_build_debug_headers( @@ -2513,7 +2556,15 @@ if MCP_AVAILABLE: # being stuck with a silently empty tool list. Must run after # toolset scoping so the probe list is derived from the fully- # authorized server set, not the raw user-supplied names. - await _check_passthrough_upstream_auth(scope, user_api_key_auth, mcp_servers, _sse_client_ip) + await _check_passthrough_upstream_auth( + scope, + user_api_key_auth, + mcp_servers, + _sse_client_ip, + oauth2_headers=oauth2_headers, + mcp_server_auth_headers=mcp_server_auth_headers, + raw_headers=raw_headers, + ) set_auth_context( user_api_key_auth=user_api_key_auth, mcp_auth_header=mcp_auth_header, diff --git a/tests/integration/_support/mcp.py b/tests/integration/_support/mcp.py index 17b1e74124c..68bffc85f40 100644 --- a/tests/integration/_support/mcp.py +++ b/tests/integration/_support/mcp.py @@ -427,10 +427,18 @@ def tool_names(gateway: Gateway, key: str, identity: str) -> dict[str, str]: } -def call_tool(gateway: Gateway, key: str, identity: str, name: str, arguments: dict[str, object]) -> httpx.Response: +def call_tool( + gateway: Gateway, + key: str, + identity: str, + name: str, + arguments: dict[str, object], + *, + headers: Mapping[str, str] | None = None, +) -> httpx.Response: return gateway.client.post( "/mcp-rest/tools/call", - headers={"x-litellm-api-key": key}, + headers={"x-litellm-api-key": key, **(headers or {})}, json={"server_id": identity, "name": name, "arguments": arguments}, ) diff --git a/tests/integration/mcp/test_mcp_credentials.py b/tests/integration/mcp/test_mcp_credentials.py index 16dcaab274a..da952b216d5 100644 --- a/tests/integration/mcp/test_mcp_credentials.py +++ b/tests/integration/mcp/test_mcp_credentials.py @@ -18,6 +18,7 @@ from integration._support.mcp import ( ScriptedTool, call_tool, mcp_peer, + openapi_peer, register_mcp, scripted_peer, text_result, @@ -130,6 +131,115 @@ def test_caller_headers_for_other_servers_and_unknown_headers_never_reach_the_pe assert tool_calls(other.drain()) == () +@pytest.mark.parametrize( + ("auth_type", "per_server_header"), + (("none", True), ("true_passthrough", False)), +) +def test_callers_own_litellm_key_never_reaches_the_peer_over_rest( + gateway: Gateway, auth_type: str, per_server_header: bool +) -> None: + with mcp_peer() as peer, gateway.scenario() as scenario: + alias: Final = "cred" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias, auth_type=auth_type) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + caller_headers: Final = { + "x-litellm-api-key": f"Bearer {key}", + (f"x-mcp-{alias}-authorization" if per_server_header else "Authorization"): f"Bearer {key}", + } + peer.drain() + response: Final = call_tool( + gateway, key, identity, tool_names(gateway, key, identity)["add"], ADD, headers=caller_headers + ) + assert response.status_code == 200, response.text + assert response.json()["isError"] is False + observed: Final = peer.drain() + calls: Final = tool_calls(observed) + assert len(calls) == 1, calls + header_sets: Final = tuple( + TypeAdapter(dict[bytes, bytes]).validate_python(request["headers"]) for request in observed + ) + assert all(all(key.encode() not in value for value in headers.values()) for headers in header_sets), ( + "caller key appeared in the recorded MCP header set" + ) + + +@pytest.mark.parametrize("entry", ("server_mcp", "rest")) +def test_empty_primary_header_does_not_expose_authorization_key(gateway: Gateway, entry: EntryPoint) -> None: + with mcp_peer() as peer, gateway.scenario() as scenario: + alias: Final = "cred" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias, extra_headers=["x-upstream-token", "x-tenant"]) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + caller: Final = McpCaller( + gateway, + key, + entry, + alias, + headers={ + "x-litellm-api-key": "", + "Authorization": f"Bearer {key}", + f"x-mcp-{alias}-authorization": f"Bearer {key}", + "x-upstream-token": key, + "x-tenant": "tenant-control", + }, + ) + peer.drain() + outcome: Final = caller.call(f"{alias}-add", ADD, identity if entry == "rest" else None) + assert outcome.ok, outcome.raw + observed: Final = peer.drain() + calls: Final = tool_calls(observed) + assert len(calls) == 1, calls + assert _header(calls[0], b"x-tenant") == b"tenant-control" + header_sets: Final = tuple( + TypeAdapter(dict[bytes, bytes]).validate_python(request["headers"]) for request in observed + ) + assert all(all(key.encode() not in value for value in headers.values()) for headers in header_sets), ( + "empty primary header prevented scrubbing the admitted Authorization key" + ) + + +@pytest.mark.parametrize("entry", ENTRY_POINTS) +def test_extra_headers_cannot_forward_gateway_admission_key(gateway: Gateway, entry: EntryPoint) -> None: + with mcp_peer() as peer, gateway.scenario() as scenario: + alias: Final = "cred" + uuid.uuid4().hex[:8] + identity: Final = register_mcp( + scenario, peer, alias, extra_headers=["X-LiteLLM-API-Key", "x-tenant"], + ) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + caller: Final = McpCaller(gateway, key, entry, alias, headers={"x-tenant": "tenant-control"}) + peer.drain() + outcome: Final = caller.call(f"{alias}-add", ADD, identity if entry != "server_mcp" else None) + assert outcome.ok, outcome.raw + observed: Final = peer.drain() + calls: Final = tool_calls(observed) + assert len(calls) == 1, calls + assert _header(calls[0], b"x-tenant") == b"tenant-control" + header_sets: Final = tuple( + TypeAdapter(dict[bytes, bytes]).validate_python(request["headers"]) for request in observed + ) + assert all(all(key.encode() not in value for value in headers.values()) for headers in header_sets), ( + "gateway admission key reached the upstream through Extra Headers" + ) + + +@pytest.mark.parametrize("entry", ("server_mcp", "rest")) +def test_openapi_extra_headers_cannot_forward_gateway_admission_key(gateway: Gateway, entry: EntryPoint) -> None: + with openapi_peer() as peer, gateway.scenario() as scenario: + alias: Final = "cred" + uuid.uuid4().hex[:8] + identity: Final = register_mcp( + scenario, peer, alias, extra_headers=["X-LiteLLM-API-Key", "x-tenant"], + ) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + caller: Final = McpCaller(gateway, key, entry, alias, headers={"x-tenant": "tenant-control"}) + peer.drain() + outcome: Final = caller.call(f"{alias}-getpet", {"petId": "7"}, identity if entry == "rest" else None) + assert outcome.ok, outcome.raw + observed: Final = peer.drain() + assert [(request["method"], request["path"]) for request in observed] == [("GET", "/pets/7")] + headers: Final = TypeAdapter(dict[bytes, bytes]).validate_python(observed[0]["headers"]) + assert headers.get(b"x-tenant") == b"tenant-control" + assert all(key.encode() not in value for value in headers.values()), "gateway admission key reached OpenAPI" + + def test_server_scoped_caller_header_reaches_only_its_server(gateway: Gateway) -> None: with mcp_peer() as peer, mcp_peer() as other, gateway.scenario() as scenario: alias: Final = "cred" + uuid.uuid4().hex[:8] @@ -488,3 +598,48 @@ def test_deprecated_string_x_mcp_auth_callers_on_a_user_less_key_own_separate_li ) assert seen == (f"Adds for Bearer {first_token}", f"Adds for Bearer {second_token}"), seen assert tool_calls(peer.drain()) == (), "a blocked probe reached the peer" + + +@pytest.mark.parametrize("entry", ("mcp", "server_mcp", "sse")) +@pytest.mark.parametrize("forwarded_token", (False, True)) +def test_oauth_passthrough_probe_uses_server_token_without_admission_key( + gateway: Gateway, entry: EntryPoint, forwarded_token: bool +) -> None: + with mcp_peer() as peer, gateway.scenario() as scenario: + alias: Final = "probe" + uuid.uuid4().hex[:8] + identity: Final = register_mcp( + scenario, peer, alias, auth_type="none", oauth_passthrough=True, extra_headers=["Authorization"] + ) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + upstream_token: Final = "upstream-" + uuid.uuid4().hex + caller: Final = McpCaller( + gateway, + key, + entry, + alias, + headers={ + "Authorization": f"Bearer {upstream_token if forwarded_token else key}", + f"x-mcp-{alias}-authorization": "Bearer expired-token" + if forwarded_token + else f"Bearer {upstream_token}", + }, + ) + peer.drain() + outcome: Final = caller.call(f"{alias}-add", ADD, identity if entry != "server_mcp" else None) + assert outcome.ok, outcome.raw + observed: Final = peer.drain() + probes: Final = tuple( + request + for request in observed + if TypeAdapter(dict[str, object]).validate_python(request["body"]).get("id") == "litellm-mcp-auth-probe" + ) + assert probes, "expected an upstream initialize authentication probe" + assert all(_header(probe, b"authorization") == f"Bearer {upstream_token}".encode() for probe in probes) + assert len(tool_calls(observed)) == 1 + assert all(_header(request, b"authorization") == f"Bearer {upstream_token}".encode() for request in observed) + header_sets: Final = tuple( + TypeAdapter(dict[bytes, bytes]).validate_python(request["headers"]) for request in observed + ) + assert all(all(key.encode() not in value for value in headers.values()) for headers in header_sets), ( + "upstream authentication probe exposed the gateway admission key" + ) diff --git a/tests/unit/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/unit/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index 281bb8f413b..ad94defebd7 100644 --- a/tests/unit/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/unit/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -2,7 +2,7 @@ import contextlib import json import os from datetime import datetime, timedelta, timezone -from typing import Literal +from typing import Final, Literal from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -9472,7 +9472,7 @@ class TestSessionBearerEgressScrub: Authorization: a session bearer placed in x-mcp-auth OR a per-server x-mcp-{alias}-authorization header is stripped too (the High-severity gap: those were forwarded upstream before).""" sess = "Bearer llm_session_abc" - oauth2, raw, mcp_auth, per_server = MCPRequestHandler._scrub_gateway_admission_credentials( + oauth2, raw, mcp_auth, per_server = MCPRequestHandler.scrub_gateway_admission_credentials( admitted=False, oauth2_headers={"Authorization": sess}, raw_headers={ @@ -9482,6 +9482,7 @@ class TestSessionBearerEgressScrub: }, mcp_auth_header="llm_session_xyz", mcp_server_auth_headers={"github": {"Authorization": "llm_session_ghi"}}, + admitted_credential=None, ) assert oauth2 is None assert "authorization" not in {k.lower() for k in raw} @@ -9492,12 +9493,13 @@ class TestSessionBearerEgressScrub: async def test_scrub_keeps_real_upstream_tokens(self): """A legitimate upstream token is never session-/envelope-shaped, so every context is forwarded unchanged — guards against over-stripping a real credential the caller meant for the upstream.""" - oauth2, raw, mcp_auth, per_server = MCPRequestHandler._scrub_gateway_admission_credentials( + oauth2, raw, mcp_auth, per_server = MCPRequestHandler.scrub_gateway_admission_credentials( admitted=False, oauth2_headers={"Authorization": "Bearer real-upstream-xyz"}, raw_headers={"authorization": "Bearer real-upstream-xyz", "x-mcp-github-authorization": "Bearer gh_real"}, mcp_auth_header="some-api-key-123", mcp_server_auth_headers={"github": {"Authorization": "Bearer gh_real"}}, + admitted_credential=None, ) assert oauth2 == {"Authorization": "Bearer real-upstream-xyz"} assert raw["authorization"] == "Bearer real-upstream-xyz" @@ -9507,17 +9509,477 @@ class TestSessionBearerEgressScrub: async def test_scrub_admitted_drops_authorization_but_keeps_injected_upstream_token(self): """An admitted subject's top-level Authorization is dropped unconditionally, while the real upstream token the bridge arm INJECTS into a per-server header (not gateway-shaped) survives.""" - oauth2, raw, mcp_auth, per_server = MCPRequestHandler._scrub_gateway_admission_credentials( + oauth2, raw, mcp_auth, per_server = MCPRequestHandler.scrub_gateway_admission_credentials( admitted=True, oauth2_headers={"Authorization": "Bearer llm_session_abc"}, raw_headers={"authorization": "Bearer llm_session_abc"}, mcp_auth_header=None, mcp_server_auth_headers={"github": {"Authorization": "Bearer gh_injected_upstream"}}, + admitted_credential=None, ) assert oauth2 is None assert "authorization" not in {k.lower() for k in raw} assert per_server == {"github": {"Authorization": "Bearer gh_injected_upstream"}} + @pytest.mark.parametrize( + "context", + ["per_server", "mcp_auth", "oauth2_authorization", "raw_authorization"], + ) + @pytest.mark.parametrize( + "caller_value", + ["sk-caller-admission-key-123", "Bearer sk-caller-admission-key-123"], + ) + async def test_scrub_removes_caller_admission_key_from_each_egress_context( + self, + context: Literal["per_server", "mcp_auth", "oauth2_authorization", "raw_authorization"], + caller_value: str, + ) -> None: + caller_key: Final = "sk-caller-admission-key-123" + upstream_token: Final = "Bearer real-upstream-token" + oauth2_headers: Final = { + "Authorization": caller_value if context == "oauth2_authorization" else upstream_token + } + raw_headers: Final = { + "x-litellm-api-key": f"Bearer {caller_key}", + "authorization": caller_value if context == "raw_authorization" else upstream_token, + "x-upstream-token": upstream_token, + } + mcp_auth_header: Final = caller_value if context == "mcp_auth" else upstream_token + mcp_server_auth_headers: Final = { + "echo_srv": { + "Authorization": caller_value if context == "per_server" else upstream_token, + } + } + + oauth2, raw, mcp_auth, per_server = MCPRequestHandler.scrub_gateway_admission_credentials( + admitted=False, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + mcp_auth_header=mcp_auth_header, + mcp_server_auth_headers=mcp_server_auth_headers, + admitted_credential=caller_key, + ) + + assert raw["x-litellm-api-key"] == f"Bearer {caller_key}" + assert raw["x-upstream-token"] == upstream_token + if context == "raw_authorization": + assert "authorization" not in raw + else: + assert raw["authorization"] == upstream_token + if context == "oauth2_authorization": + assert oauth2 is None + else: + assert oauth2 == {"Authorization": upstream_token} + if context == "mcp_auth": + assert mcp_auth is None + else: + assert mcp_auth == upstream_token + if context == "per_server": + assert not per_server + else: + assert per_server == {"echo_srv": {"Authorization": upstream_token}} + + async def test_scrub_keeps_non_ascii_per_server_token(self) -> None: + upstream_token: Final = "Bearer t\u00f6ken" + _oauth2, _raw, _mcp_auth, per_server = MCPRequestHandler.scrub_gateway_admission_credentials( + admitted=False, + admitted_credential="sk-caller-admission-key-123", + oauth2_headers=None, + raw_headers={}, + mcp_auth_header=None, + mcp_server_auth_headers={"echo_srv": {"Authorization": upstream_token}}, + ) + + assert per_server == {"echo_srv": {"Authorization": upstream_token}} + + async def test_scrub_keeps_caller_key_when_admission_credential_is_missing(self) -> None: + caller_key: Final = "Bearer sk-caller-admission-key-123" + oauth2, raw, mcp_auth, per_server = MCPRequestHandler.scrub_gateway_admission_credentials( + admitted=False, + oauth2_headers={"Authorization": caller_key}, + raw_headers={"x-litellm-api-key": caller_key, "authorization": caller_key}, + mcp_auth_header=caller_key, + mcp_server_auth_headers={"echo_srv": {"Authorization": caller_key}}, + admitted_credential=None, + ) + + assert oauth2 == {"Authorization": caller_key} + assert raw == {"x-litellm-api-key": caller_key, "authorization": caller_key} + assert mcp_auth == caller_key + assert per_server == {"echo_srv": {"Authorization": caller_key}} + + @pytest.mark.parametrize( + "admission_value,per_server_value", + ( + ("Bearer sk-caller-admission-key-123", "BEARER sk-caller-admission-key-123"), + ("Bearer sk-caller-admission-key-123", "Bearer sk-caller-admission-key-123"), + ("Bearer sk-caller-admission-key-123", "Basic sk-caller-admission-key-123"), + ("Basic sk-caller-admission-key-123", "sk-caller-admission-key-123"), + ("Bearer sk-caller-admission-key-123", "bearer sk-caller-admission-key-123"), + ), + ) + async def test_caller_admission_credential_scrubs_bearer_variants_and_basic( + self, + admission_value: str, + per_server_value: str, + ) -> None: + upstream_authorization: Final = "Bearer unrelated-upstream-token" + headers: Final = Headers( + { + "x-litellm-api-key": admission_value, + "authorization": upstream_authorization, + "x-mcp-auth": per_server_value, + "x-mcp-echo_srv-authorization": per_server_value, + "x-upstream-token": upstream_authorization, + } + ) + user_api_key_auth: Final = UserAPIKeyAuth(api_key="stored-key-hash") + admitted_credential: Final = MCPRequestHandler.caller_admission_credential( + headers, + user_api_key_auth, + custom_key_header_name=None, + ) + result: Final = MCPRequestHandler.scrub_gateway_admission_credentials( + admitted=False, + oauth2_headers={"Authorization": upstream_authorization}, + raw_headers=dict(headers), + mcp_auth_header=per_server_value, + mcp_server_auth_headers={ + "echo_srv": {"Authorization": per_server_value}, + "other_srv": {"Authorization": upstream_authorization}, + }, + admitted_credential=admitted_credential, + ) + + assert result == ( + {"Authorization": upstream_authorization}, + { + "x-litellm-api-key": admission_value, + "authorization": upstream_authorization, + "x-upstream-token": upstream_authorization, + }, + None, + {"other_srv": {"Authorization": upstream_authorization}}, + ) + + @pytest.mark.parametrize( + ("empty_header", "admission_header"), + (("x-litellm-api-key", "authorization"), ("authorization", "api-key")), + ) + async def test_empty_admission_header_falls_back_to_authenticated_key( + self, empty_header: str, admission_header: str + ) -> None: + caller_key: Final = "sk-caller-admission-key-123" + headers: Final = Headers({empty_header: "", admission_header: caller_key}) + credential: Final = MCPRequestHandler.caller_admission_credential( + headers, UserAPIKeyAuth(api_key="stored-key-hash"), custom_key_header_name=None + ) + assert credential == caller_key + result: Final = MCPRequestHandler.scrub_gateway_admission_credentials( + admitted=False, + oauth2_headers={"Authorization": f"Bearer {caller_key}"}, + raw_headers={**dict(headers), "x-upstream-token": caller_key, "x-tenant": "tenant-control"}, + mcp_auth_header=caller_key, + mcp_server_auth_headers={"echo": {"Authorization": caller_key}}, + admitted_credential=credential, + ) + assert result == (None, {empty_header: "", "x-tenant": "tenant-control"}, None, {}) + + async def test_caller_admission_credential_prefers_x_litellm_header(self) -> None: + caller_key: Final = "sk-caller-admission-key-123" + upstream_authorization: Final = "Bearer unrelated-upstream-token" + provider_key: Final = "provider-key-not-admitted" + headers: Final = Headers( + { + "x-litellm-api-key": f"Bearer {caller_key}", + "authorization": upstream_authorization, + SpecialHeaders.azure_authorization.value: provider_key, + "x-mcp-auth": f"Bearer {caller_key}", + "x-mcp-echo_srv-authorization": f"Bearer {caller_key}", + } + ) + admitted_credential: Final = MCPRequestHandler.caller_admission_credential( + headers, + UserAPIKeyAuth(api_key="stored-key-hash"), + custom_key_header_name=None, + ) + result: Final = MCPRequestHandler.scrub_gateway_admission_credentials( + admitted=False, + oauth2_headers={"Authorization": upstream_authorization}, + raw_headers=dict(headers), + mcp_auth_header=f"Bearer {caller_key}", + mcp_server_auth_headers={ + "echo_srv": {"Authorization": f"Bearer {caller_key}"}, + "other_srv": {"Authorization": upstream_authorization}, + }, + admitted_credential=admitted_credential, + ) + + assert result == ( + {"Authorization": upstream_authorization}, + { + "x-litellm-api-key": f"Bearer {caller_key}", + "authorization": upstream_authorization, + "api-key": provider_key, + }, + None, + {"other_srv": {"Authorization": upstream_authorization}}, + ) + + async def test_caller_admission_credential_prefers_authorization_before_provider_headers(self) -> None: + caller_key: Final = "sk-caller-admission-key-123" + provider_key: Final = "provider-key-not-admitted" + headers: Final = Headers( + { + "authorization": f"Bearer {caller_key}", + SpecialHeaders.azure_authorization.value: provider_key, + "x-mcp-auth": f"Bearer {caller_key}", + "x-mcp-echo_srv-authorization": f"Bearer {caller_key}", + } + ) + admitted_credential: Final = MCPRequestHandler.caller_admission_credential( + headers, + UserAPIKeyAuth(api_key="stored-key-hash"), + custom_key_header_name=None, + ) + result: Final = MCPRequestHandler.scrub_gateway_admission_credentials( + admitted=False, + oauth2_headers={"Authorization": f"Bearer {caller_key}"}, + raw_headers=dict(headers), + mcp_auth_header=f"Bearer {caller_key}", + mcp_server_auth_headers={"echo_srv": {"Authorization": f"Bearer {caller_key}"}}, + admitted_credential=admitted_credential, + ) + + assert result == ( + None, + {"api-key": provider_key}, + None, + {}, + ) + + async def test_caller_admission_credential_uses_configured_custom_header(self) -> None: + caller_key: Final = "sk-caller-admission-key-123" + standard_key: Final = "Bearer standard-key-not-admitted" + upstream_authorization: Final = "Bearer unrelated-upstream-token" + headers: Final = Headers( + { + "x-custom-api-key": f"Bearer {caller_key}", + "x-litellm-api-key": standard_key, + "authorization": upstream_authorization, + "x-mcp-auth": f"Bearer {caller_key}", + "x-mcp-echo_srv-authorization": f"Bearer {caller_key}", + } + ) + admitted_credential: Final = MCPRequestHandler.caller_admission_credential( + headers, + UserAPIKeyAuth(api_key="stored-key-hash"), + custom_key_header_name="x-custom-api-key", + ) + result: Final = MCPRequestHandler.scrub_gateway_admission_credentials( + admitted=False, + oauth2_headers={"Authorization": upstream_authorization}, + raw_headers=dict(headers), + mcp_auth_header=f"Bearer {caller_key}", + mcp_server_auth_headers={"echo_srv": {"Authorization": f"Bearer {caller_key}"}}, + admitted_credential=admitted_credential, + ) + + assert result == ( + {"Authorization": upstream_authorization}, + { + "x-litellm-api-key": standard_key, + "authorization": upstream_authorization, + }, + None, + {}, + ) + + async def test_caller_admission_credential_does_not_fall_back_when_custom_header_is_absent(self) -> None: + caller_value: Final = "Bearer sk-caller-admission-key-123" + headers: Final = Headers( + { + "x-litellm-api-key": caller_value, + "authorization": caller_value, + "x-mcp-auth": caller_value, + "x-mcp-echo_srv-authorization": caller_value, + } + ) + admitted_credential: Final = MCPRequestHandler.caller_admission_credential( + headers, + UserAPIKeyAuth(api_key="stored-key-hash"), + custom_key_header_name="x-custom-api-key", + ) + result: Final = MCPRequestHandler.scrub_gateway_admission_credentials( + admitted=False, + oauth2_headers={"Authorization": caller_value}, + raw_headers=dict(headers), + mcp_auth_header=caller_value, + mcp_server_auth_headers={"echo_srv": {"Authorization": caller_value}}, + admitted_credential=admitted_credential, + ) + + assert result == ( + {"Authorization": caller_value}, + { + "x-litellm-api-key": caller_value, + "authorization": caller_value, + "x-mcp-auth": caller_value, + "x-mcp-echo_srv-authorization": caller_value, + }, + caller_value, + {"echo_srv": {"Authorization": caller_value}}, + ) + + async def test_caller_admission_credential_uses_master_key_alias_as_admission_gate(self) -> None: + master_key: Final = "sk-" + "1234" + headers: Final = Headers( + { + "x-litellm-api-key": f"Bearer {master_key}", + "authorization": f"Bearer {master_key}", + "x-mcp-auth": f"Bearer {master_key}", + "x-mcp-echo_srv-authorization": f"Bearer {master_key}", + } + ) + admitted_credential: Final = MCPRequestHandler.caller_admission_credential( + headers, + UserAPIKeyAuth(api_key="litellm_proxy_master_key"), + custom_key_header_name=None, + ) + result: Final = MCPRequestHandler.scrub_gateway_admission_credentials( + admitted=False, + oauth2_headers={"Authorization": f"Bearer {master_key}"}, + raw_headers=dict(headers), + mcp_auth_header=f"Bearer {master_key}", + mcp_server_auth_headers={"echo_srv": {"Authorization": f"Bearer {master_key}"}}, + admitted_credential=admitted_credential, + ) + + assert result == ( + None, + {"x-litellm-api-key": f"Bearer {master_key}"}, + None, + {}, + ) + + @pytest.mark.parametrize( + "provider_header_name", + ( + SpecialHeaders.azure_authorization.value, + SpecialHeaders.anthropic_authorization.value, + SpecialHeaders.google_ai_studio_authorization.value, + SpecialHeaders.azure_apim_authorization.value, + ), + ) + async def test_caller_admission_credential_scrubs_provider_header( + self, + provider_header_name: str, + ) -> None: + provider_key: Final = "provider-admission-key" + headers: Final = Headers( + { + provider_header_name: provider_key, + "x-mcp-auth": f"Bearer {provider_key}", + "x-mcp-echo_srv-authorization": f"Bearer {provider_key}", + "x-upstream-token": "Bearer unrelated-upstream-token", + } + ) + admitted_credential: Final = MCPRequestHandler.caller_admission_credential( + headers, + UserAPIKeyAuth(api_key="stored-key-hash"), + custom_key_header_name=None, + ) + result: Final = MCPRequestHandler.scrub_gateway_admission_credentials( + admitted=False, + oauth2_headers=None, + raw_headers=dict(headers), + mcp_auth_header=f"Bearer {provider_key}", + mcp_server_auth_headers={ + "echo_srv": {"Authorization": f"Bearer {provider_key}"}, + "other_srv": {"Authorization": "Bearer unrelated-upstream-token"}, + }, + admitted_credential=admitted_credential, + ) + + assert result == ( + None, + {"x-upstream-token": "Bearer unrelated-upstream-token"}, + None, + {"other_srv": {"Authorization": "Bearer unrelated-upstream-token"}}, + ) + + async def test_caller_admission_credential_follows_provider_header_precedence(self) -> None: + azure_key: Final = "azure-admission-key" + headers: Final = Headers( + { + SpecialHeaders.azure_authorization.value: azure_key, + SpecialHeaders.anthropic_authorization.value: "anthropic-not-admitted", + SpecialHeaders.google_ai_studio_authorization.value: "google-not-admitted", + SpecialHeaders.azure_apim_authorization.value: "apim-not-admitted", + "x-mcp-echo_srv-authorization": f"Bearer {azure_key}", + } + ) + admitted_credential: Final = MCPRequestHandler.caller_admission_credential( + headers, + UserAPIKeyAuth(api_key="stored-key-hash"), + custom_key_header_name=None, + ) + result: Final = MCPRequestHandler.scrub_gateway_admission_credentials( + admitted=False, + oauth2_headers=None, + raw_headers=dict(headers), + mcp_auth_header=f"Bearer {azure_key}", + mcp_server_auth_headers={"echo_srv": {"Authorization": f"Bearer {azure_key}"}}, + admitted_credential=admitted_credential, + ) + + assert result == ( + None, + { + "x-api-key": "anthropic-not-admitted", + "x-goog-api-key": "google-not-admitted", + "ocp-apim-subscription-key": "apim-not-admitted", + }, + None, + {}, + ) + + async def test_caller_admission_credential_returns_none_without_admitted_api_key(self) -> None: + caller_value: Final = "Bearer sk-caller-admission-key-123" + headers: Final = Headers( + { + "x-litellm-api-key": caller_value, + "authorization": caller_value, + "x-mcp-auth": caller_value, + "x-mcp-echo_srv-authorization": caller_value, + } + ) + admitted_credential: Final = MCPRequestHandler.caller_admission_credential( + headers, + UserAPIKeyAuth(api_key=None), + custom_key_header_name=None, + ) + result: Final = MCPRequestHandler.scrub_gateway_admission_credentials( + admitted=False, + oauth2_headers={"Authorization": caller_value}, + raw_headers=dict(headers), + mcp_auth_header=caller_value, + mcp_server_auth_headers={"echo_srv": {"Authorization": caller_value}}, + admitted_credential=admitted_credential, + ) + + assert result == ( + {"Authorization": caller_value}, + { + "x-litellm-api-key": caller_value, + "authorization": caller_value, + "x-mcp-auth": caller_value, + "x-mcp-echo_srv-authorization": caller_value, + }, + caller_value, + {"echo_srv": {"Authorization": caller_value}}, + ) + # --------------------------------------------------------------------------- # Internal-user (human) MCP entitlement tests 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 5ae3c595c53..449e437865b 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py @@ -2150,8 +2150,8 @@ async def test_list_tool_rest_api_with_server_specific_auth(): ) as mock_get_server_auth: mock_get_auth.return_value = "Bearer default_token" mock_get_server_auth.return_value = { - "zapier": "Bearer zapier_token", - "slack": "Bearer slack_token", + "zapier": {"Authorization": "Bearer zapier_token"}, + "slack": {"Authorization": "Bearer slack_token"}, } # Mock the global_mcp_server_manager @@ -2218,7 +2218,7 @@ async def test_list_tool_rest_api_with_server_specific_auth(): call_args = mock_get_tools.call_args assert call_args[0][0] == mock_server # server assert ( - call_args[0][1] == "Bearer zapier_token" + call_args[0][1] == {"Authorization": "Bearer zapier_token"} ) # server_auth_header @@ -2342,8 +2342,8 @@ async def test_list_tool_rest_api_all_servers_with_auth(): ) as mock_get_server_auth: mock_get_auth.return_value = "Bearer default_token" mock_get_server_auth.return_value = { - "zapier": "Bearer zapier_token", - "slack": "Bearer slack_token", + "zapier": {"Authorization": "Bearer zapier_token"}, + "slack": {"Authorization": "Bearer slack_token"}, } # Mock the global_mcp_server_manager @@ -2436,10 +2436,10 @@ async def test_list_tool_rest_api_all_servers_with_auth(): } assert ( - server_auth_map.get(mock_zapier_server) == "Bearer zapier_token" + server_auth_map.get(mock_zapier_server) == {"Authorization": "Bearer zapier_token"} ) assert ( - server_auth_map.get(mock_slack_server) == "Bearer slack_token" + server_auth_map.get(mock_slack_server) == {"Authorization": "Bearer slack_token"} ) 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 b67e4d24f47..c1d64bf3d3a 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 @@ -19281,6 +19281,20 @@ async def test_repeated_stale_discovery_uses_current_callers_endpoint(endpoint: assert resolved is replacement +@pytest.mark.parametrize("header_name", ("x-litellm-api-key", "X-LiteLLM-API-Key")) +def test_openapi_extra_headers_exclude_gateway_admission_key(header_name: str) -> None: + caller_key: Final = "Bearer sk-admission-only" + raw_headers: Final = {"x-litellm-api-key": caller_key, "x-tenant": "tenant-control"} + server: Final = MCPServer( + server_id="header-boundary", name="header-boundary", transport=MCPTransport.http, + spec_path="/spec.yaml", auth_type=MCPAuth.none, extra_headers=[header_name, "X-Tenant"], + ) + assert resolve_openapi_tool_auth( + server, None, None, raw_headers, UserAPIKeyAuth(api_key="sk-admission-only"), + ) == (None, {"X-Tenant": "tenant-control"}, None) + assert raw_headers == {"x-litellm-api-key": caller_key, "x-tenant": "tenant-control"} + + @pytest.mark.asyncio async def test_legacy_upstream_elicitation_rejects_modern_downstream_without_consent() -> None: from types import SimpleNamespace diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py index 49d21924d9c..d7deb1cf774 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py @@ -3506,7 +3506,8 @@ async def test_initialize_request_with_existing_session_tracks_new_session(): patch( # test-quality-ok: registry is empty in unit tests; key owns one server "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers", new_callable=AsyncMock, - return_value=[MagicMock()], + return_value=[MCPServer(server_id="new-server", name="new-server", url="http://upstream/mcp", + transport=MCPTransport.http, auth_type=MCPAuth.none)], ), patch( "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", @@ -7114,6 +7115,9 @@ async def test_legacy_delegate_bare_token_is_not_probed_upstream(): # test-qual user_api_key_auth=UserAPIKeyAuth(), mcp_servers=["delegate_test"], client_ip=None, + oauth2_headers=None, + mcp_server_auth_headers=None, + raw_headers=None, ) probe.assert_not_awaited() @@ -7150,6 +7154,9 @@ async def test_legacy_delegate_dual_credentials_are_not_probed_upstream(): # te user_api_key_auth=UserAPIKeyAuth(user_id="admitted-user"), mcp_servers=["delegate_test"], client_ip=None, + oauth2_headers={"Authorization": "Bearer upstream-token"}, + mcp_server_auth_headers=None, + raw_headers={"x-litellm-api-key": "sk-litellm-proxy-key", "Authorization": "Bearer upstream-token"}, ) probe.assert_not_awaited() @@ -7198,6 +7205,9 @@ async def test_oauth_passthrough_preflight_preserves_status_contract(probe_statu user_api_key_auth=UserAPIKeyAuth(user_id="admitted-user"), mcp_servers=["passthrough_server"], client_ip=None, + oauth2_headers={"Authorization": "Bearer upstream-token"}, + mcp_server_auth_headers=None, + raw_headers={"x-litellm-api-key": "sk-litellm-proxy-key", "Authorization": "Bearer upstream-token"}, ) assert result is None else: @@ -7207,6 +7217,9 @@ async def test_oauth_passthrough_preflight_preserves_status_contract(probe_statu user_api_key_auth=UserAPIKeyAuth(user_id="admitted-user"), mcp_servers=["passthrough_server"], client_ip=None, + oauth2_headers={"Authorization": "Bearer upstream-token"}, + mcp_server_auth_headers=None, + raw_headers={"x-litellm-api-key": "sk-litellm-proxy-key", "Authorization": "Bearer upstream-token"}, ) assert exc_info.value.status_code == expected_status if expected_status == 401: @@ -7243,6 +7256,9 @@ async def test_delegate_tokenless_request_not_probed(): user_api_key_auth=UserAPIKeyAuth(), mcp_servers=["delegate_test"], client_ip=None, + oauth2_headers=None, + mcp_server_auth_headers=None, + raw_headers=None, ) probe.assert_not_awaited() @@ -7276,6 +7292,9 @@ async def test_delegate_preflight_skipped_on_multi_server_routes(): user_api_key_auth=UserAPIKeyAuth(), mcp_servers=["delegate_test", "other_server"], client_ip=None, + oauth2_headers=None, + mcp_server_auth_headers=None, + raw_headers=None, ) probe.assert_not_awaited() @@ -7319,6 +7338,9 @@ async def test_bare_authorization_never_probes_passthrough_servers(): user_api_key_auth=UserAPIKeyAuth(), mcp_servers=["pt_server"], client_ip=None, + oauth2_headers=None, + mcp_server_auth_headers=None, + raw_headers=None, ) probe.assert_not_awaited() @@ -7365,6 +7387,9 @@ async def test_delegate_not_probed_when_named_only_via_server_id(): user_api_key_auth=UserAPIKeyAuth(user_id="u1", api_key="hashed-sk"), mcp_servers=["delegate-secret-id"], client_ip=None, + oauth2_headers=None, + mcp_server_auth_headers=None, + raw_headers=None, ) probe.assert_not_awaited() @@ -7398,6 +7423,9 @@ async def test_delegate_probe_not_fanned_out_to_access_group_members(): user_api_key_auth=UserAPIKeyAuth(), mcp_servers=["prod_tools_group"], client_ip=None, + oauth2_headers=None, + mcp_server_auth_headers=None, + raw_headers=None, ) probe.assert_not_awaited() @@ -11265,7 +11293,7 @@ async def test_modern_oauth_challenge_follows_continuation_authorization( [route_name], None, {"Authorization": "Bearer expired-upstream-token"}, - None, + {"Authorization": "Bearer expired-upstream-token"}, ) ), ), @@ -11437,3 +11465,133 @@ async def test_modern_preflight_leaves_invalid_envelopes_to_sdk_without_upstream ) assert result is None resolve.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.usefixtures("httpx_transport") +@pytest.mark.parametrize("binding", ("probe", "PROBE", "shared")) +@pytest.mark.parametrize("forwarded_token", (False, True)) +async def test_modern_passthrough_probe_uses_scrubbed_server_credential(binding: str, forwarded_token: bool): + import respx + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + from litellm.proxy._experimental.mcp_server import server + from litellm.proxy._experimental.mcp_server.contracts import OperationContext + + target: Final = MCPServer( + server_id="probe-id", + name="probe", + alias="probe", + access_groups=["shared"], + url="http://upstream/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + oauth_passthrough=True, + extra_headers=["Authorization"], + ) + scope: Final = _delegate_scope( + [ + (b"x-litellm-api-key", b"sk-admission-key"), + (b"authorization", b"Bearer distinct-upstream-token" if forwarded_token else b"Bearer sk-admission-key"), + (b"mcp-protocol-version", b"2026-07-28"), + (b"mcp-method", b"tools/call"), + (b"mcp-name", b"probe-add"), + ] + ) + context: Final = OperationContext( + _caller=UserAPIKeyAuth( + api_key="hashed-admission-key", + object_permission=LiteLLM_ObjectPermissionTable( + object_permission_id="probe-permission", mcp_servers=[target.server_id] + ), + ), + mcp_servers=("probe",), + mcp_server_auth_headers={ + binding: {"Authorization": "Bearer expired-token" if forwarded_token else "Bearer distinct-upstream-token"} + }, + oauth2_headers={"Authorization": "Bearer distinct-upstream-token"} if forwarded_token else None, + raw_headers={ + "x-litellm-api-key": "sk-admission-key", + **({"Authorization": "Bearer distinct-upstream-token"} if forwarded_token else {}), + }, + ) + body: Final = json.dumps( + { + "jsonrpc": "2.0", + "id": 1, + "method": "tools/call", + "params": { + "name": "probe-add", + "arguments": {"a": 2, "b": 3}, + "_meta": { + "io.modelcontextprotocol/protocolVersion": "2026-07-28", + "io.modelcontextprotocol/clientCapabilities": {}, + }, + }, + } + ).encode() + mcp_operations.global_mcp_server_manager.registry[target.server_id] = target + with respx.mock as upstream: + probe: Final = upstream.post(target.url).respond(200) + result: Final = await server._preflight_modern_interaction(scope, body, context) + assert result is None + assert probe.call_count == 1 + sent_headers: Final = probe.calls[0].request.headers + assert sent_headers["Authorization"] == "Bearer distinct-upstream-token" + assert all("sk-admission-key" not in value for value in sent_headers.values()) + + +@pytest.mark.asyncio +@pytest.mark.usefixtures("httpx_transport") +@pytest.mark.parametrize("allowed", (True, False)) +@pytest.mark.parametrize("header_shape", ("mapping", "non_authorization")) +async def test_passthrough_probes_bind_each_token_to_its_authorized_server(allowed: bool, header_shape: str): + import respx + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + from litellm.proxy._experimental.mcp_server import server + + targets: Final = [ + MCPServer( + server_id=name, + name=name, + alias=name, + url=f"http://{name}/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + oauth_passthrough=True, + extra_headers=["Authorization"], + ) + for name in ("first", "second") + ] + mcp_operations.global_mcp_server_manager.registry.update({target.server_id: target for target in targets}) + with respx.mock(assert_all_called=False) as upstream: + first: Final = upstream.post(targets[0].url).respond(200) + second: Final = upstream.post(targets[1].url).respond(200) + result: Final = await server._check_passthrough_upstream_auth( + raw_headers=None, + scope=_delegate_scope([(b"authorization", b"Bearer sk-admission-key")]), + user_api_key_auth=UserAPIKeyAuth( + api_key="hashed-admission-key", + object_permission=LiteLLM_ObjectPermissionTable( + object_permission_id="probe-permission", + mcp_servers=[target.server_id for target in targets] if allowed else [], + ), + ), + mcp_servers=["first", "second"], + client_ip=None, + oauth2_headers=None, + mcp_server_auth_headers={ + "first": {"x-other" if header_shape == "non_authorization" else "aUtHoRiZaTiOn": "Bearer first-token"}, + "second": { + "x-other" if header_shape == "non_authorization" else "Authorization": "Bearer second-token" + }, + }, + ) + assert result is None + assert tuple( + (str(call.request.url), call.request.headers["Authorization"]) + for call in (*first.calls, *second.calls) + ) == ( + (("http://first/mcp", "Bearer first-token"), ("http://second/mcp", "Bearer second-token")) + if allowed and header_shape != "non_authorization" + else () + ) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_operations.py b/tests/unit/proxy/_experimental/mcp_server/test_operations.py index d9997ec0fce..25c80398e46 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_operations.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_operations.py @@ -1424,3 +1424,20 @@ async def test_initial_tool_listing_preserves_legacy_error_fallback(monkeypatch: assert listing.tools == [] assert listing.next_cursor is None fetch.assert_awaited_once() + + +@pytest.mark.parametrize("header_name", ("x-litellm-api-key", "X-LiteLLM-API-Key")) +def test_discovery_extra_headers_exclude_gateway_admission_key(header_name: str) -> None: + caller_key: Final = "Bearer sk-admission-only" + raw_headers: Final = {"x-litellm-api-key": caller_key, "x-tenant": "tenant-control"} + server: Final = MCPServer( + server_id="header-boundary", name="header-boundary", transport=MCPTransport.http, + url="https://example.invalid/mcp", auth_type=MCPAuth.none, + extra_headers=[header_name, "X-Tenant"], + ) + auth_header, extra_headers = operations._prepare_mcp_server_headers( + server, None, None, None, raw_headers, UserAPIKeyAuth(api_key="sk-admission-only"), + ) + assert auth_header is None + assert extra_headers == {"X-Tenant": "tenant-control"} + assert raw_headers == {"x-litellm-api-key": caller_key, "x-tenant": "tenant-control"} diff --git a/tests/unit/proxy/_experimental/mcp_server/test_rest_endpoints.py b/tests/unit/proxy/_experimental/mcp_server/test_rest_endpoints.py index 791e6d18ea8..fc61ec66b30 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_rest_endpoints.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_rest_endpoints.py @@ -2652,6 +2652,66 @@ class TestCallToolRestAPI: assert captured["oauth2_headers"] is None fire_logging.assert_awaited_once() + async def test_extract_mcp_headers_scrubs_admitted_caller_credential(self) -> None: + caller_key: Final = "sk-rest-caller-admission-key-123" + upstream_token: Final = "Bearer unrelated-upstream-token" + request: Final = _build_request( + headers={ + "x-litellm-api-key": f"Bearer {caller_key}", + "authorization": f"Bearer {caller_key}", + "x-mcp-echo_srv-authorization": f"Bearer {caller_key}", + "x-mcp-upstream-authorization": upstream_token, + } + ) + + extracted_headers: Final = rest_endpoints._extract_mcp_headers_from_request( + request, + UserAPIKeyAuth(api_key="stored-key-hash"), + ) + + assert extracted_headers == ( + None, + {"upstream": {"Authorization": upstream_token}}, + { + "x-litellm-api-key": f"Bearer {caller_key}", + "x-mcp-upstream-authorization": upstream_token, + }, + None, + ) + + async def test_extract_mcp_headers_preserves_caller_credential_without_authenticated_key(self) -> None: + caller_key: Final = "sk-rest-caller-admission-key-123" + caller_authorization: Final = f"Bearer {caller_key}" + upstream_token: Final = "Bearer unrelated-upstream-token" + request: Final = _build_request( + headers={ + "x-litellm-api-key": caller_authorization, + "authorization": caller_authorization, + "x-mcp-echo_srv-authorization": caller_authorization, + "x-mcp-upstream-authorization": upstream_token, + } + ) + + extracted_headers: Final = rest_endpoints._extract_mcp_headers_from_request( + request, + UserAPIKeyAuth(), + ) + + assert extracted_headers == ( + None, + { + "echo_srv": {"Authorization": caller_authorization}, + "upstream": {"Authorization": upstream_token}, + }, + { + "x-litellm-api-key": caller_authorization, + "authorization": caller_authorization, + "x-mcp-echo_srv-authorization": caller_authorization, + "x-mcp-upstream-authorization": upstream_token, + }, + {"Authorization": caller_authorization}, + ) + @pytest.mark.parametrize( ("structured", "expected_structured", "expected_texts"), [