From 7beb4778ab9cbd8018566a0955f352ff3806e065 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 20 Aug 2026 16:53:05 +0000 Subject: [PATCH] fix(mcp): allow oauth token to target a custom upstream header Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/experimental_mcp_client/client.py | 13 +++- .../mcp_server/mcp_server_manager.py | 30 ++++++--- .../outbound_credentials/adapter.py | 3 +- .../client_credentials.py | 14 +++- .../outbound_credentials/resolver.py | 2 +- .../mcp_server/outbound_credentials/types.py | 7 +- .../proxy/_experimental/mcp_server/server.py | 4 +- .../mcp_management_endpoints.py | 1 + litellm/types/mcp.py | 20 +++++- .../types/mcp_server/mcp_server_manager.py | 4 ++ .../test_mcp_client.py | 38 +++++++++++ .../outbound_credentials/test_adapter.py | 31 +++++++++ .../test_client_credentials.py | 26 ++++++++ .../outbound_credentials/test_resolver.py | 13 ++++ .../mcp_server/test_mcp_server_manager.py | 64 +++++++++++++++++-- .../src/components/mcp_tools/types.tsx | 2 +- 16 files changed, 244 insertions(+), 28 deletions(-) diff --git a/litellm/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py index 7bd0a847ad8..81b620e1648 100644 --- a/litellm/experimental_mcp_client/client.py +++ b/litellm/experimental_mcp_client/client.py @@ -38,11 +38,13 @@ from litellm.constants import MCP_CLIENT_TIMEOUT, MCP_NPM_CACHE_DIR from litellm.llms.custom_httpx.http_handler import get_ssl_configuration from litellm.types.llms.custom_http import VerifyTypes from litellm.types.mcp import ( + DEFAULT_OAUTH_TOKEN_HEADER, MCPAuth, MCPAuthType, MCPStdioConfig, MCPTransport, MCPTransportType, + resolve_oauth_token_header, ) @@ -227,6 +229,7 @@ class MCPClient: timeout: float | None = None, stdio_config: MCPStdioConfig | None = None, extra_headers: dict[str, str] | None = None, + oauth_token_header: str | None = None, ssl_verify: VerifyTypes | None = None, aws_auth: httpx.Auth | None = None, resolved_auth: httpx.Auth | None = None, @@ -241,6 +244,7 @@ class MCPClient: self._mcp_auth_value: str | dict[str, str] | None = None self.stdio_config: MCPStdioConfig | None = stdio_config self.extra_headers: dict[str, str] | None = extra_headers + self.oauth_token_header: str = resolve_oauth_token_header(oauth_token_header) self.ssl_verify: VerifyTypes | None = ssl_verify self._aws_auth: httpx.Auth | None = aws_auth # A pre-resolved httpx.Auth (e.g. from the v2 credential resolver) attached to the @@ -467,19 +471,24 @@ class MCPClient: elif self.auth_type == MCPAuth.authorization: headers["Authorization"] = self._mcp_auth_value elif self.auth_type == MCPAuth.oauth2: - headers["Authorization"] = f"Bearer {self._mcp_auth_value}" + headers[self.oauth_token_header] = f"Bearer {self._mcp_auth_value}" elif self.auth_type == MCPAuth.token: headers["Authorization"] = f"token {self._mcp_auth_value}" elif self.auth_type == MCPAuth.oauth2_token_exchange: - headers["Authorization"] = f"Bearer {self._mcp_auth_value}" + headers[self.oauth_token_header] = f"Bearer {self._mcp_auth_value}" elif isinstance(self._mcp_auth_value, dict): headers.update(self._mcp_auth_value) # Note: aws_sigv4 auth is not handled here — SigV4 requires per-request # signing (including the body hash), so it uses httpx.Auth flow instead # of static headers. See MCPSigV4Auth and _create_httpx_client_factory(). # update the headers with the extra headers + minted_oauth_token: Final = headers.get(self.oauth_token_header) if self.extra_headers: headers.update(self.extra_headers) + # A custom oauth_token_header exists so the minted token can ride beside a static + # Authorization, so it owns that header instead of being shadowed by a static entry. + if minted_oauth_token is not None and self.oauth_token_header != DEFAULT_OAUTH_TOKEN_HEADER: + headers[self.oauth_token_header] = minted_oauth_token return _strip_header_whitespace(headers) def _create_httpx_client_factory(self) -> Callable[..., httpx.AsyncClient]: diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 26a6f8d1251..5f2cfcb79e9 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -34,6 +34,7 @@ from mcp.types import ( ) from mcp.types import Tool as MCPTool from pydantic import AnyUrl, BaseModel +from typing_extensions import ReadOnly import litellm from litellm._logging import verbose_logger @@ -149,6 +150,7 @@ from litellm.proxy.utils import PrismaClient, ProxyLogging, get_server_root_path from litellm.repositories.table_repositories import MCPServerRepository from litellm.types.llms.custom_http import httpxSpecialProvider from litellm.types.mcp import ( + DEFAULT_OAUTH_TOKEN_HEADER, DEFAULT_SUBJECT_TOKEN_TYPE, MCPAuth, MCPStdioConfig, @@ -349,6 +351,7 @@ class MCPServerConfig(TypedDict, total=False): audience: str subject_token_type: str upstream_resource: str + oauth_token_header: ReadOnly[str] id_jag_resource_token_endpoint: str id_jag_resource: str client_private_key: str @@ -828,15 +831,17 @@ def _should_strip_caller_authorization( ) -def _without_authorization( +def _without_header( headers: dict[str, str] | None, + name: str = DEFAULT_OAUTH_TOKEN_HEADER, ) -> dict[str, str] | None: - """A copy of ``headers`` with any ``Authorization`` key removed (case-insensitive), or - None if nothing remains. Drops only the credential, keeping other forwarded headers. + """A copy of ``headers`` with ``name`` (``Authorization`` by default) removed + (case-insensitive), or None if nothing remains. Drops only that one header, keeping every + other forwarded header. """ if not headers: return None - filtered: Final = {k: v for k, v in headers.items() if k.lower() != "authorization"} + filtered: Final = {k: v for k, v in headers.items() if k.lower() != name.lower()} return filtered or None @@ -909,7 +914,7 @@ def _resolve_openapi_tool_auth( if isinstance(per_server, dict): authorization: Final = next((v for k, v in per_server.items() if k.lower() == "authorization"), None) - merged: Final = merge_mcp_headers(extra_headers=forwarded, static_headers=_without_authorization(per_server)) + merged: Final = merge_mcp_headers(extra_headers=forwarded, static_headers=_without_header(per_server)) if authorization is None: byok: Final = _format_byok_openapi_auth_header(mcp_server, mcp_auth_header) if mcp_auth_header else None return byok, merged, mcp_auth_header @@ -976,7 +981,7 @@ def _client_forwarded_authorization_headers( raw_headers=raw_headers, user_api_key_auth=user_api_key_auth, ): - return _without_authorization(extra_headers) + return _without_header(extra_headers) return extra_headers @@ -989,7 +994,7 @@ def _take_forwarded_authorization( if not headers: return None, headers value: Final = next((v for k, v in headers.items() if k.lower() == "authorization"), None) - return value, _without_authorization(headers) + return value, _without_header(headers) def _passthrough_token_from_mcp_auth_header( @@ -2161,6 +2166,7 @@ class MCPServerManager: DEFAULT_SUBJECT_TOKEN_TYPE, ), upstream_resource=server_config.get("upstream_resource", None), + oauth_token_header=server_config.get("oauth_token_header", None), # ID-JAG fields id_jag_resource_token_endpoint=server_config.get("id_jag_resource_token_endpoint", None), id_jag_resource=server_config.get("id_jag_resource", None), @@ -2693,6 +2699,7 @@ class MCPServerManager: or (credentials_dict.get("subject_token_type") if credentials_dict else None) or DEFAULT_SUBJECT_TOKEN_TYPE, upstream_resource=(credentials_dict.get("upstream_resource") if credentials_dict else None), + oauth_token_header=(credentials_dict.get("oauth_token_header") if credentials_dict else None), # ID-JAG fields — read from credentials JSON blob id_jag_resource_token_endpoint=( credentials_dict.get("id_jag_resource_token_endpoint") if credentials_dict else None @@ -3549,8 +3556,10 @@ class MCPServerManager: # Authorization must NOT shadow it (otherwise the upstream gets e.g. the # signer's JWT instead of the minted token and rejects it, and for M2M the # one-shot 401 refetch is lost with it). Drop the conflicting header so the - # resolved token reaches upstream. - return auth, _without_authorization(extra_headers) + # resolved token reaches upstream. Only that header is dropped, since a + # client_credentials server can mint onto a custom oauth_token_header while a + # static Authorization the upstream also requires rides alongside it. + return auth, _without_header(extra_headers, header_name or DEFAULT_OAUTH_TOKEN_HEADER) # Other modes: an Authorization already supplied via extra_headers (a forwarded caller # header or static_headers) is intentional and wins; v1 applies those last. return None, extra_headers @@ -3768,6 +3777,7 @@ class MCPServerManager: auth_value=auth_value, timeout=(resolved_server.timeout if resolved_server.timeout is not None else MCP_CLIENT_TIMEOUT), extra_headers=extra_headers, + oauth_token_header=resolved_server.oauth_token_header, aws_auth=aws_auth, sampling_callback=sampling_cb, elicitation_callback=elicitation_cb, @@ -5312,7 +5322,7 @@ class MCPServerManager: raw_headers=raw_headers, user_api_key_auth=user_api_key_auth, ): - extra_headers = _without_authorization(extra_headers) + extra_headers = _without_header(extra_headers) elif mcp_server.is_client_forwarded_token: extra_headers = _client_forwarded_authorization_headers( mcp_server=mcp_server, diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/adapter.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/adapter.py index be8ec1b8eb3..6688d014739 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/adapter.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/adapter.py @@ -35,7 +35,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, resolve_oauth_token_header if TYPE_CHECKING: from litellm.proxy._types import UserAPIKeyAuth @@ -147,6 +147,7 @@ def _client_credentials_spec(server: MCPServer, resource: str) -> ServerSpec: audience=server.audience, upstream_resource=resolve_upstream_resource(server), token_endpoint_auth_method=server.token_endpoint_auth_method, + token_header=resolve_oauth_token_header(server.oauth_token_header), ), ) diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/client_credentials.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/client_credentials.py index d0053fbe0a8..74b8fdb87c5 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/client_credentials.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/client_credentials.py @@ -51,6 +51,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.types import ( ClientCredentialsConfig, CredError, ) +from litellm.types.mcp import DEFAULT_OAUTH_TOKEN_HEADER class TokenEndpointSuccess(BaseModel): @@ -326,10 +327,19 @@ class ClientCredentialsBearerAuth(httpx.Auth): The initial token was already resolved (so config/IdP failures surfaced as typed errors before any upstream request); ``refetch`` is the source's 401-recovery callback. If the refetch fails, or the retried request 401s again, the upstream's response stands. + + ``header_name`` is the header the token is written to. It defaults to ``Authorization``; + a server whose upstream reads the minted token beside a static ``Authorization`` of its own + configures another header, and only that header is touched here so the static one survives. """ - def __init__(self, access_token: str, refetch: Callable[[str], Awaitable[str | None]]) -> None: - self.header_name = "Authorization" + def __init__( + self, + access_token: str, + refetch: Callable[[str], Awaitable[str | None]], + header_name: str = DEFAULT_OAUTH_TOKEN_HEADER, + ) -> None: + self.header_name = header_name self._access_token = SecretStr(access_token) self._refetch = refetch diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/resolver.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/resolver.py index 94c59962b70..85c48f05f09 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/resolver.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/resolver.py @@ -307,7 +307,7 @@ class UpstreamCredentialProvider: match await self._client_credentials_source.get(server_id, config): case Ok(token): refetch: Final = partial(self._client_credentials_source.refetch, server_id, config) - return Ok(ClientCredentialsBearerAuth(token.access_token, refetch)) + return Ok(ClientCredentialsBearerAuth(token.access_token, refetch, header_name=config.token_header)) case Error(err): return Error(err) diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/types.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/types.py index ce9948f0448..17a5fa36e7e 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/types.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/types.py @@ -39,7 +39,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.result import ( Ok, Result, ) -from litellm.types.mcp import DEFAULT_SUBJECT_TOKEN_TYPE +from litellm.types.mcp import DEFAULT_OAUTH_TOKEN_HEADER, DEFAULT_SUBJECT_TOKEN_TYPE class AuthSpecKind(str, Enum): @@ -190,6 +190,10 @@ class ClientCredentialsConfig(BaseModel): the client_credentials grant (sent as `audience` in the token request when set). `token_endpoint_auth_method` selects how the client authenticates to the token endpoint (RFC 6749 section 2.3.1); `None` defaults to `client_secret_post`. + + `token_header` is the upstream header the minted token is written to. It defaults to + `Authorization`, and an upstream that reads the minted token beside a static `Authorization` + credential of its own configures another header so the two do not collide. """ model_config = ConfigDict(frozen=True) @@ -201,6 +205,7 @@ class ClientCredentialsConfig(BaseModel): audience: str | None = None upstream_resource: str | None = None token_endpoint_auth_method: Literal["client_secret_post", "client_secret_basic"] | None = None + token_header: str = DEFAULT_OAUTH_TOKEN_HEADER class TokenExchangeConfig(BaseModel): diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 0dc85c0318c..4b4eac6426d 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -430,7 +430,7 @@ if MCP_AVAILABLE: _client_forwarded_authorization_headers, _resolve_openapi_tool_auth, _should_strip_caller_authorization, - _without_authorization, + _without_header, global_mcp_server_manager, ) from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( @@ -1730,7 +1730,7 @@ if MCP_AVAILABLE: raw_headers=raw_headers, user_api_key_auth=user_api_key_auth, ): - extra_headers = _without_authorization(extra_headers) + extra_headers = _without_header(extra_headers) elif is_client_forwarded_mode: if not withhold_forwarded_authorization: extra_headers = _client_forwarded_authorization_headers( diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 54a591a5e1a..074f0ff0ef0 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -739,6 +739,7 @@ if MCP_AVAILABLE: ("aws_region_name", "aws_region_name"), ("aws_service_name", "aws_service_name"), ("upstream_resource", "upstream_resource"), + ("oauth_token_header", "oauth_token_header"), ) def _has_non_admin_config_credentials(credentials: "MCPCredentials | None") -> bool: diff --git a/litellm/types/mcp.py b/litellm/types/mcp.py index 57437ea7e54..c02e461c1a3 100644 --- a/litellm/types/mcp.py +++ b/litellm/types/mcp.py @@ -2,7 +2,7 @@ import enum from typing import TYPE_CHECKING, Any, Final, Literal from pydantic import BaseModel -from typing_extensions import TypedDict +from typing_extensions import ReadOnly, TypedDict from litellm.types.llms.base import HiddenParams @@ -222,8 +222,24 @@ class MCPCredentials(TypedDict, total=False): top-level request field. """ + oauth_token_header: ReadOnly[str | None] + """ + Header the gateway-minted upstream OAuth token is written to. Defaults to ``Authorization``. + Set it to another header (e.g. ``x-upstream-oauth``) when the upstream reads the minted token + beside a different static ``Authorization`` credential of its own, which would otherwise + collide with it. Not a secret; stored unencrypted. + """ -MCP_ADMIN_CONFIG_CREDENTIAL_KEYS: Final[tuple[str, ...]] = ("upstream_resource",) + +DEFAULT_OAUTH_TOKEN_HEADER: Final = "Authorization" + + +def resolve_oauth_token_header(configured: str | None) -> str: + """The header a gateway-minted OAuth token is written to, defaulting to ``Authorization``.""" + return (configured or "").strip() or DEFAULT_OAUTH_TOKEN_HEADER + + +MCP_ADMIN_CONFIG_CREDENTIAL_KEYS: Final[tuple[str, ...]] = ("upstream_resource", "oauth_token_header") """Non-secret credential keys returned on read so the admin form can show and clear them. Mirrors ``ADMIN_CONFIG_CREDENTIAL_KEYS`` in ``ui/litellm-dashboard/src/components/mcp_tools/types.tsx``.""" diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index d09503cdc4d..6cc88e2774d 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -86,6 +86,10 @@ class MCPServer(BaseModel): # today's behavior; "auto" derives the canonical URI from ``url``; any other value is sent # verbatim. Resolved by ``oauth_utils.resolve_upstream_resource``. upstream_resource: str | None = None + # Header the gateway-minted upstream OAuth token is written to. None means ``Authorization``, + # today's behavior. A custom header keeps the minted token from colliding with a static + # ``Authorization`` the upstream also requires. + oauth_token_header: str | None = None # AWS SigV4 fields aws_access_key_id: str | None = None aws_secret_access_key: str | None = None diff --git a/tests/test_litellm/experimental_mcp_client/test_mcp_client.py b/tests/test_litellm/experimental_mcp_client/test_mcp_client.py index 7beb1c43a94..34a4c3bab33 100644 --- a/tests/test_litellm/experimental_mcp_client/test_mcp_client.py +++ b/tests/test_litellm/experimental_mcp_client/test_mcp_client.py @@ -325,6 +325,44 @@ class TestMCPClient: assert headers["Authorization"] == "token my-token" assert headers["X-Custom-Header"] == "custom-value" + def test_oauth_token_lands_on_the_configured_header_beside_a_static_authorization(self): + """A gateway can require its own static Authorization plus the minted OAuth token on another + header; the static entry must not shadow the token and the token must not displace it.""" + client = MCPClient( + server_url="http://example.com/mcp", + transport_type="http", + auth_type=MCPAuth.oauth2, + auth_value="minted-m2m", + oauth_token_header="esb-oauth", + extra_headers={"Authorization": "Bearer static-pat", "esb-oauth": "placeholder", "envlbl": "prod"}, + ) + + headers = client._get_auth_headers() + + assert headers == { + "Authorization": "Bearer static-pat", + "esb-oauth": "Bearer minted-m2m", + "envlbl": "prod", + } + + def test_oauth_token_defaults_to_authorization_and_still_defers_to_static_headers(self): + client = MCPClient( + server_url="http://example.com/mcp", + transport_type="http", + auth_type=MCPAuth.oauth2, + auth_value="minted-m2m", + ) + assert client._get_auth_headers() == {"Authorization": "Bearer minted-m2m"} + + with_static = MCPClient( + server_url="http://example.com/mcp", + transport_type="http", + auth_type=MCPAuth.oauth2, + auth_value="minted-m2m", + extra_headers={"Authorization": "Bearer static-pat"}, + ) + assert with_static._get_auth_headers() == {"Authorization": "Bearer static-pat"} + def test_get_auth_headers_strips_static_header_whitespace(self): """ Static header names/values must be stripped of surrounding whitespace. diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_adapter.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_adapter.py index e336bdc80c2..40f137d2c6e 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_adapter.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_adapter.py @@ -185,6 +185,37 @@ def test_client_credentials_resolves_upstream_resource_onto_the_config(): assert spec.config.upstream_resource == "https://up.example.com/mcp" +def test_client_credentials_carries_the_configured_token_header(): + spec = to_server_spec( + _server( + auth_type=MCPAuth.oauth2, + oauth2_flow="client_credentials", + token_url="https://idp.example.com/token", + client_id="cid", + client_secret="csec", + oauth_token_header="esb-oauth", + ) + ) + assert spec is not None + assert isinstance(spec.config, ClientCredentialsConfig) + assert spec.config.token_header == "esb-oauth" + + +def test_client_credentials_token_header_defaults_to_authorization(): + spec = to_server_spec( + _server( + auth_type=MCPAuth.oauth2, + oauth2_flow="client_credentials", + token_url="https://idp.example.com/token", + client_id="cid", + client_secret="csec", + ) + ) + assert spec is not None + assert isinstance(spec.config, ClientCredentialsConfig) + assert spec.config.token_header == "Authorization" + + def test_client_credentials_with_incomplete_grant_fields_is_owned_for_fail_closed(): # An M2M server missing its grant fields is still owned by v2 (spec, not None) so it fails # closed at the source (misconfigured, 500) rather than deferring to v1, which would connect diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_client_credentials.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_client_credentials.py index a5d17428b37..06df092655a 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_client_credentials.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_client_credentials.py @@ -425,3 +425,29 @@ def test_bearer_auth_rejects_sync_clients(): with httpx.Client(transport=httpx.MockTransport(lambda request: httpx.Response(200)), auth=auth) as client: with pytest.raises(RuntimeError): client.get("https://upstream.example.com/mcp") + + +@pytest.mark.asyncio +async def test_bearer_auth_writes_a_custom_header_and_leaves_authorization_alone(): + # A gateway that reads the minted token beside its own static Authorization: the minted token + # must land on the configured header and must not displace the static credential. + seen: "list[tuple[str, str]]" = [] + + def handler(request: httpx.Request) -> httpx.Response: + seen.append((request.headers.get("Authorization", ""), request.headers.get("esb-oauth", ""))) + return httpx.Response(401 if len(seen) == 1 else 200) + + async def refetch(failed: str) -> "str | None": + return "fresh-token" + + auth = ClientCredentialsBearerAuth("stale-token", refetch, header_name="esb-oauth") + async with httpx.AsyncClient(transport=httpx.MockTransport(handler), auth=auth) as client: + response = await client.get( + "https://upstream.example.com/mcp", + headers={"Authorization": "Bearer static-pat"}, + ) + assert response.status_code == 200 + assert seen == [ + ("Bearer static-pat", "Bearer stale-token"), + ("Bearer static-pat", "Bearer fresh-token"), + ] diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_resolver.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_resolver.py index 0d130767bd5..500ec87f1c5 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_resolver.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_resolver.py @@ -400,6 +400,19 @@ async def test_client_credentials_emits_the_minted_bearer(): assert source.gets == ["s"] +@pytest.mark.asyncio +async def test_client_credentials_emits_the_minted_token_on_the_configured_header(): + source = _FakeM2MSource(Ok(OAuthToken(access_token="m2m-at"))) + config = _M2M.model_copy(update={"token_header": "esb-oauth"}) + result = await UpstreamCredentialProvider(client_credentials_source=source).resolve_credentials( + _SUBJECT, _spec(config) + ) + assert isinstance(result, Ok) + headers, _ = await _emitted_async(result.ok) + assert headers["esb-oauth"] == "Bearer m2m-at" + assert "Authorization" not in headers + + @pytest.mark.asyncio async def test_client_credentials_ignores_the_subject(): # The contract's no-user-context clause: every caller shares the one client identity. diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 5ee8143fb8e..7e07a4a109d 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -44,7 +44,7 @@ from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( _obo_retry_applies, _resolve_openapi_tool_auth, _should_strip_caller_authorization, - _without_authorization, + _without_header, ) from litellm.proxy._types import ( LiteLLM_MCPServerTable, @@ -968,6 +968,18 @@ class TestMCPServerManager: assert server.oauth2_flow == "client_credentials" assert server.has_client_credentials is True + @pytest.mark.asyncio + async def test_load_servers_from_config_reads_the_oauth_token_header(self): + manager = MCPServerManager() + + with patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=None)): + await manager.load_servers_from_config( + self._oauth2_config(oauth2_flow="client_credentials", oauth_token_header="esb-oauth") + ) + + server = next(iter(manager.config_mcp_servers.values())) + assert server.oauth_token_header == "esb-oauth" + @pytest.mark.asyncio async def test_load_servers_from_config_accepts_explicit_authorization_code(self): manager = MCPServerManager() @@ -2406,6 +2418,42 @@ class TestMCPServerManager: assert client._resolved_auth is not None assert "authorization" not in {k.lower() for k in (client.extra_headers or {})} + @pytest.mark.asyncio + async def test_minted_token_on_a_custom_header_leaves_a_static_authorization_intact(self): + """A gateway can require its own static Authorization plus the minted token on another + header. Only the header the resolved auth writes is dropped from the static set, so the + static credential the upstream also needs still ships.""" + from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import ( + StaticHeaderAuth, + ) + from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Ok + + class _FakeProvider: + async def resolve_credentials(self, subject, server): + return Ok(StaticHeaderAuth("Bearer MINTED-M2M", header_name="esb-oauth")) + + manager = MCPServerManager(cred_provider=_FakeProvider()) + server = MCPServer( + server_id="m2m-custom-header", + name="m2m-custom-header-server", + url="https://up.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow="client_credentials", + client_id="cid", + client_secret="csec", + token_url="https://idp.example.com/token", + oauth_token_header="esb-oauth", + ) + + client = await manager._create_mcp_client( + server, + extra_headers={"Authorization": "Bearer static-pat", "esb-oauth": "placeholder", "envlbl": "prod"}, + ) + + assert client._resolved_auth is not None + assert client.extra_headers == {"Authorization": "Bearer static-pat", "envlbl": "prod"} + @pytest.mark.asyncio async def test_preflight_token_exchange_challenges_on_rejected_subject(self): """A subject the IdP rejects must raise the RFC 9728 401 challenge from the preflight, so a @@ -2625,14 +2673,18 @@ class TestMCPServerManager: if captured_extra_headers: assert "authorization" not in {k.lower() for k in captured_extra_headers} - def test_without_authorization_drops_only_the_credential(self): + def test_without_header_drops_only_the_credential(self): # None / empty -> None - assert _without_authorization(None) is None - assert _without_authorization({}) is None + assert _without_header(None) is None + assert _without_header({}) is None # Only Authorization present -> nothing left -> None (case-insensitive) - assert _without_authorization({"authorization": "Bearer x"}) is None + assert _without_header({"authorization": "Bearer x"}) is None # Authorization dropped, other headers kept - assert _without_authorization({"Authorization": "Bearer x", "X-Trace-Id": "t"}) == {"X-Trace-Id": "t"} + assert _without_header({"Authorization": "Bearer x", "X-Trace-Id": "t"}) == {"X-Trace-Id": "t"} + # A custom header name leaves a static Authorization intact + assert _without_header({"Authorization": "Bearer pat", "esb-oauth": "Bearer stale"}, "ESB-OAuth") == { + "Authorization": "Bearer pat" + } @pytest.mark.asyncio async def test_call_regular_mcp_tool_passthrough_forwards_authorization_with_admission_header( diff --git a/ui/litellm-dashboard/src/components/mcp_tools/types.tsx b/ui/litellm-dashboard/src/components/mcp_tools/types.tsx index 67e51ab23aa..35adfce9ad9 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/types.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/types.tsx @@ -151,7 +151,7 @@ const DECLARED_APP_CREDENTIAL_KEYS = ["client_id", "client_secret"] as const; // would destroy admin input), but it must stay OUT of the declared-app set: whether an app exists is // a distinct question that gates the "app may not match upstream" warning, and a server using dynamic // client registration can set a resource indicator while having no app at all. -export const ADMIN_CONFIG_CREDENTIAL_KEYS = ["upstream_resource"] as const; +export const ADMIN_CONFIG_CREDENTIAL_KEYS = ["upstream_resource", "oauth_token_header"] as const; // Minted token material the oauth2 authorize path writes beside the app keys; stripped from restored // snapshots and from any credentials that transit to the temp-session preview so a stale token never