Merge pull request #41364 from BerriAI/litellm_fix_mcp_auth_fail_closed_4501

fix(mcp): fail closed on missing upstream credentials
This commit is contained in:
joshua-berri 2026-09-16 23:26:38 +00:00 • committed by GitHub
commit 41410e9556
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
8 changed files with 723 additions and 58 deletions

View file

@ -346,13 +346,17 @@ class MCPClient:
self.update_auth_value(auth_value)
async def discovery_auth_fingerprint(self) -> str:
return self._hash_discovery_auth(await self.prepare_request_auth())
async def prepare_request_auth(self) -> httpx.Request:
"""Preview the authenticated request without sending it, closing the auth flow afterwards."""
request: Final = httpx.Request("POST", self.server_url or "http://localhost/", headers=self._get_auth_headers())
if self._resolved_auth is None:
return self._hash_discovery_auth(request)
return request
flow: Final = self._resolved_auth.async_auth_flow(request)
try:
authenticated: Final = await flow.__anext__()
return self._hash_discovery_auth(authenticated)
return authenticated
finally:
await flow.aclose()

View file

@ -102,6 +102,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials import (
UpstreamCredentialProvider,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import (
prepare_mcp_client,
raise_public,
raise_token_exchange_challenge,
raise_user_oauth_challenge,
@ -2804,6 +2805,8 @@ class MCPServerManager:
headers=headers,
server_label=server.name or server.server_name or server.alias or server.server_id,
relays_upstream_auth=server.is_client_forwarded_token,
auth_type=server.auth_type,
upstream_token_header=server.upstream_token_header,
)
tool_func.__name__ = prefixed_tool_name
tool_func.__doc__ = description
@ -4259,15 +4262,20 @@ class MCPServerManager:
user_api_key_auth=user_api_key_auth,
extra_headers=extra_headers,
)
return MCPClient(
server_url=server_url,
transport_type=transport,
auth_type=resolved_server.auth_type,
timeout=(resolved_server.timeout if resolved_server.timeout is not None else MCP_CLIENT_TIMEOUT),
extra_headers=extra_headers,
resolved_auth=resolved_auth,
sampling_callback=sampling_cb,
elicitation_callback=elicitation_cb,
return await prepare_mcp_client(
resolved_server,
MCPClient(
server_url=server_url,
transport_type=transport,
auth_type=resolved_server.auth_type,
timeout=(
resolved_server.timeout if resolved_server.timeout is not None else MCP_CLIENT_TIMEOUT
),
extra_headers=extra_headers,
resolved_auth=resolved_auth,
sampling_callback=sampling_cb,
elicitation_callback=elicitation_cb,
),
)
# Create SigV4 auth if configured
@ -4297,17 +4305,20 @@ class MCPServerManager:
else AuthResolution.no_auth
)
record_auth_resolution(server.server_id, legacy_source)
return MCPClient(
server_url=server_url,
transport_type=transport,
auth_type=resolved_server.auth_type,
auth_value=auth_value,
auth_header_name=auth_header_name,
timeout=(resolved_server.timeout if resolved_server.timeout is not None else MCP_CLIENT_TIMEOUT),
extra_headers=extra_headers,
aws_auth=aws_auth,
sampling_callback=sampling_cb,
elicitation_callback=elicitation_cb,
return await prepare_mcp_client(
resolved_server,
MCPClient(
server_url=server_url,
transport_type=transport,
auth_type=resolved_server.auth_type,
auth_value=auth_value,
auth_header_name=auth_header_name,
timeout=(resolved_server.timeout if resolved_server.timeout is not None else MCP_CLIENT_TIMEOUT),
extra_headers=extra_headers,
aws_auth=aws_auth,
sampling_callback=sampling_cb,
elicitation_callback=elicitation_cb,
),
)
async def _get_tools_from_server(

View file

@ -54,7 +54,7 @@ from litellm.llms.custom_httpx.http_handler import (
from litellm.proxy._experimental.mcp_server.tool_registry import (
global_mcp_tool_registry,
)
from litellm.types.mcp import credential_redirect_hook, custom_credential_slot
from litellm.types.mcp import MCPAuthType, credential_redirect_hook, custom_credential_slot
class _OpenAPIJSONSchema(TypedDict, total=False):
@ -471,6 +471,8 @@ def create_tool_function(
headers: dict[str, str] | None = None,
server_label: str | None = None,
relays_upstream_auth: bool = False,
auth_type: MCPAuthType = None,
upstream_token_header: str | None = None,
):
"""Create a tool function for an OpenAPI operation.
@ -503,6 +505,18 @@ def create_tool_function(
by using **kwargs instead of named parameters.
"""
effective_headers: Final = _merge_openapi_tool_request_headers(headers)
if auth_type is not None:
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import (
raise_public,
validate_static_credential,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Error, Ok
match validate_static_credential(auth_type, effective_headers, upstream_token_header):
case Error(error):
raise_public(error)
case Ok():
pass
# Build URL from base_url and path
url = base_url + path

View file

@ -13,15 +13,17 @@ from __future__ import annotations
import base64
import os
from collections.abc import Mapping
from typing import TYPE_CHECKING, Final, Literal, NoReturn
from fastapi import HTTPException
from pydantic import SecretStr
from typing_extensions import assert_never
from litellm.experimental_mcp_client.client import strip_auth_scheme, to_basic_credentials
from litellm.experimental_mcp_client.client import MCPClient, strip_auth_scheme, to_basic_credentials
from litellm.proxy._experimental.mcp_server.exceptions import MCPServerURLCredentialsError
from litellm.proxy._experimental.mcp_server.oauth_utils import resolve_upstream_resource
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Error, Ok, Result
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
DEFAULT_CREDENTIAL_HEADER,
ApiKeyConfig,
@ -39,7 +41,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
Subject,
TokenExchangeConfig,
)
from litellm.types.mcp import DEFAULT_SUBJECT_TOKEN_TYPE, MCPAuth
from litellm.types.mcp import DEFAULT_SUBJECT_TOKEN_TYPE, MCPAuth, MCPAuthType, MCPTransport
if TYPE_CHECKING:
from litellm.proxy._types import UserAPIKeyAuth
@ -79,7 +81,7 @@ def to_server_spec(server: MCPServer) -> ServerSpec | None:
BYOK is the per-user source of the ``api_key`` mode; its scheme rides on ``auth_type`` just
like a shared key, but the value is per-user and not migrated yet, so a BYOK server defers
to v1 regardless of ``auth_type`` (this guard is the seam the BYOK arm replaces later).
to v1 for its static schemes. Declared OBO always stays with the exchange arm.
Dispatches on the declared ``auth_type``. The match is exhaustive over ``MCPAuthType`` with
an ``assert_never`` tail, so a newly added auth mode fails the type gate here until it is
@ -90,8 +92,8 @@ def to_server_spec(server: MCPServer) -> ServerSpec | None:
modes ``true_passthrough`` / ``oauth_delegate`` (``PassthroughConfig``); delegated/passthrough
oauth2 and SigV4 return None and stay on v1.
"""
if server.is_byok:
return None # per-user BYOK source not migrated yet -> defer to v1 (any auth_type)
if server.is_byok and server.auth_type != MCPAuth.oauth2_token_exchange:
return None # per-user BYOK source not migrated yet -> defer to v1
resource: Final = server.url or server.server_id
auth_type: Final = server.auth_type
match auth_type:
@ -165,21 +167,9 @@ def _client_credentials_spec(server: MCPServer, resource: str) -> ServerSpec:
)
def _token_exchange_spec(server: MCPServer, resource: str) -> ServerSpec | None:
"""Build a token_exchange (OBO) spec, or defer (None) when it is not OBO-configured.
An OBO server with ``client_id``/``client_secret`` is owned by the v2 arm even if the
``token_exchange_endpoint``/``token_url`` is absent: a missing endpoint then fails closed (412) at
the exchanger rather than silently deferring to v1 and connecting unauthenticated, since the
gateway must not guess the IdP or fall back to a weaker source. Without client credentials there is
nothing to own, so the server stays on v1 (parity-safe). ``profile`` selects the wire dialect
(``rfc8693`` default, ``entra_obo`` for Microsoft Entra On-Behalf-Of); an unrecognized value
normalizes to ``rfc8693`` so a bad config value cannot crash spec-building. ``audience`` is
forwarded only when the operator set it; a missing one is omitted, not derived.
"""
def _token_exchange_spec(server: MCPServer, resource: str) -> ServerSpec:
"""Keep declared OBO owned by the resolver, including incomplete client configuration."""
endpoint: Final = server.token_exchange_endpoint or server.effective_token_url
if not server.client_id or not server.client_secret:
return None
profile: Final[Literal["rfc8693", "entra_obo"]] = (
"entra_obo" if server.token_exchange_profile == "entra_obo" else "rfc8693"
)
@ -193,7 +183,7 @@ def _token_exchange_spec(server: MCPServer, resource: str) -> ServerSpec | None:
token_exchange_endpoint=endpoint,
audience=server.audience,
client_id=server.client_id,
client_secret=SecretStr(server.client_secret),
client_secret=SecretStr(server.client_secret) if server.client_secret else None,
token_endpoint_auth_method=server.token_endpoint_auth_method,
scopes=tuple(server.scopes or ()),
),
@ -397,3 +387,69 @@ def raise_token_exchange_challenge(
detail="Unauthorized",
headers={"WWW-Authenticate": www_authenticate},
)
_STATIC_MODES: Final = frozenset(
(MCPAuth.api_key, MCPAuth.bearer_token, MCPAuth.basic, MCPAuth.token, MCPAuth.authorization)
)
def _usable_credential_value(auth_type: MCPAuthType, name: str, value: str) -> bool:
if not value:
return False
if auth_type == MCPAuth.api_key and name != "authorization":
return True
if value.lower() in ("bearer", "basic", "token", "apikey"):
return False
if auth_type == MCPAuth.api_key:
api_scheme: Final = value.split(None, 1)[0]
if api_scheme.lower() in ("bearer", "token", "apikey"):
api_credential: Final = strip_auth_scheme(value, api_scheme).strip()
return api_credential.lower() != api_scheme.lower()
if auth_type in (MCPAuth.bearer_token, MCPAuth.token):
scheme: Final = "Bearer" if auth_type == MCPAuth.bearer_token else "token"
credential: Final = strip_auth_scheme(value, scheme).strip()
return bool(credential) and credential.lower() != scheme.lower()
if auth_type == MCPAuth.basic:
parts: Final = value.split(None, 1)
if len(parts) != 2 or parts[0].lower() != "basic":
return False
try:
decoded: Final = base64.b64decode(parts[1], validate=True).strip()
return b":" in decoded
except ValueError:
return False
return True
def validate_static_credential(
auth_type: MCPAuthType,
headers: Mapping[str, str],
upstream_token_header: str | None = None,
) -> Result[None, CredError]:
if auth_type not in _STATIC_MODES:
return Ok(None)
default_slot: Final = "X-API-Key" if auth_type == MCPAuth.api_key else "Authorization"
slots: Final = frozenset(
name.lower()
for name in (
upstream_token_header or default_slot,
default_slot,
"Authorization",
)
)
values: Final = tuple((name.lower(), value.strip()) for name, value in headers.items() if name.lower() in slots)
if any(_usable_credential_value(auth_type, name, value) for name, value in values):
return Ok(None)
return Error(CredError.of_misconfigured(f"{auth_type} requires a usable upstream credential"))
async def prepare_mcp_client(server: MCPServer, client: MCPClient) -> MCPClient:
if server.auth_type not in _STATIC_MODES or client.transport_type == MCPTransport.stdio:
return client
request: Final = await client.prepare_request_auth()
match validate_static_credential(server.auth_type, request.headers, server.upstream_token_header):
case Error(error):
raise_public(error)
case Ok():
return client

View file

@ -1934,3 +1934,18 @@ async def test_discovery_auth_fingerprint_tracks_effective_credentials(resolved:
assert original != replaced
assert len(original) == 64
assert "private-original-credential" not in original
@pytest.mark.asyncio
async def test_request_auth_preview_uses_the_same_effective_headers_as_egress() -> None:
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import StaticHeaderAuth
client: Final = MCPClient(
server_url="https://upstream.example/mcp", auth_type=MCPAuth.bearer_token,
resolved_auth=StaticHeaderAuth("Bearer resolved"), extra_headers={"X-Trace": "trace"},
)
request: Final = await client.prepare_request_auth()
assert request.method == "POST"
assert str(request.url) == "https://upstream.example/mcp"
assert request.headers["Authorization"] == "Bearer resolved"
assert request.headers["X-Trace"] == "trace"

View file

@ -7,6 +7,7 @@ maps each CredError onto its HTTP status. These pin the parity-critical mapping
import base64
from types import SimpleNamespace
from typing import Final
import pytest
from fastapi import HTTPException
@ -20,7 +21,9 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import
raise_user_oauth_challenge,
to_server_spec,
to_subject,
validate_static_credential,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Ok
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
ApiKeyConfig,
AuthorizationCodeConfig,
@ -34,10 +37,28 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
SharedKey,
TokenExchangeConfig,
)
from litellm.types.mcp import MCPAuth, MCPTransport
from litellm.types.mcp import MCPAuth, MCPAuthType, MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
@pytest.mark.parametrize("auth_type,header,value", [
(MCPAuth.api_key, "Authorization", "Bearer fixture-key"),
(MCPAuth.api_key, "Authorization", "ApiKey fixture-key"),
(MCPAuth.api_key, "Authorization", "token fixture-key"),
(MCPAuth.api_key, "Authorization", "Bearer token"),
(MCPAuth.api_key, "Authorization", "opaque-key"),
(MCPAuth.api_key, "Authorization", "Custom Custom"),
(MCPAuth.api_key, "X-API-Key", "Bearer Bearer"),
(MCPAuth.api_key, "X-Custom", "ApiKey ApiKey"),
(MCPAuth.authorization, "Authorization", "opaque-secret-value"),
])
def test_static_credential_preserves_supported_api_key_and_raw_headers(
auth_type: MCPAuthType, header: str, value: str,
) -> None:
result: Final = validate_static_credential(auth_type, {header: value}, upstream_token_header=header)
assert isinstance(result, Ok)
def _server(**kwargs) -> MCPServer:
return MCPServer(server_id="s", name="n", transport=MCPTransport.http, **kwargs)
@ -155,12 +176,6 @@ def test_oauth2_user_token_maps_to_authorization_code(oauth2_flow):
_server(auth_type=MCPAuth.api_key), # no token configured
_server(auth_type=MCPAuth.bearer_token), # no token configured
_server(auth_type=MCPAuth.oauth2, delegate_auth_to_upstream=True), # delegated upstream OAuth -> v1
_server(auth_type=MCPAuth.oauth2_token_exchange), # no endpoint/client creds -> incomplete -> v1
_server(
auth_type=MCPAuth.oauth2_token_exchange,
token_exchange_endpoint="https://idp/token",
client_id="cid",
), # missing client_secret -> incomplete -> v1
_server(auth_type=MCPAuth.aws_sigv4),
_server(auth_type=None, oauth_passthrough=True, extra_headers=["Authorization"]),
],
@ -802,3 +817,14 @@ def test_a_blank_header_name_means_unset_rather_than_an_error(blank):
spec = to_server_spec(server)
assert spec is not None
assert spec.config.header_name == "Authorization"
@pytest.mark.parametrize("client_secret", [None, ""])
@pytest.mark.parametrize("is_byok", [False, True])
def test_incomplete_obo_keeps_exchange_ownership(client_secret: str | None, is_byok: bool) -> None:
spec = to_server_spec(_server(auth_type=MCPAuth.oauth2_token_exchange, client_id="client",
client_secret=client_secret, is_byok=is_byok))
assert spec is not None
assert isinstance(spec.config, TokenExchangeConfig)
assert spec.config.client_id == "client"
assert spec.config.client_secret is None

View file

@ -5,11 +5,13 @@ import logging
import os
import sys
from datetime import datetime
from pathlib import Path
from typing import Any, Dict, Final, Literal, Optional
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import HTTPException
from respx import MockRouter
from litellm.proxy._experimental.mcp_server.exceptions import (
MCPServerListError,
@ -5127,7 +5129,8 @@ class TestMCPServerManager:
captured: dict = {}
def fake_create_tool_function(
path, method, operation, base_url, headers=None, server_label=None, relays_upstream_auth=False
path, method, operation, base_url, headers=None, server_label=None, relays_upstream_auth=False,
auth_type=None, upstream_token_header=None,
):
captured["headers"] = headers
captured["server_label"] = server_label
@ -5212,7 +5215,8 @@ class TestMCPServerManager:
captured: dict = {}
def fake_create_tool_function(
path, method, operation, base_url, headers=None, server_label=None, relays_upstream_auth=False
path, method, operation, base_url, headers=None, server_label=None, relays_upstream_auth=False,
auth_type=None, upstream_token_header=None,
):
captured["headers"] = headers
@ -9401,12 +9405,13 @@ class TestCreateMcpClientV2Graft:
assert "misconfigured" in str(exc_info.value.detail)
assert "token_url" in str(exc_info.value.detail)
async def test_static_token_missing_defers_to_v1(self):
client = await MCPServerManager()._create_mcp_client(
self._http_server(auth_type=MCPAuth.api_key, authentication_token=None)
)
assert client._resolved_auth is None
async def test_static_token_missing_rejects_before_connecting(self):
with pytest.raises(HTTPException) as exc:
await MCPServerManager()._create_mcp_client(
self._http_server(auth_type=MCPAuth.api_key, authentication_token=None)
)
assert exc.value.status_code == 500
assert "credential" in str(exc.value.detail)
async def test_stdio_migrated_auth_type_still_defers_to_v1(self):
client = await MCPServerManager()._create_mcp_client(
@ -13467,3 +13472,400 @@ async def test_discovery_cache_returns_oversized_results_without_retaining_them(
result: Final = await cache.get(("server", None), fetch)
assert result[0].description == description
assert fetch.await_count == 2
class TestProtectedCredentialPreparation:
@pytest.mark.asyncio
@pytest.mark.parametrize("auth_type,credential", [
(MCPAuth.bearer_token, None),
(MCPAuth.bearer_token, "Bearer"),
(MCPAuth.api_key, None),
(MCPAuth.basic, "Basic"),
])
@pytest.mark.parametrize("dispatch", ["managed", "local"])
async def test_openapi_dispatch_rejects_unusable_effective_credentials(
self, tmp_path: Path, respx_mock: MockRouter, monkeypatch: pytest.MonkeyPatch,
auth_type: MCPAuthType, credential: str | None, dispatch: str,
) -> None:
from litellm.proxy._experimental.mcp_server.server import _handle_local_mcp_tool
from litellm.proxy._experimental.mcp_server.utils import add_server_prefix_to_name, get_server_prefix
spec_path: Final = tmp_path / "openapi.json"
spec_path.write_text(json.dumps({"openapi": "3.0.0", "info": {"title": "Auth", "version": "1"},
"paths": {"/echo": {"get": {"operationId": "echo"}}}}))
server: Final = MCPServer(
server_id="dispatch-auth", name="dispatch-auth", url="https://upstream.example",
transport=MCPTransport.http, auth_type=auth_type, authentication_token=credential,
)
manager: Final = MCPServerManager()
await manager._register_openapi_tools(str(spec_path), server, server.url)
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
destination: Final = respx_mock.get("https://upstream.example/echo").respond(200, text="unexpected success")
result: Final = (
await manager._call_openapi_tool_handler(server, "echo", {})
if dispatch == "managed"
else await _handle_local_mcp_tool(add_server_prefix_to_name("echo", get_server_prefix(server)), {})
)
assert result.isError is True
assert "requires a usable upstream credential" in result.content[0].text
assert destination.call_count == 0
@pytest.mark.asyncio
@pytest.mark.parametrize("transport", [MCPTransport.http, MCPTransport.sse])
@pytest.mark.parametrize("client_secret", [None, ""])
@pytest.mark.parametrize("subject", [None, "caller-subject"])
async def test_incomplete_obo_rejects_caller_and_static_fallback(
self, transport: MCPTransport, client_secret: str | None, subject: str | None
) -> None:
server = MCPServer(
server_id="incomplete-obo", name="incomplete-obo", url="https://upstream.example/mcp",
transport=transport, auth_type=MCPAuth.oauth2_token_exchange,
client_id="gateway", client_secret=client_secret,
token_exchange_endpoint="https://idp.example/token", authentication_token="static-fallback",
)
with pytest.raises(HTTPException) as exc:
await MCPServerManager()._create_mcp_client(
server, mcp_auth_header="Bearer override", subject_token=subject,
)
assert exc.value.status_code == (401 if subject is None else 500)
assert "static-fallback" not in str(exc.value.detail)
assert "override" not in str(exc.value.detail)
@pytest.mark.asyncio
@pytest.mark.parametrize("auth_type", [MCPAuth.api_key, MCPAuth.bearer_token])
@pytest.mark.parametrize("credential", [None, "", " ", {"X-Trace": "trace"}])
async def test_static_auth_without_usable_credential_rejects(
self, auth_type: MCPAuthType, credential: str | dict[str, str] | None
) -> None:
server = MCPServer(
server_id="empty-static", name="empty-static", url="https://upstream.example/mcp",
transport=MCPTransport.http, auth_type=auth_type,
)
with pytest.raises(HTTPException) as exc:
await MCPServerManager()._create_mcp_client(server, mcp_auth_header=credential)
assert exc.value.status_code == 500
assert "credential" in str(exc.value.detail).lower()
@pytest.mark.asyncio
@pytest.mark.parametrize("auth_type,headers", [
(MCPAuth.api_key, {"X-API-Key": "key"}),
(MCPAuth.bearer_token, {"Authorization": "Bearer token"}),
])
async def test_static_auth_accepts_actual_forwarded_credential(
self, auth_type: MCPAuthType, headers: dict[str, str]
) -> None:
server = MCPServer(
server_id="header-static", name="header-static", url="https://upstream.example/mcp",
transport=MCPTransport.http, auth_type=auth_type,
)
client = await MCPServerManager()._create_mcp_client(server, extra_headers=headers)
assert client._get_auth_headers() == headers
@pytest.mark.asyncio
@pytest.mark.parametrize("auth_type", [MCPAuth.oauth2_token_exchange])
async def test_openapi_protected_auth_rejects_missing_credentials(self, auth_type: MCPAuthType) -> None:
server = MCPServer(
server_id="openapi-empty", name="openapi-empty", url="https://upstream.example/mcp",
transport=MCPTransport.http, auth_type=auth_type,
token_exchange_endpoint="https://idp.example/token",
)
with pytest.raises(HTTPException) as exc:
await MCPServerManager().resolve_openapi_upstream_auth(
mcp_server=server, oauth2_headers=None, raw_headers=None, mcp_auth_header=None,
user_api_key_auth=None, forwarded_headers=None,
)
assert exc.value.status_code in (401, 500)
@pytest.mark.asyncio
@pytest.mark.parametrize("auth_type,slot,value", [
(MCPAuth.api_key, "X-API-Key", "token"),
(MCPAuth.authorization, "Authorization", "opaque-secret-value"),
(MCPAuth.authorization, "Authorization", "Bearer abc"),
(MCPAuth.authorization, "Authorization", "Custom abc"),
])
async def test_raw_static_credentials_are_forwarded_unchanged(
self, auth_type: MCPAuthType, slot: str, value: str,
) -> None:
server = MCPServer(server_id="raw-key", name="raw-key", url="https://upstream.example/mcp",
transport=MCPTransport.http, auth_type=auth_type, authentication_token=value)
client = await MCPServerManager()._create_mcp_client(server)
assert client._resolved_auth is not None
request = httpx.Request("GET", server.url)
flow = client._resolved_auth.auth_flow(request)
try:
assert next(flow).headers[slot] == value
finally:
flow.close()
@pytest.mark.asyncio
@pytest.mark.parametrize("value", ["Bearer", "basic", "token", "ApiKey", " bEaReR ", "\tTOKEN\t"])
@pytest.mark.parametrize("source", ["configured", "caller", "forwarded"])
async def test_raw_authorization_rejects_bare_schemes_before_dispatch(
self, respx_mock: MockRouter, value: str, source: str,
) -> None:
server: Final = MCPServer(
server_id="raw-empty", name="raw-empty", url="https://upstream.example/mcp",
transport=MCPTransport.http, auth_type=MCPAuth.authorization,
authentication_token=value if source == "configured" else None,
)
destination: Final = respx_mock.route().respond(200)
with pytest.raises(HTTPException, match="requires a usable upstream credential") as exc:
await MCPServerManager()._create_mcp_client(
server, mcp_auth_header=value if source == "caller" else None,
extra_headers={"Authorization": value} if source == "forwarded" else None,
)
assert exc.value.status_code == 500
assert destination.call_count == 0
@pytest.mark.asyncio
async def test_byok_flag_cannot_bypass_incomplete_obo(self) -> None:
server = MCPServer(server_id="obo-byok", name="obo-byok", url="https://upstream.example/mcp",
transport=MCPTransport.http, auth_type=MCPAuth.oauth2_token_exchange, is_byok=True,
token_exchange_endpoint="https://idp.example/token")
with pytest.raises(HTTPException) as exc:
await MCPServerManager()._create_mcp_client(server, mcp_auth_header="Bearer override")
assert exc.value.status_code == 401
@pytest.mark.asyncio
@pytest.mark.parametrize("configured,override", [(None, "Bearer usable"), ("shared", "Bearer usable")])
async def test_bearer_override_remains_usable(self, configured: str | None, override: str) -> None:
server = MCPServer(server_id="override", name="override", url="https://upstream.example/mcp",
transport=MCPTransport.http, auth_type=MCPAuth.bearer_token, authentication_token=configured)
client = await MCPServerManager()._create_mcp_client(server, mcp_auth_header=override)
assert client._get_auth_headers()["Authorization"] == override
@pytest.mark.asyncio
@pytest.mark.parametrize("token", [None, "shared"])
async def test_empty_injected_header_cannot_satisfy_protected_auth(self, token: str | None) -> None:
server = MCPServer(server_id="empty-header", name="empty-header", url="https://upstream.example/mcp",
transport=MCPTransport.http, auth_type=MCPAuth.bearer_token, authentication_token=token)
with pytest.raises(HTTPException) as exc:
await MCPServerManager()._create_mcp_client(server, extra_headers={"authorization": " "})
assert exc.value.status_code == 500
@pytest.mark.asyncio
async def test_custom_slot_uses_its_actual_credential(self) -> None:
server = MCPServer(server_id="custom", name="custom", url="https://upstream.example/mcp",
transport=MCPTransport.http, auth_type=MCPAuth.api_key,
upstream_token_header="X-Custom", authentication_token="key")
client = await MCPServerManager()._create_mcp_client(server, extra_headers={"X-Trace": "trace"})
assert client._credential_slot == "X-Custom"
assert await client.discovery_auth_fingerprint()
@pytest.mark.asyncio
@pytest.mark.parametrize("static,forwarded,caller", [
({"X-API-Key": "static"}, {"x-api-key": "forwarded"}, None),
({}, {"X-API-Key": "forwarded"}, None),
({}, None, "ApiKey caller"),
({"X-API-Key": "static"}, {"Authorization": ""}, None),
])
async def test_openapi_static_credentials_remain_supported(
self, respx_mock: MockRouter, monkeypatch: pytest.MonkeyPatch,
static: dict[str, str], forwarded: dict[str, str] | None, caller: str | None
) -> None:
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
_request_auth_header, _request_extra_headers, create_tool_function,
)
tool: Final = create_tool_function(
"/echo", "get", {}, "https://upstream.example", headers=static, auth_type=MCPAuth.api_key,
)
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
destination: Final = respx_mock.get("https://upstream.example/echo").respond(200, text="authenticated")
caller_token: Final = _request_auth_header.set(caller)
extra_token: Final = _request_extra_headers.set(forwarded)
try:
assert await tool() == "authenticated"
sent: Final = destination.calls.last.request.headers
assert sent.get("x-api-key") == static.get("X-API-Key", (forwarded or {}).get("X-API-Key"))
if caller:
assert sent["authorization"] == caller
assert destination.call_count == 1
finally:
_request_auth_header.reset(caller_token)
_request_extra_headers.reset(extra_token)
@pytest.mark.asyncio
async def test_static_resolution_cancellation_closes_flow(self) -> None:
from collections.abc import AsyncGenerator
from litellm.experimental_mcp_client.client import MCPClient
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import prepare_mcp_client
class CancelledAuth(httpx.Auth):
closed = False
async def async_auth_flow(self, request: httpx.Request) -> AsyncGenerator[httpx.Request, httpx.Response]:
try:
raise asyncio.CancelledError()
yield request
finally:
self.closed = True
auth = CancelledAuth()
server = MCPServer(server_id="cancel", name="cancel", url="https://upstream.example/mcp",
transport=MCPTransport.http, auth_type=MCPAuth.api_key)
client = MCPClient(server_url=server.url, auth_type=MCPAuth.api_key, resolved_auth=auth)
with pytest.raises(asyncio.CancelledError):
await prepare_mcp_client(server, client)
assert auth.closed
@pytest.mark.asyncio
@pytest.mark.parametrize("auth_type", [MCPAuth.basic, MCPAuth.token, MCPAuth.authorization])
async def test_other_static_schemes_reject_whitespace_credentials(self, auth_type: MCPAuthType) -> None:
server = MCPServer(server_id="blank-static", name="blank-static", url="https://upstream.example/mcp",
transport=MCPTransport.http, auth_type=auth_type, authentication_token=" ")
with pytest.raises(HTTPException) as exc:
await MCPServerManager()._create_mcp_client(server)
assert exc.value.status_code == 500
@pytest.mark.asyncio
@pytest.mark.parametrize("header", ["Basic", "Basic @@@", "Other abc", "Basic QmFzaWM=", "Basic bm8tY29sb24="])
async def test_basic_headers_without_usable_credentials_reject(self, header: str) -> None:
server = MCPServer(server_id="bad-basic", name="bad-basic", url="https://upstream.example/mcp",
transport=MCPTransport.http, auth_type=MCPAuth.basic)
with pytest.raises(HTTPException) as exc:
await MCPServerManager()._create_mcp_client(server, extra_headers={"Authorization": header})
assert exc.value.status_code == 500
@pytest.mark.asyncio
@pytest.mark.parametrize("value", ["Basic", "Basic ", "basic"])
@pytest.mark.parametrize("source", ["configured", "caller"])
async def test_basic_scheme_alone_is_not_a_credential(self, value: str, source: str) -> None:
server = MCPServer(server_id="basic-scheme", name="basic-scheme", url="https://upstream.example/mcp",
transport=MCPTransport.http, auth_type=MCPAuth.basic,
authentication_token=value if source == "configured" else None)
with pytest.raises(HTTPException) as exc:
await MCPServerManager()._create_mcp_client(server, mcp_auth_header=value if source == "caller" else None)
assert exc.value.status_code == 500
@pytest.mark.asyncio
@pytest.mark.parametrize("auth_type,value,default_slot", [
(MCPAuth.api_key, "fixture-key", "X-API-Key"),
(MCPAuth.bearer_token, "fixture-key", "Authorization"),
(MCPAuth.basic, "user:pass", "Authorization"),
(MCPAuth.token, "fixture-key", "Authorization"),
(MCPAuth.authorization, "fixture-key", "Authorization"),
])
@pytest.mark.parametrize("source", ["configured", "caller"])
async def test_usable_credential_survives_an_empty_alternate_header(
self, auth_type: MCPAuthType, value: str, default_slot: str, source: str
) -> None:
server: Final = MCPServer(
server_id="alternate", name="alternate", url="https://upstream.example/mcp",
transport=MCPTransport.http, auth_type=auth_type, upstream_token_header="X-Custom",
authentication_token=value if source == "configured" else None,
)
empty_slot: Final = default_slot if source == "configured" else "X-Custom"
selected_slot: Final = "X-Custom" if source == "configured" else default_slot
client: Final = await MCPServerManager()._create_mcp_client(
server, mcp_auth_header=value if source == "caller" else None, extra_headers={empty_slot: ""},
)
request: Final = await client.prepare_request_auth()
assert request.headers[selected_slot]
assert request.headers[empty_slot] == ""
@pytest.mark.asyncio
async def test_empty_custom_and_default_headers_do_not_satisfy_auth(self) -> None:
server: Final = MCPServer(
server_id="both-empty", name="both-empty", url="https://upstream.example/mcp",
transport=MCPTransport.http, auth_type=MCPAuth.api_key, upstream_token_header="X-Custom",
)
with pytest.raises(HTTPException) as exc:
await MCPServerManager()._create_mcp_client(server, extra_headers={"X-Custom": "", "X-API-Key": ""})
assert exc.value.status_code == 500
@pytest.mark.asyncio
@pytest.mark.parametrize("custom_slot", [None, "X-Custom"])
@pytest.mark.parametrize("source", ["caller", "forwarded"])
async def test_api_key_preserves_explicit_authorization_credential(
self, custom_slot: str | None, source: str
) -> None:
server: Final = MCPServer(
server_id="caller-auth", name="caller-auth", url="https://upstream.example/mcp",
transport=MCPTransport.http, auth_type=MCPAuth.api_key, upstream_token_header=custom_slot,
)
headers: Final = {"Authorization": "Bearer caller-credential", "X-API-Key": ""}
client: Final = await MCPServerManager()._create_mcp_client(
server, mcp_auth_header=headers if source == "caller" else None,
extra_headers=headers if source == "forwarded" else None,
)
request: Final = await client.prepare_request_auth()
assert request.headers["Authorization"] == "Bearer caller-credential"
assert request.headers["X-API-Key"] == ""
assert custom_slot is None or custom_slot not in request.headers
@pytest.mark.asyncio
@pytest.mark.parametrize("value", [
"", " ", "Bearer", "Basic", "token", "ApiKey",
"Bearer Bearer", "ApiKey ApiKey", "token token", "bEaReR BEARER", "aPiKeY\tAPIKEY",
])
async def test_api_key_rejects_authorization_without_a_credential(self, value: str) -> None:
server: Final = MCPServer(
server_id="caller-empty", name="caller-empty", url="https://upstream.example/mcp",
transport=MCPTransport.http, auth_type=MCPAuth.api_key,
)
with pytest.raises(HTTPException) as exc:
await MCPServerManager()._create_mcp_client(server, mcp_auth_header={"Authorization": value})
assert exc.value.status_code == 500
@pytest.mark.asyncio
@pytest.mark.parametrize("value", ["no-colon", "Basic bm8tY29sb24="])
@pytest.mark.parametrize("source", ["configured", "caller"])
async def test_basic_requires_a_username_password_separator(self, value: str, source: str) -> None:
server: Final = MCPServer(
server_id="basic-pair", name="basic-pair", url="https://upstream.example/mcp",
transport=MCPTransport.http, auth_type=MCPAuth.basic,
authentication_token=value if source == "configured" else None,
)
with pytest.raises(HTTPException) as exc:
await MCPServerManager()._create_mcp_client(server, mcp_auth_header=value if source == "caller" else None)
assert exc.value.status_code == 500
@pytest.mark.asyncio
@pytest.mark.parametrize("value", ["user:pass", "user:", ":pass", ":"])
async def test_basic_preserves_username_password_pairs(self, value: str) -> None:
import base64
server: Final = MCPServer(
server_id="basic-valid", name="basic-valid", url="https://upstream.example/mcp",
transport=MCPTransport.http, auth_type=MCPAuth.basic, authentication_token=value,
)
client: Final = await MCPServerManager()._create_mcp_client(server)
request: Final = await client.prepare_request_auth()
scheme, encoded = request.headers["Authorization"].split(" ", 1)
assert scheme == "Basic"
assert base64.b64decode(encoded) == value.encode()
@pytest.mark.asyncio
@pytest.mark.parametrize("auth_type,value", [
(MCPAuth.bearer_token, "Bearer"), (MCPAuth.bearer_token, "Bearer "), (MCPAuth.bearer_token, "bearer"),
(MCPAuth.token, "token"), (MCPAuth.token, "token "), (MCPAuth.token, "TOKEN"),
])
@pytest.mark.parametrize("source", ["configured", "caller"])
async def test_static_scheme_only_input_cannot_hide_behind_rendered_prefix(
self, auth_type: MCPAuthType, value: str, source: str
) -> None:
server: Final = MCPServer(
server_id="empty-scheme", name="empty-scheme", url="https://upstream.example/mcp",
transport=MCPTransport.http, auth_type=auth_type,
authentication_token=value if source == "configured" else None,
)
with pytest.raises(HTTPException) as exc:
await MCPServerManager()._create_mcp_client(server, mcp_auth_header=value if source == "caller" else None)
assert exc.value.status_code == 500
@pytest.mark.asyncio
@pytest.mark.parametrize("auth_type,value,expected", [
(MCPAuth.bearer_token, "token", "Bearer token"),
(MCPAuth.bearer_token, "Bearertoken", "Bearer Bearertoken"),
(MCPAuth.token, "tokenish", "token tokenish"),
])
async def test_static_credentials_that_resemble_schemes_remain_usable(
self, auth_type: MCPAuthType, value: str, expected: str
) -> None:
server: Final = MCPServer(
server_id="real-token", name="real-token", url="https://upstream.example/mcp",
transport=MCPTransport.http, auth_type=auth_type, authentication_token=value,
)
client: Final = await MCPServerManager()._create_mcp_client(server)
request: Final = await client.prepare_request_auth()
assert request.headers["Authorization"] == expected

View file

@ -10,9 +10,14 @@ This test suite ensures that:
"""
from types import SimpleNamespace
from typing import Final
from unittest.mock import AsyncMock, patch
import pytest
from fastapi import HTTPException
from respx import MockRouter
from litellm.types.mcp import MCPAuth, MCPAuthType
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
_request_auth_header,
@ -35,6 +40,120 @@ from litellm.proxy._experimental.mcp_server.exceptions import (
GET_ASYNC_CLIENT_TARGET = "litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator.get_async_httpx_client"
@pytest.mark.asyncio
@pytest.mark.parametrize("auth_type,value,accepted", [
(MCPAuth.api_key, "Bearer Bearer", False), (MCPAuth.api_key, "ApiKey ApiKey", False),
(MCPAuth.api_key, "token token", False), (MCPAuth.api_key, "bEaReR BEARER", False),
(MCPAuth.api_key, "aPiKeY\tAPIKEY", False), (MCPAuth.api_key, "Bearer fixture-key", True),
(MCPAuth.api_key, "ApiKey fixture-key", True), (MCPAuth.api_key, "token fixture-key", True),
(MCPAuth.authorization, "Bearer", False), (MCPAuth.authorization, "basic", False),
(MCPAuth.authorization, "token", False), (MCPAuth.authorization, "ApiKey", False),
(MCPAuth.authorization, " bEaReR ", False), (MCPAuth.authorization, "\tTOKEN\t", False),
(MCPAuth.authorization, "opaque-secret-value", True), (MCPAuth.authorization, "Bearer abc", True),
(MCPAuth.authorization, "Custom abc", True),
])
async def test_authorization_validates_credentials_before_http(
respx_mock: MockRouter, monkeypatch: pytest.MonkeyPatch, auth_type: MCPAuthType, value: str, accepted: bool,
) -> None:
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
tool: Final = create_tool_function(
"/echo", "get", {}, "https://upstream.example", auth_type=auth_type,
)
destination: Final = respx_mock.get("https://upstream.example/echo").respond(200, text="authenticated")
caller_token: Final = _request_auth_header.set(value)
try:
if accepted:
assert await tool() == "authenticated"
assert destination.call_count == 1
assert destination.calls.last.request.headers["authorization"] == value
else:
with pytest.raises(HTTPException, match="requires a usable upstream credential") as exc:
await tool()
assert exc.value.status_code == 500
assert destination.call_count == 0
finally:
_request_auth_header.reset(caller_token)
@pytest.mark.asyncio
@pytest.mark.parametrize("static,forwarded,caller,resolved,expected", [
({"Authorization": "Bearer configured"}, {"authorization": "Bearer forwarded"}, None, None, "Bearer configured"),
({"Authorization": "Bearer configured"}, None, "Bearer caller", None, "Bearer caller"),
({"Authorization": "Bearer configured"}, None, "Bearer", None, None),
({"Authorization": "Bearer configured"}, None, "Bearer caller", {"authorization": " "}, None),
({"Authorization": "Bearer configured"}, None, "Bearer", {"authorization": "Bearer resolved"}, "Bearer resolved"),
])
async def test_static_auth_validates_headers_after_existing_precedence(
respx_mock: MockRouter, monkeypatch: pytest.MonkeyPatch,
static: dict[str, str], forwarded: dict[str, str] | None, caller: str | None,
resolved: dict[str, str] | None, expected: str | None,
) -> None:
tool: Final = create_tool_function(
"/echo", "get", {}, "https://upstream.example", headers=static, auth_type=MCPAuth.bearer_token,
)
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
destination: Final = respx_mock.get("https://upstream.example/echo").respond(200, text="authenticated")
caller_token: Final = _request_auth_header.set(caller)
extra_token: Final = _request_extra_headers.set(forwarded)
resolved_token: Final = _request_resolved_auth_headers.set(resolved)
try:
if expected is None:
with pytest.raises(HTTPException, match="requires a usable upstream credential") as exc:
await tool()
assert exc.value.status_code == 500
assert destination.call_count == 0
else:
assert await tool() == "authenticated"
assert destination.call_count == 1
assert destination.calls.last.request.headers["authorization"] == expected
finally:
_request_auth_header.reset(caller_token)
_request_extra_headers.reset(extra_token)
_request_resolved_auth_headers.reset(resolved_token)
@pytest.mark.asyncio
@pytest.mark.parametrize("credential", ["custom-key", ""])
async def test_static_auth_uses_configured_custom_header(
respx_mock: MockRouter, monkeypatch: pytest.MonkeyPatch, credential: str,
) -> None:
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
tool: Final = create_tool_function(
"/echo", "get", {}, "https://upstream.example", headers={"x-custom": credential},
auth_type=MCPAuth.api_key, upstream_token_header="X-Custom",
)
destination: Final = respx_mock.get("https://upstream.example/echo").respond(200, text="authenticated")
if credential:
assert await tool() == "authenticated"
assert destination.call_count == 1
assert destination.calls.last.request.headers["x-custom"] == credential
else:
with pytest.raises(HTTPException, match="requires a usable upstream credential"):
await tool()
assert destination.call_count == 0
@pytest.mark.asyncio
@pytest.mark.parametrize("auth_type,resolved", [
(MCPAuth.none, None),
(MCPAuth.oauth2, {"Authorization": "Bearer user-oauth"}),
])
async def test_static_validation_preserves_no_auth_and_resolved_oauth(
respx_mock: MockRouter, monkeypatch: pytest.MonkeyPatch,
auth_type: MCPAuthType, resolved: dict[str, str] | None,
) -> None:
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
tool: Final = create_tool_function("/echo", "get", {}, "https://upstream.example", auth_type=auth_type)
destination: Final = respx_mock.get("https://upstream.example/echo").respond(200, text="echo")
token: Final = _request_resolved_auth_headers.set(resolved)
try:
assert await tool() == "echo"
assert destination.call_count == 1
assert destination.calls.last.request.headers.get("authorization") == (resolved or {}).get("Authorization")
finally:
_request_resolved_auth_headers.reset(token)
def _create_mock_client(method: str, response_text: str, status_code: int = 200) -> AsyncMock:
"""Utility to create a mocked async httpx client for the given method.
@ -1458,3 +1577,21 @@ class TestBoundedOpenAPISpecLoading:
else:
assert await load_openapi_spec_async("https://93.184.216.34/spec.json", max_bytes=100) == {"paths": {}}
assert destination.call_count == 1
def test_openapi_generator_import_does_not_require_mcp_sdk() -> None:
import subprocess
import sys
script = """
import builtins
original_import = builtins.__import__
def without_mcp(name, *args, **kwargs):
if name == 'mcp' or name.startswith('mcp.'):
raise ModuleNotFoundError('MCP SDK unavailable')
return original_import(name, *args, **kwargs)
builtins.__import__ = without_mcp
import litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator
"""
result = subprocess.run([sys.executable, "-c", script], capture_output=True, text=True)
assert result.returncode == 0, result.stderr