mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
f028f32c9e
commit
84b77fc72c
6 changed files with 170 additions and 56 deletions
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue