mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
a3236ba947
commit
a4fd58501a
13 changed files with 1093 additions and 79 deletions
|
|
@ -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:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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},
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 ()
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
[
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue