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>
This commit is contained in:
joshua-berri 2026-10-09 03:21:07 -07:00 • committed by GitHub
parent a3236ba947
commit a4fd58501a
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
13 changed files with 1093 additions and 79 deletions

View file

@ -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:
"""

View file

@ -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

View file

@ -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

View file

@ -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
)

View file

@ -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,

View file

@ -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},
)

View file

@ -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"
)

View file

@ -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

View file

@ -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"}
)

View file

@ -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

View file

@ -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 ()
)

View file

@ -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"}

View file

@ -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"),
[