fix(mcp): stop caller-supplied auth from overriding stored authorization_code tokens

A caller-supplied per-request override (mcp_auth_header / x-mcp-auth / x-mcp-<alias>-authorization) disabled the v2 resolver in _create_mcp_client for any spec, so an authenticated user with a stored authorization_code token could force an arbitrary upstream bearer and bypass the stored credential and its save-time validation. _create_mcp_client now keeps the v2 spec for authorization_code and ignores the override; other modes keep the client-side-credentials override

The create/test tools preview no longer relies on that override path. It resolves the just-authorized, not-yet-persisted token through the v2 resolver via a one-shot PresentedOAuthTokenStore passed as cred_provider - the same path runtime uses for the stored token - so preview and runtime resolve identically. This replaces the mcp_auth_header routing added earlier

Adds tests: a caller override cannot bypass the v2 resolver for authorization_code; the interactive preview resolves via the presented store rather than a caller header; M2M and token-exchange build no presented provider
This commit is contained in:
Tin Chi Lo 2026-06-26 20:35:55 -07:00
parent f028f32c9e
commit 84b77fc72c
6 changed files with 170 additions and 56 deletions

View file

@ -70,6 +70,9 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import
from litellm.proxy._experimental.mcp_server.outbound_credentials.per_user_oauth_store import (
LazyPerUserOAuthTokenStore,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
AuthorizationCodeConfig,
)
from litellm.proxy._experimental.mcp_server.utils import (
MCP_TOOL_PREFIX_SEPARATOR,
MCPMissingUserEnvVarsError,
@ -1965,6 +1968,7 @@ class MCPServerManager:
stdio_env: Optional[Dict[str, str]] = None,
subject_token: Optional[str] = None,
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
cred_provider: Optional[UpstreamCredentialProvider] = None,
) -> MCPClient:
"""
Create an MCPClient instance for the given server.
@ -1988,11 +1992,17 @@ class MCPServerManager:
"""
transport = server.transport or MCPTransport.sse
spec = None if transport == MCPTransport.stdio else to_server_spec(server)
# 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:
provider = cred_provider or self._cred_provider
# A caller-supplied per-request override (mcp_auth_header / x-mcp-*) defers to the v1 path
# so it wins - except for authorization_code, whose per-user token the v2 resolver owns. A
# caller must not be able to substitute another user's stored credential, so we keep the v2
# spec and ignore the override there; the REST tools preview supplies its not-yet-persisted
# token through the resolver (cred_provider), never this path.
if (
spec is not None
and mcp_auth_header
and not isinstance(spec.config, AuthorizationCodeConfig)
):
spec = None
auth_value = (
await resolve_mcp_auth(server, mcp_auth_header, subject_token=subject_token)
@ -2070,7 +2080,7 @@ class MCPServerManager:
server_url = server.url or ""
if spec is not None:
match await self._cred_provider.resolve_credentials(
match await provider.resolve_credentials(
to_subject(user_api_key_auth, subject_token), spec
):
case Ok(auth):

View file

@ -0,0 +1,26 @@
"""One-shot ``OAuthTokenStore`` for the create/test tools preview.
The preview tests an unsaved server, so no per-user credential is persisted yet. The operator holds
the just-authorized token; this serves it through the same v2 resolver path runtime uses for the
stored token, so the preview never relies on the caller-credential-override path that
``_create_mcp_client`` refuses for ``authorization_code``. It backs a single preview call, so it
returns its one token regardless of the lookup key.
"""
from __future__ import annotations
from dataclasses import dataclass
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
OAuthToken,
)
@dataclass(frozen=True, slots=True)
class PresentedOAuthTokenStore:
"""Serves one in-hand token for the single preview call it backs (no DB, no cache)."""
token: OAuthToken
async def fetch(self, user_id: str, server_id: str) -> OAuthToken | None:
return self.token

View file

@ -1034,42 +1034,63 @@ if MCP_AVAILABLE:
None if server_model.has_client_credentials else oauth2_headers
)
merged_headers = merge_mcp_headers(
extra_headers=effective_oauth2_headers,
static_headers=request.static_headers,
# Interactive authorization_code tools preview: the operator holds a just-authorized
# token but it is not persisted yet. Resolve it through the v2 resolver via a one-shot
# presented store - the same path runtime uses for the stored token - rather than the
# caller-override path _create_mcp_client refuses for authorization_code. The bare token
# becomes the upstream credential, so it is not also forwarded as a caller header. Gated
# to the v2-mapped oauth2 case (to_server_spec non-None); M2M (client_credentials),
# delegate/passthrough, and token-exchange are unaffected.
from litellm.proxy._experimental.mcp_server.outbound_credentials import ( # noqa: PLC0415
UpstreamCredentialProvider,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import ( # noqa: PLC0415
to_server_spec,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import ( # noqa: PLC0415
OAuthToken,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.presented_token_store import ( # noqa: PLC0415
PresentedOAuthTokenStore,
)
# Interactive authorization_code tools preview: the UI forwards the just-authorized token
# in oauth2_headers before any per-user credential is persisted. Route it through
# mcp_auth_header so _create_mcp_client's per-request-override deferral takes the v1 path
# (use the forwarded token directly - no resolver, no fail-closed challenge), exactly as
# v1's preview did. Gated to the v2-mapped oauth2 case (to_server_spec non-None ==
# authorization_code): M2M (client_credentials) and delegate/passthrough return None from
# to_server_spec, and token-exchange is a different auth_type, so all three keep their
# existing v1 preview behavior untouched.
if (
mcp_auth_header is None
and server_model.auth_type == MCPAuth.oauth2
and oauth2_headers
):
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import ( # noqa: PLC0415
to_server_spec,
)
forwarded_authorization = oauth2_headers.get("Authorization")
if forwarded_authorization and to_server_spec(server_model) is not None:
# The oauth2 client re-adds the "Bearer " scheme, so pass the bare token.
mcp_auth_header = (
forwarded_authorization[7:]
if forwarded_authorization[:7].lower() == "bearer "
else forwarded_authorization
forwarded_authorization = (
effective_oauth2_headers.get("Authorization")
if effective_oauth2_headers
else None
)
is_interactive_authz_code = (
server_model.auth_type == MCPAuth.oauth2
and forwarded_authorization is not None
and to_server_spec(server_model) is not None
)
preview_cred_provider = (
UpstreamCredentialProvider(
oauth_token_store=PresentedOAuthTokenStore(
OAuthToken(
access_token=forwarded_authorization[7:]
if forwarded_authorization[:7].lower() == "bearer "
else forwarded_authorization
)
)
)
if is_interactive_authz_code
else None
)
merged_headers = merge_mcp_headers(
extra_headers=(
None if preview_cred_provider else effective_oauth2_headers
),
static_headers=request.static_headers,
)
client = await global_mcp_server_manager._create_mcp_client(
server=server_model,
mcp_auth_header=mcp_auth_header,
extra_headers=merged_headers,
stdio_env=stdio_env,
cred_provider=preview_cred_provider,
)
return await operation(client)

View file

@ -0,0 +1,19 @@
"""Tests for the create/test-preview presented OAuth token store."""
import pytest
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
OAuthToken,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.presented_token_store import (
PresentedOAuthTokenStore,
)
@pytest.mark.asyncio
async def test_serves_the_presented_token_regardless_of_key():
token = OAuthToken(access_token="at", scopes=("read",))
store = PresentedOAuthTokenStore(token)
# one-shot store backs a single preview call, so the lookup key is irrelevant
assert await store.fetch("alice", "srv-1") is token
assert await store.fetch("", "") is token

View file

@ -139,6 +139,43 @@ class TestMCPServerManager:
assert client.stdio_config["env"]["NODE_ENV"] == "test"
assert client.stdio_config["env"]["NPM_CONFIG_CACHE"] == MCP_NPM_CACHE_DIR
@pytest.mark.asyncio
async def test_caller_auth_header_cannot_bypass_v2_for_authorization_code(self):
"""A caller-supplied per-request override must not substitute the stored authorization_code
token: _create_mcp_client keeps the v2 spec and resolves through the injected provider
rather than deferring to the v1 caller-override path."""
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import (
StaticHeaderAuth,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import (
Ok,
)
from litellm.types.mcp import MCPAuth
calls = []
class _FakeProvider:
async def resolve_credentials(self, subject, server):
calls.append((subject.subject_id, server.server_id))
return Ok(StaticHeaderAuth("stored-token"))
manager = MCPServerManager(cred_provider=_FakeProvider())
server = MCPServer(
server_id="authz-srv",
name="authz",
url="https://upstream.example/mcp",
transport=MCPTransport.sse,
auth_type=MCPAuth.oauth2, # oauth2 + no client creds + not delegate -> authorization_code
)
client = await manager._create_mcp_client(
server, mcp_auth_header="Bearer caller-supplied-token"
)
# the v2 resolver ran (the caller override did NOT defer to v1); the stored token wins
assert calls == [("", "authz-srv")]
assert client is not None
async def test_create_mcp_client_stdio_injects_npm_config_cache(self):
"""Test that _create_mcp_client injects NPM_CONFIG_CACHE when not already set,
and preserves user-provided NPM_CONFIG_CACHE when present."""

View file

@ -271,21 +271,20 @@ class TestExecuteWithMcpClient:
)
@pytest.mark.asyncio
async def test_interactive_oauth_routes_forwarded_token_to_mcp_auth_header(
async def test_interactive_oauth_resolves_forwarded_token_via_presented_store(
self, monkeypatch
):
"""Interactive authorization_code preview (oauth2, no client credentials): the forwarded
just-authorized token must be routed to mcp_auth_header so _create_mcp_client takes the v1
per-request-override path and uses it directly, instead of the v2 resolver fail-closing on
the not-yet-persisted credential. The "Bearer " scheme is stripped since the oauth2 client
re-adds it."""
just-authorized token is resolved THROUGH the v2 resolver via a one-shot presented store
(cred_provider), not the caller-override path. The bare token (Bearer stripped) is the
upstream credential and is not also forwarded in extra_headers."""
captured: dict = {}
def fake_build_stdio_env(server, raw_headers):
return None
async def fake_create_client(*args, **kwargs):
captured["mcp_auth_header"] = kwargs.get("mcp_auth_header")
captured.update(kwargs)
return object()
monkeypatch.setattr(
@ -318,22 +317,27 @@ class TestExecuteWithMcpClient:
)
assert result["status"] == "ok"
assert captured["mcp_auth_header"] == "forwarded-user-token"
# Resolved via the v2 resolver, never the caller-override header
assert captured["mcp_auth_header"] is None
provider = captured["cred_provider"]
assert provider is not None
token = await provider._oauth_token_store.fetch("u", "s")
assert token is not None and token.access_token == "forwarded-user-token"
# The resolver supplies the bearer, so it is not also forwarded as a caller header
extra_headers = captured.get("extra_headers") or {}
assert not any(k.lower() == "authorization" for k in extra_headers)
@pytest.mark.asyncio
async def test_m2m_does_not_route_forwarded_token_to_mcp_auth_header(
self, monkeypatch
):
"""M2M (client_credentials) must NOT route the forwarded header to mcp_auth_header - it stays
on the auto-fetch path. to_server_spec returns None for M2M, so the interactive shortcut is
skipped and mcp_auth_header remains None (the incoming header is dropped, as before)."""
async def test_m2m_does_not_build_presented_store(self, monkeypatch):
"""M2M (client_credentials): to_server_spec returns None, so no presented provider is built;
the auto-fetch path is unchanged (no cred_provider, the incoming header dropped as before)."""
captured: dict = {}
def fake_build_stdio_env(server, raw_headers):
return None
async def fake_create_client(*args, **kwargs):
captured["mcp_auth_header"] = kwargs.get("mcp_auth_header")
captured.update(kwargs)
return object()
monkeypatch.setattr(
@ -367,23 +371,20 @@ class TestExecuteWithMcpClient:
)
assert result["status"] == "ok"
assert captured.get("cred_provider") is None
assert captured["mcp_auth_header"] is None
@pytest.mark.asyncio
async def test_token_exchange_does_not_route_forwarded_token_to_mcp_auth_header(
self, monkeypatch
):
"""OBO / token-exchange (auth_type oauth2_token_exchange, not oauth2) must NOT route the
forwarded header to mcp_auth_header - resolve_mcp_auth performs the exchange on the v1 path.
The guard's auth_type == oauth2 check excludes it, so mcp_auth_header stays None (setting it
would short-circuit the exchange, since resolve_mcp_auth returns mcp_auth_header first)."""
async def test_token_exchange_does_not_build_presented_store(self, monkeypatch):
"""OBO / token-exchange (auth_type oauth2_token_exchange, not oauth2): excluded by the
auth_type == oauth2 guard, so no presented provider is built and the v1 exchange path runs."""
captured: dict = {}
def fake_build_stdio_env(server, raw_headers):
return None
async def fake_create_client(*args, **kwargs):
captured["mcp_auth_header"] = kwargs.get("mcp_auth_header")
captured.update(kwargs)
return object()
monkeypatch.setattr(
@ -416,7 +417,7 @@ class TestExecuteWithMcpClient:
)
assert result["status"] == "ok"
assert captured["mcp_auth_header"] is None
assert captured.get("cred_provider") is None
@pytest.mark.asyncio
async def test_catches_exception_group(self, monkeypatch):