refactor(mcp): resolve inbound-header conflict on v2 instead of deferring

For an Authorization already supplied via extra_headers (a guardrail hook such as the
JWT signer, static_headers, or a forwarded caller header), keep the request on the v2
path and skip resolved_auth rather than deferring to v1. The inbound header still wins
since nothing overwrites it, but hooks no longer pin a v1 fallback, which is what lets
resolve_mcp_auth be retired once the remaining modes migrate.

The mcp_auth_header per-request override still defers to v1, since that value becomes
the upstream credential rather than sitting in extra_headers; that defer falls away
once the per-user modes stop writing mcp_auth_header.
This commit is contained in:
Tin Chi Lo 2026-06-23 18:45:22 -07:00
parent 62ab28eba8
commit 9c8be2cbba
2 changed files with 23 additions and 23 deletions

View file

@ -57,7 +57,6 @@ from litellm.proxy._experimental.mcp_server.sampling_handler import (
)
from litellm.proxy._experimental.mcp_server.oauth2_token_cache import resolve_mcp_auth
from litellm.proxy._experimental.mcp_server.outbound_credentials import (
ApiKeyConfig,
Error,
Ok,
UpstreamCredentialProvider,
@ -1956,23 +1955,11 @@ class MCPServerManager:
"""
transport = server.transport or MCPTransport.sse
spec = None if transport == MCPTransport.stdio else to_server_spec(server)
# Credential-isolation invariant (mirrors the v2 egress path): the resolved credential
# rides the httpx auth flow, which writes its header after extra_headers, so it would
# overwrite an inbound credential. Defer to v1 when a per-request override is present, or
# when the credential's header is already supplied via extra_headers (guardrail hook,
# static_headers, or a forwarded caller header) — v1 lets those win. ``none`` writes no
# header, so it never conflicts.
if spec is not None and (
mcp_auth_header
or (
isinstance(spec.config, ApiKeyConfig)
and extra_headers
and any(
key.lower() == spec.config.header_name.lower()
for key in extra_headers
)
)
):
# A per-request override is the caller-supplied credential v1 turns into the upstream
# auth, so it must win; defer those to v1 (this defer falls away once the per-user modes
# stop writing mcp_auth_header). An inbound header already in extra_headers is handled on
# the v2 path below, not here.
if spec is not None and mcp_auth_header:
spec = None
auth_value = (
await resolve_mcp_auth(server, mcp_auth_header, subject_token=subject_token)
@ -2055,6 +2042,20 @@ class MCPServerManager:
):
case Ok(auth):
resolved_auth = auth
# Do not override an Authorization already supplied via extra_headers
# (a guardrail hook such as the JWT signer, static_headers, or a
# forwarded caller header): v1 applies those last, so they win. NoOpAuth
# has no header_name and so never skips.
header_name = getattr(resolved_auth, "header_name", None)
if (
header_name
and extra_headers
and any(
key.lower() == header_name.lower()
for key in extra_headers
)
):
resolved_auth = None
case Error(err):
raise_public(err)
return MCPClient(

View file

@ -4705,10 +4705,10 @@ class TestCreateMcpClientV2Graft:
assert client._resolved_auth is None
assert client._mcp_auth_value == "caller-override"
async def test_conflicting_extra_header_defers_to_v1(self):
async def test_conflicting_extra_header_skips_resolved_auth_on_v2(self):
# An Authorization already supplied via extra_headers (guardrail hook like the JWT
# signer, static_headers, or a forwarded caller header) must not be clobbered by the
# resolved static credential, so the static server defers to v1.
# signer, static_headers, or a forwarded caller header) must win. The server stays on
# the v2 path but skips resolved_auth, so nothing overwrites the inbound header.
client = await MCPServerManager()._create_mcp_client(
self._http_server(
auth_type=MCPAuth.bearer_token, authentication_token="shared-tok"
@ -4717,8 +4717,7 @@ class TestCreateMcpClientV2Graft:
)
assert client._resolved_auth is None
assert client._mcp_auth_value == "shared-tok"
# v1 applies extra_headers last, so the inbound header wins on the wire.
assert client._mcp_auth_value is None
assert client._get_auth_headers()["Authorization"] == "Bearer hook-jwt"
async def test_none_with_extra_header_stays_v2_without_clobbering(self):