Merge pull request #34407 from BerriAI/litellm_mcp_v1_obo_deadcode_cleanup

refactor(mcp): delete unreachable v1 OBO handler and gate REST OAuth on v2 resolver
This commit is contained in:
tin-berri 2026-07-23 13:06:47 -07:00 committed by GitHub
commit df3050f538
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
9 changed files with 151 additions and 814 deletions

View file

@ -1,192 +0,0 @@
"""
OAuth 2.0 Token Exchange (RFC 8693) handler for MCP servers.
Exchanges a user's incoming JWT (subject_token) for a scoped access token
at an IDP's token exchange endpoint. The exchanged token is then used to
authenticate requests to the upstream MCP server.
See: https://datatracker.ietf.org/doc/html/rfc8693
"""
import asyncio
import hashlib
import weakref
from typing import TYPE_CHECKING, Dict, Tuple
import httpx
from litellm._logging import verbose_logger
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.constants import (
MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL,
MCP_OAUTH2_TOKEN_CACHE_MIN_TTL,
MCP_OAUTH2_TOKEN_EXPIRY_BUFFER_SECONDS,
MCP_TOKEN_EXCHANGE_CACHE_MAX_SIZE,
)
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import (
build_token_endpoint_client_auth,
)
from litellm.types.llms.custom_http import httpxSpecialProvider
from litellm.types.mcp import DEFAULT_SUBJECT_TOKEN_TYPE
if TYPE_CHECKING:
from litellm.types.mcp_server.mcp_server_manager import MCPServer
# RFC 8693 grant type constant
TOKEN_EXCHANGE_GRANT_TYPE = "urn:ietf:params:oauth:grant-type:token-exchange"
class TokenExchangeHandler:
"""Handles OAuth 2.0 Token Exchange (RFC 8693) for MCP servers.
Caches exchanged tokens keyed by ``hash(subject_token + server_id)`` so
repeated calls with the same user token skip the IDP round-trip.
"""
def __init__(self) -> None:
self._cache = InMemoryCache(
max_size_in_memory=MCP_TOKEN_EXCHANGE_CACHE_MAX_SIZE,
default_ttl=MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL,
)
# WeakValueDictionary so locks are GC'd once no coroutine holds a reference,
# preventing unbounded growth with many rotating user tokens.
self._locks: weakref.WeakValueDictionary[str, asyncio.Lock] = weakref.WeakValueDictionary()
def _get_lock(self, cache_key: str) -> asyncio.Lock:
lock = self._locks.get(cache_key)
if lock is None:
lock = asyncio.Lock()
self._locks[cache_key] = lock
return lock
@staticmethod
def _cache_key(subject_token: str, server_id: str) -> str:
raw = f"{subject_token}:{server_id}"
return hashlib.sha256(raw.encode()).hexdigest()
async def exchange_token(
self,
subject_token: str,
server: "MCPServer",
) -> str:
"""Exchange *subject_token* for a scoped access token.
Returns the exchanged ``access_token`` string (suitable for a
``Bearer`` header).
Raises ``ValueError`` on configuration or IDP errors.
"""
cache_key = self._cache_key(subject_token, server.server_id)
# Fast path
cached = self._cache.get_cache(cache_key)
if cached is not None:
return cached
# Slow path — one exchange at a time per (user, server) pair
async with self._get_lock(cache_key):
cached = self._cache.get_cache(cache_key)
if cached is not None:
return cached
token, ttl = await self._do_exchange(subject_token, server)
self._cache.set_cache(cache_key, token, ttl=ttl)
return token
async def _do_exchange(
self,
subject_token: str,
server: "MCPServer",
) -> Tuple[str, int]:
"""POST to the token exchange endpoint with RFC 8693 parameters.
Returns ``(access_token, ttl_seconds)``.
"""
endpoint = server.token_exchange_endpoint or server.token_url
if not endpoint:
raise ValueError(
f"MCP server '{server.server_id}' has auth_type=oauth2_token_exchange "
f"but no token_exchange_endpoint or token_url configured"
)
if not server.client_id or not server.client_secret:
raise ValueError(
f"MCP server '{server.server_id}' has auth_type=oauth2_token_exchange "
f"but missing client_id or client_secret"
)
client_auth = build_token_endpoint_client_auth(
auth_method=server.token_endpoint_auth_method,
client_id=server.client_id,
client_secret=server.client_secret,
)
data: Dict[str, str] = {
"grant_type": TOKEN_EXCHANGE_GRANT_TYPE,
"subject_token": subject_token,
"subject_token_type": server.subject_token_type or DEFAULT_SUBJECT_TOKEN_TYPE,
**client_auth.body,
}
if server.audience:
data["audience"] = server.audience
if server.scopes:
data["scope"] = " ".join(server.scopes)
verbose_logger.debug(
"Exchanging token for MCP server %s at %s (audience=%s)",
server.server_id,
endpoint,
server.audience,
)
client = get_async_httpx_client(llm_provider=httpxSpecialProvider.MCP)
post_kwargs = {"data": data, **({"headers": client_auth.headers} if client_auth.headers else {})}
try:
response = await client.post(endpoint, **post_kwargs)
response.raise_for_status()
except httpx.HTTPStatusError as exc:
verbose_logger.debug(
"Token exchange IDP error for MCP server %s (status %d)",
server.server_id,
exc.response.status_code,
)
raise ValueError(
f"Token exchange for MCP server '{server.server_id}' failed with status {exc.response.status_code}"
) from exc
body = response.json()
if not isinstance(body, dict):
raise ValueError(
f"Token exchange response for MCP server '{server.server_id}' "
f"returned non-object JSON (got {type(body).__name__})"
)
access_token = body.get("access_token")
if not access_token:
raise ValueError(f"Token exchange response for MCP server '{server.server_id}' missing 'access_token'")
raw_expires_in = body.get("expires_in")
try:
expires_in = int(raw_expires_in) if raw_expires_in is not None else MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL
except (TypeError, ValueError):
expires_in = MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL
ttl = max(
expires_in - MCP_OAUTH2_TOKEN_EXPIRY_BUFFER_SECONDS,
MCP_OAUTH2_TOKEN_CACHE_MIN_TTL,
)
verbose_logger.info(
"Token exchange succeeded for MCP server %s (expires in %ds)",
server.server_id,
expires_in,
)
return access_token, ttl
def invalidate(self, subject_token: str, server_id: str) -> None:
"""Remove a cached exchanged token (e.g. after a 401)."""
cache_key = self._cache_key(subject_token, server_id)
self._cache.delete_cache(cache_key)
# Module-level singleton
mcp_token_exchange_handler = TokenExchangeHandler()

View file

@ -3086,9 +3086,7 @@ class MCPServerManager:
)
):
spec = None
auth_value = (
await resolve_mcp_auth(server, mcp_auth_header, subject_token=subject_token) if spec is None else None
)
auth_value = await resolve_mcp_auth(server, mcp_auth_header) if spec is None else None
# Create sampling and elicitation callbacks for this client
sampling_cb = _create_sampling_callback(user_api_key_auth=user_api_key_auth) if server.allow_sampling else None

View file

@ -26,7 +26,6 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import (
decrypt_value_helper,
encrypt_value_helper,
)
from litellm.proxy._experimental.mcp_server.auth import token_exchange
from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import (
build_token_endpoint_client_auth,
)
@ -58,17 +57,12 @@ class MCPOAuth2TokenCache(InMemoryCache):
def _has_client_credentials_config(server: "MCPServer") -> bool:
return bool(server.client_id and server.client_secret and server.token_url)
async def async_get_token(
self,
server: "MCPServer",
*,
require_client_credentials_flow: bool = True,
) -> Optional[str]:
async def async_get_token(self, server: "MCPServer") -> Optional[str]:
"""Return a valid access token, fetching or refreshing as needed.
Returns ``None`` when the server lacks client credentials config.
"""
if require_client_credentials_flow and not server.has_client_credentials:
if not server.has_client_credentials:
return None
if not self._has_client_credentials_config(server):
return None
@ -278,36 +272,16 @@ mcp_per_user_token_cache = MCPPerUserTokenCache()
async def resolve_mcp_auth(
server: "MCPServer",
mcp_auth_header: Optional[Union[str, Dict[str, str]]] = None,
subject_token: Optional[str] = None,
) -> Optional[Union[str, Dict[str, str]]]:
"""Resolve the auth value for an MCP server.
Priority:
1. ``mcp_auth_header`` per-request/per-user override
2. OAuth2 Token Exchange (OBO / RFC 8693) exchange user token for scoped token
3. OAuth2 client_credentials token auto-fetched and cached
4. ``server.authentication_token`` static token from config/DB
2. OAuth2 client_credentials token auto-fetched and cached
3. ``server.authentication_token`` static token from config/DB
"""
if mcp_auth_header:
return mcp_auth_header
if server.has_token_exchange_config:
if subject_token:
return await token_exchange.mcp_token_exchange_handler.exchange_token(subject_token, server)
# No subject_token — fall back to client_credentials using the same client
# credentials and token_url so M2M scenarios still work.
if server.client_id and server.client_secret and server.token_url:
return await mcp_oauth2_token_cache.async_get_token(
server,
require_client_credentials_flow=False,
)
# OBO configured but no subject_token and missing client credentials — warn
# rather than silently proceeding unauthenticated.
verbose_logger.warning(
"MCP server '%s' is configured for token exchange (OBO) but no subject_token "
"was provided and client credentials (client_id/client_secret/token_url) are "
"incomplete. The request will proceed without authentication.",
server.server_id,
)
if server.has_client_credentials:
return await mcp_oauth2_token_cache.async_get_token(server)
return server.authentication_token

View file

@ -230,16 +230,33 @@ if MCP_AVAILABLE:
return server_auth
return mcp_auth_header
def _get_oauth2_server_ids(allowed_server_ids: List[str]) -> Set[str]:
"""Return the subset of *allowed_server_ids* whose servers use OAuth2 auth.
def _is_v1_resolved_oauth2_server(server: Optional[MCPServer]) -> bool:
"""Whether this server's per-user OAuth2 token is still resolved by v1.
Used as a cheap pre-flight check to skip bulk credential fetching when no
OAuth2 servers are involved in the current request.
A server the v2 resolver owns reads its stored token from the resolver at connect
time and drops any Authorization built for it here, so the v1 lookup would be a DB
round-trip whose result is discarded. Mirrors the same guard on the protocol listing
path and in ``_resolve_oauth2_headers_for_tool_call``.
"""
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import (
to_server_spec,
)
if getattr(server, "auth_type", None) != MCPAuth.oauth2:
return False
return to_server_spec(server) is None
def _v1_resolved_oauth2_server_ids(allowed_server_ids: List[str]) -> Set[str]:
"""Return the subset of *allowed_server_ids* whose per-user OAuth2 token is still
resolved by v1.
Used as a cheap pre-flight check to skip bulk credential fetching when no such
server is involved in the current request.
"""
return {
sid
for sid in allowed_server_ids
if getattr(global_mcp_server_manager.get_mcp_server_by_id(sid), "auth_type", None) == MCPAuth.oauth2
if _is_v1_resolved_oauth2_server(global_mcp_server_manager.get_mcp_server_by_id(sid))
}
async def _get_user_oauth_extra_headers(
@ -253,11 +270,13 @@ if MCP_AVAILABLE:
the MCP server the same way the admin "Add MCP / Authorize and Fetch" flow does.
Returns None for non-OAuth2 servers or when no credential is stored.
A server the v2 resolver owns is skipped; see ``_is_v1_resolved_oauth2_server``.
Args:
prefetched_creds: Optional dict keyed by server_id with credential payloads.
When provided, avoids a per-server DB round-trip.
"""
if getattr(server, "auth_type", None) != MCPAuth.oauth2:
if not _is_v1_resolved_oauth2_server(server):
return None
user_id = getattr(user_api_key_dict, "user_id", None)
server_id = getattr(server, "server_id", None)
@ -320,38 +339,6 @@ if MCP_AVAILABLE:
verbose_logger.warning(f"_prefetch_user_oauth_creds: failed to prefetch for user={user_id}: {e}")
return {}
async def _get_bulk_user_oauth_headers(
user_api_key_dict: UserAPIKeyAuth,
) -> Dict[str, Dict[str, str]]:
"""
Fetch ALL OAuth2 credentials for the current user in a single DB query and
return a mapping of server_id {"Authorization": "Bearer <token>"}.
This is the batch alternative to calling _get_user_oauth_extra_headers
per-server inside a loop (N+1 DB queries).
"""
user_id = getattr(user_api_key_dict, "user_id", None)
if not user_id:
return {}
try:
from litellm.proxy._experimental.mcp_server.db import (
list_user_oauth_credentials,
)
from litellm.proxy.utils import get_prisma_client_or_throw
prisma_client = get_prisma_client_or_throw(
"Database not connected. Connect a database to use OAuth2 MCP tools."
)
creds = await list_user_oauth_credentials(prisma_client, user_id)
return {
c["server_id"]: {"Authorization": f"Bearer {c['access_token']}"}
for c in creds
if c.get("access_token") and c.get("server_id")
}
except Exception:
verbose_logger.debug("Failed to bulk-fetch OAuth credentials", exc_info=True)
return {}
def _create_tool_response_objects(tools, server: MCPServer):
"""Helper function to create tool response objects.
@ -825,7 +812,7 @@ if MCP_AVAILABLE:
# to avoid an unnecessary DB round-trip on requests with no OAuth2 MCP servers.
prefetched_oauth_creds = (
await _prefetch_user_oauth_creds(user_api_key_dict)
if _get_oauth2_server_ids(allowed_server_ids)
if _v1_resolved_oauth2_server_ids(allowed_server_ids)
else {}
)

View file

@ -261,12 +261,3 @@ class MCPServer(BaseModel):
if self.oauth_passthrough is not True:
return False
return any(h.lower() == "authorization" for h in self.extra_headers)
@property
def has_token_exchange_config(self) -> bool:
"""True if this server is configured for OAuth2 token exchange (OBO / RFC 8693)."""
return (
self.auth_type == MCPAuth.oauth2_token_exchange
and bool(self.client_id and self.client_secret)
and bool(self.token_exchange_endpoint or self.token_url)
)

View file

@ -1,539 +0,0 @@
"""
Tests for OAuth 2.0 Token Exchange (RFC 8693) handler for MCP servers.
Covers: exchange flow, caching, error handling, resolve_mcp_auth integration,
bearer token extraction, and config loading.
"""
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
from litellm.proxy._experimental.mcp_server.auth.token_exchange import (
TOKEN_EXCHANGE_GRANT_TYPE,
TokenExchangeHandler,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
MCPServerManager,
)
from litellm.proxy._experimental.mcp_server.oauth2_token_cache import (
resolve_mcp_auth,
)
from litellm.proxy._types import LiteLLM_MCPServerTable, MCPTransport
from litellm.types.mcp import MCPAuth
from litellm.types.mcp_server.mcp_server_manager import MCPServer
def _obo_server(**overrides) -> MCPServer:
defaults = dict(
server_id="srv-obo-1",
name="test-obo",
url="https://mcp.example.com/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2_token_exchange,
client_id="litellm-client-id",
client_secret="litellm-client-secret",
token_exchange_endpoint="https://idp.example.com/oauth2/token",
audience="api://mcp-server",
scopes=["mcp.tools.read", "mcp.tools.execute"],
)
defaults.update(overrides)
return MCPServer(**defaults)
def _exchange_response(token="exchanged-tok-abc", expires_in=3600):
resp = MagicMock()
resp.json.return_value = {
"access_token": token,
"token_type": "Bearer",
"expires_in": expires_in,
}
resp.raise_for_status = MagicMock()
resp.text = ""
return resp
# ── Exchange Flow ──
@pytest.mark.asyncio
async def test_exchange_token_success():
"""Token exchange sends correct RFC 8693 parameters and returns access_token."""
handler = TokenExchangeHandler()
server = _obo_server()
mock_client = AsyncMock()
mock_client.post.return_value = _exchange_response("scoped-token-1")
with patch(
"litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client",
return_value=mock_client,
):
result = await handler.exchange_token("user-jwt-xyz", server)
assert result == "scoped-token-1"
mock_client.post.assert_called_once()
_, kwargs = mock_client.post.call_args
data = kwargs["data"]
assert data["grant_type"] == TOKEN_EXCHANGE_GRANT_TYPE
assert data["subject_token"] == "user-jwt-xyz"
assert data["subject_token_type"] == "urn:ietf:params:oauth:token-type:access_token"
assert data["audience"] == "api://mcp-server"
assert data["scope"] == "mcp.tools.read mcp.tools.execute"
assert data["client_id"] == "litellm-client-id"
assert data["client_secret"] == "litellm-client-secret"
@pytest.mark.asyncio
async def test_exchange_token_no_audience():
"""When audience is None, it is omitted from the request."""
handler = TokenExchangeHandler()
server = _obo_server(audience=None)
mock_client = AsyncMock()
mock_client.post.return_value = _exchange_response()
with patch(
"litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client",
return_value=mock_client,
):
await handler.exchange_token("user-jwt", server)
_, kwargs = mock_client.post.call_args
assert "audience" not in kwargs["data"]
@pytest.mark.asyncio
async def test_exchange_token_no_scopes():
"""When scopes is None, scope param is omitted from the request."""
handler = TokenExchangeHandler()
server = _obo_server(scopes=None)
mock_client = AsyncMock()
mock_client.post.return_value = _exchange_response()
with patch(
"litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client",
return_value=mock_client,
):
await handler.exchange_token("user-jwt", server)
_, kwargs = mock_client.post.call_args
assert "scope" not in kwargs["data"]
# ── Caching ──
@pytest.mark.asyncio
async def test_exchange_token_cached():
"""Second call with same user token uses cache — only 1 HTTP POST."""
handler = TokenExchangeHandler()
server = _obo_server()
mock_client = AsyncMock()
mock_client.post.return_value = _exchange_response("cached-exchange-tok")
with patch(
"litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client",
return_value=mock_client,
):
t1 = await handler.exchange_token("same-jwt", server)
t2 = await handler.exchange_token("same-jwt", server)
assert t1 == t2 == "cached-exchange-tok"
assert mock_client.post.call_count == 1
@pytest.mark.asyncio
async def test_different_user_tokens_not_shared():
"""Different user JWTs get different exchanged tokens."""
handler = TokenExchangeHandler()
server = _obo_server()
call_count = 0
async def mock_post(url, data=None):
nonlocal call_count
call_count += 1
resp = MagicMock()
resp.json.return_value = {
"access_token": f"exchanged-{call_count}",
"expires_in": 3600,
}
resp.raise_for_status = MagicMock()
return resp
mock_client = AsyncMock()
mock_client.post = mock_post
with patch(
"litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client",
return_value=mock_client,
):
t1 = await handler.exchange_token("user-a-jwt", server)
t2 = await handler.exchange_token("user-b-jwt", server)
assert t1 == "exchanged-1"
assert t2 == "exchanged-2"
assert call_count == 2
# ── Error Handling ──
@pytest.mark.asyncio
async def test_exchange_token_http_error():
"""HTTP errors from the IDP are wrapped in a ValueError."""
handler = TokenExchangeHandler()
server = _obo_server()
mock_response = MagicMock()
mock_response.status_code = 400
mock_response.text = "invalid_grant"
mock_response.raise_for_status.side_effect = httpx.HTTPStatusError(
"Bad Request",
request=MagicMock(),
response=mock_response,
)
mock_client = AsyncMock()
mock_client.post.return_value = mock_response
with (
patch(
"litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client",
return_value=mock_client,
),
pytest.raises(ValueError, match="failed with status 400"),
):
await handler.exchange_token("bad-jwt", server)
@pytest.mark.asyncio
async def test_exchange_token_http_error_does_not_log_response_body():
"""Raw IDP error bodies are not logged because they can contain credentials."""
handler = TokenExchangeHandler()
server = _obo_server()
raw_response_body = "client_secret=do-not-log"
mock_response = MagicMock()
mock_response.status_code = 401
mock_response.text = raw_response_body
mock_response.raise_for_status.side_effect = httpx.HTTPStatusError(
"Unauthorized",
request=MagicMock(),
response=mock_response,
)
mock_client = AsyncMock()
mock_client.post.return_value = mock_response
with (
patch(
"litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client",
return_value=mock_client,
),
patch(
"litellm.proxy._experimental.mcp_server.auth.token_exchange.verbose_logger.debug"
) as mock_debug,
pytest.raises(ValueError, match="failed with status 401"),
):
await handler.exchange_token("bad-jwt", server)
logged_values = " ".join(
str(value)
for call in mock_debug.call_args_list
for value in [*call.args, *call.kwargs.values()]
)
assert raw_response_body not in logged_values
@pytest.mark.asyncio
async def test_exchange_token_missing_access_token():
"""Response without access_token raises ValueError."""
handler = TokenExchangeHandler()
server = _obo_server()
resp = MagicMock()
resp.json.return_value = {"token_type": "Bearer"}
resp.raise_for_status = MagicMock()
mock_client = AsyncMock()
mock_client.post.return_value = resp
with (
patch(
"litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client",
return_value=mock_client,
),
pytest.raises(ValueError, match="missing 'access_token'"),
):
await handler.exchange_token("jwt", server)
@pytest.mark.asyncio
async def test_exchange_token_missing_endpoint():
"""Missing token_exchange_endpoint and token_url raises ValueError."""
handler = TokenExchangeHandler()
server = _obo_server(token_exchange_endpoint=None, token_url=None)
with pytest.raises(ValueError, match="no token_exchange_endpoint or token_url"):
await handler.exchange_token("jwt", server)
@pytest.mark.asyncio
async def test_exchange_token_missing_credentials():
"""Missing client_id or client_secret raises ValueError."""
handler = TokenExchangeHandler()
server = _obo_server(client_id=None, client_secret=None)
# has_token_exchange_config will be False, so we call _do_exchange directly
with pytest.raises(ValueError, match="missing client_id or client_secret"):
await handler._do_exchange("jwt", server)
# ── resolve_mcp_auth Integration ──
@pytest.mark.asyncio
async def test_resolve_mcp_auth_with_token_exchange():
"""resolve_mcp_auth delegates to token exchange when server has OBO config and subject_token provided."""
server = _obo_server()
mock_handler = AsyncMock()
mock_handler.exchange_token.return_value = "obo-scoped-token"
with patch(
"litellm.proxy._experimental.mcp_server.auth.token_exchange.mcp_token_exchange_handler",
mock_handler,
):
result = await resolve_mcp_auth(server, subject_token="user-jwt")
assert result == "obo-scoped-token"
mock_handler.exchange_token.assert_called_once_with("user-jwt", server)
@pytest.mark.asyncio
async def test_resolve_mcp_auth_obo_without_subject_token_falls_through():
"""Without a subject_token, resolve_mcp_auth falls through to client_credentials."""
server = _obo_server(
token_url="https://auth.example.com/token",
)
mock_client = AsyncMock()
mock_client.post.return_value = _exchange_response("cc-token")
with patch(
"litellm.proxy._experimental.mcp_server.oauth2_token_cache.get_async_httpx_client",
return_value=mock_client,
):
result = await resolve_mcp_auth(server, subject_token=None)
# Falls through to client_credentials since subject_token is None
# The server has client_id/client_secret/token_url so has_client_credentials is True
assert result == "cc-token"
@pytest.mark.asyncio
async def test_resolve_mcp_auth_obo_without_subject_token_uses_cached_client_credentials():
"""The M2M fallback for OBO servers reuses the client_credentials cache."""
server = _obo_server(
server_id="srv-obo-m2m-cache",
token_url="https://auth.example.com/token",
)
mock_client = AsyncMock()
mock_client.post.return_value = _exchange_response("cached-cc-token")
with patch(
"litellm.proxy._experimental.mcp_server.oauth2_token_cache.get_async_httpx_client",
return_value=mock_client,
):
first = await resolve_mcp_auth(server, subject_token=None)
second = await resolve_mcp_auth(server, subject_token=None)
assert first == second == "cached-cc-token"
mock_client.post.assert_called_once()
@pytest.mark.asyncio
async def test_resolve_mcp_auth_header_beats_obo():
"""An explicit mcp_auth_header takes priority over OBO token exchange."""
server = _obo_server()
result = await resolve_mcp_auth(
server, mcp_auth_header="Bearer override", subject_token="user-jwt"
)
assert result == "Bearer override"
# ── Bearer Token Extraction ──
def test_extract_bearer_token_from_oauth2_headers():
"""Extracts token from oauth2_headers Authorization header."""
result = MCPServerManager._extract_bearer_token(
oauth2_headers={"Authorization": "Bearer my-jwt-token"},
raw_headers=None,
)
assert result == "my-jwt-token"
def test_extract_bearer_token_from_raw_headers():
"""Falls back to raw_headers when oauth2_headers missing."""
result = MCPServerManager._extract_bearer_token(
oauth2_headers=None,
raw_headers={"authorization": "Bearer raw-jwt"},
)
assert result == "raw-jwt"
def test_extract_bearer_token_no_bearer_prefix():
"""Returns token as-is when no Bearer prefix."""
result = MCPServerManager._extract_bearer_token(
oauth2_headers={"Authorization": "some-opaque-token"},
raw_headers=None,
)
assert result == "some-opaque-token"
def test_extract_bearer_token_none():
"""Returns None when no auth headers present."""
result = MCPServerManager._extract_bearer_token(
oauth2_headers=None,
raw_headers=None,
)
assert result is None
# ── MCPServer Properties ──
def test_has_token_exchange_config_true():
"""has_token_exchange_config is True for a fully configured OBO server."""
server = _obo_server()
assert server.has_token_exchange_config is True
def test_has_token_exchange_config_false_wrong_auth_type():
"""has_token_exchange_config is False when auth_type is not oauth2_token_exchange."""
server = _obo_server(auth_type=MCPAuth.oauth2)
assert server.has_token_exchange_config is False
def test_has_token_exchange_config_false_missing_creds():
"""has_token_exchange_config is False when client_id/client_secret missing."""
server = _obo_server(client_id=None)
assert server.has_token_exchange_config is False
def test_has_token_exchange_config_uses_token_url_fallback():
"""has_token_exchange_config is True when token_url is set instead of token_exchange_endpoint."""
server = _obo_server(
token_exchange_endpoint=None,
token_url="https://idp.example.com/token",
)
assert server.has_token_exchange_config is True
# ── Config Loading ──
@pytest.mark.asyncio
async def test_config_loading_token_exchange_fields():
"""load_servers_from_config correctly maps OBO config fields to MCPServer."""
manager = MCPServerManager()
config = {
"my_obo_server": {
"url": "https://mcp.example.com/mcp",
"transport": "http",
"auth_type": "oauth2_token_exchange",
"client_id": "my-client",
"client_secret": "my-secret",
"token_exchange_endpoint": "https://idp.example.com/oauth2/token",
"audience": "api://my-mcp",
"scopes": ["read", "write"],
"subject_token_type": "urn:ietf:params:oauth:token-type:jwt",
}
}
await manager.load_servers_from_config(config)
servers = list(manager.config_mcp_servers.values())
assert len(servers) == 1
server = servers[0]
assert server.auth_type == MCPAuth.oauth2_token_exchange
assert server.token_exchange_endpoint == "https://idp.example.com/oauth2/token"
assert server.audience == "api://my-mcp"
assert server.subject_token_type == "urn:ietf:params:oauth:token-type:jwt"
assert server.client_id == "my-client"
assert server.client_secret == "my-secret"
assert server.scopes == ["read", "write"]
assert server.has_token_exchange_config is True
@pytest.mark.asyncio
async def test_config_loading_default_subject_token_type():
"""subject_token_type defaults to access_token when not specified in config."""
manager = MCPServerManager()
config = {
"obo_defaults": {
"url": "https://mcp.example.com/mcp",
"transport": "http",
"auth_type": "oauth2_token_exchange",
"client_id": "cid",
"client_secret": "csec",
"token_exchange_endpoint": "https://idp.example.com/token",
}
}
await manager.load_servers_from_config(config)
server = list(manager.config_mcp_servers.values())[0]
assert server.subject_token_type == "urn:ietf:params:oauth:token-type:access_token"
@pytest.mark.asyncio
async def test_database_loading_token_exchange_scopes_from_credentials():
"""DB-loaded OBO server credentials retain configured scopes."""
manager = MCPServerManager()
db_server = LiteLLM_MCPServerTable(
server_id="srv-obo-db",
server_name="obo_db_server",
url="https://mcp.example.com/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2_token_exchange,
credentials={
"client_id": "db-client",
"client_secret": "db-secret",
"token_exchange_endpoint": "https://idp.example.com/oauth2/token",
"audience": "api://db-mcp",
"scopes": ["db.read", "db.write"],
},
)
server = await manager.build_mcp_server_from_table(
db_server,
credentials_are_encrypted=False,
)
assert server.auth_type == MCPAuth.oauth2_token_exchange
assert server.client_id == "db-client"
assert server.client_secret == "db-secret"
assert server.token_exchange_endpoint == "https://idp.example.com/oauth2/token"
assert server.audience == "api://db-mcp"
assert server.scopes == ["db.read", "db.write"]
@pytest.mark.asyncio
async def test_exchange_token_uses_client_secret_basic_when_configured():
"""LIT-4091: token exchange with token_endpoint_auth_method=client_secret_basic sends the
client credentials as HTTP Basic and omits client_secret from the body."""
import base64
handler = TokenExchangeHandler()
server = _obo_server(
server_id="srv-obo-basic", token_endpoint_auth_method="client_secret_basic"
)
mock_client = AsyncMock()
mock_client.post.return_value = _exchange_response("scoped-basic")
with patch(
"litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client",
return_value=mock_client,
):
result = await handler.exchange_token("user-jwt-basic", server)
assert result == "scoped-basic"
_, kwargs = mock_client.post.call_args
expected = "Basic " + base64.b64encode(b"litellm-client-id:litellm-client-secret").decode()
assert kwargs["headers"]["Authorization"] == expected
assert "client_secret" not in kwargs["data"]
assert "client_id" not in kwargs["data"]
assert kwargs["data"]["grant_type"] == TOKEN_EXCHANGE_GRANT_TYPE

View file

@ -2333,6 +2333,59 @@ class TestMCPServerManager:
assert emitted.headers["Authorization"] == "Bearer upstream-token"
assert not kwargs["extra_headers"] or "authorization" not in {k.lower() for k in kwargs["extra_headers"]}
@pytest.mark.asyncio
async def test_create_mcp_client_token_exchange_never_falls_back_to_v1(self):
"""A configured OBO server is owned end to end by the v2 token_exchange arm, even when the
caller supplies an x-mcp-* override. This is what makes the v1 OBO handler unreachable, so if
it ever defers to v1 again the deleted handler is silently needed back."""
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
OAuthToken,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.resolver import (
UpstreamCredentialProvider,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Ok
class _StubExchanger:
def __init__(self):
self.subject_tokens = []
async def exchange(self, subject_token, server, config, *, tenant_id=""):
self.subject_tokens.append(subject_token)
return Ok(OAuthToken(access_token="exchanged-token"))
async def invalidate(self, subject_token, server, config, *, tenant_id=""):
return None
exchanger = _StubExchanger()
manager = MCPServerManager()
server = MCPServer(
server_id="obo-egress",
name="obo",
url="https://example.com",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2_token_exchange,
client_id="gateway-client",
client_secret="gateway-secret",
token_exchange_endpoint="https://idp.example.com/oauth2/token",
)
with (
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.resolve_mcp_auth",
new_callable=AsyncMock,
) as mock_resolve,
patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient") as mock_client_cls,
):
await manager._create_mcp_client(
server=server,
mcp_auth_header="Bearer caller-override",
subject_token="eyJ-subject-token",
cred_provider=UpstreamCredentialProvider(token_exchanger=exchanger),
)
mock_resolve.assert_not_awaited()
assert exchanger.subject_tokens == ["eyJ-subject-token"]
assert self._emitted_authorization(mock_client_cls) == "Bearer exchanged-token"
@staticmethod
def _emitted_authorization(mock_client_cls) -> str:
kwargs = mock_client_cls.call_args.kwargs

View file

@ -237,7 +237,7 @@ class TestListToolRestApiWithToolSearch:
return_value={},
),
patch(
"litellm.proxy._experimental.mcp_server.rest_endpoints._get_oauth2_server_ids",
"litellm.proxy._experimental.mcp_server.rest_endpoints._v1_resolved_oauth2_server_ids",
return_value=[],
),
patch(
@ -316,7 +316,7 @@ class TestListToolRestApiWithToolSearch:
return_value={},
),
patch(
"litellm.proxy._experimental.mcp_server.rest_endpoints._get_oauth2_server_ids",
"litellm.proxy._experimental.mcp_server.rest_endpoints._v1_resolved_oauth2_server_ids",
return_value=[],
),
patch(

View file

@ -2783,3 +2783,68 @@ class TestRestListToolsetFiltering:
)
assert [tool.name for tool in result] == ["lookup_status"]
class TestV1ResolvedOauth2Gate:
"""The REST surface must stop resolving per-user OAuth2 tokens for servers the v2 resolver owns.
``_resolve_v2_auth`` drops any Authorization built here for an ``authorization_code`` server and
injects the resolver's own token, so the v1 lookup was a DB round-trip whose result was discarded.
A server that still defers to v1 (upstream-delegated oauth2) must keep resolving, which is what
makes these assertions non-vacuous.
"""
@staticmethod
def _oauth2_server(*, delegate_auth_to_upstream: bool) -> Any:
from litellm.proxy._experimental.mcp_server.server import MCPServer
from litellm.types.mcp import MCPTransport
return MCPServer(
server_id="oauth2-srv",
name="oauth2-srv",
url="https://upstream.example.com/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
delegate_auth_to_upstream=delegate_auth_to_upstream,
)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"delegate_auth_to_upstream, expected_headers, expected_lookups",
[
(False, None, 0),
(True, {"Authorization": "Bearer stored-token"}, 1),
],
)
async def test_user_oauth_headers_skip_v2_owned_servers(
self, delegate_auth_to_upstream, expected_headers, expected_lookups, monkeypatch
):
from litellm.proxy._experimental.mcp_server import db as mcp_db
server = self._oauth2_server(delegate_auth_to_upstream=delegate_auth_to_upstream)
resolve_token = AsyncMock(return_value={"access_token": "stored-token"})
monkeypatch.setattr(mcp_db, "resolve_valid_user_oauth_token", resolve_token)
headers = await rest_endpoints._get_user_oauth_extra_headers(
server,
UserAPIKeyAuth(user_id="alice", api_key="sk-1234"),
prefetched_creds={"oauth2-srv": {"access_token": "stored-token"}},
)
assert headers == expected_headers
assert resolve_token.await_count == expected_lookups
def test_prefetch_preflight_only_counts_v1_resolved_servers(self, monkeypatch):
v2_owned = self._oauth2_server(delegate_auth_to_upstream=False)
v1_resolved = self._oauth2_server(delegate_auth_to_upstream=True)
v1_resolved.server_id = "delegate-srv"
registry = {"oauth2-srv": v2_owned, "delegate-srv": v1_resolved}
monkeypatch.setattr(
rest_endpoints.global_mcp_server_manager,
"get_mcp_server_by_id",
lambda server_id: registry.get(server_id),
)
assert rest_endpoints._v1_resolved_oauth2_server_ids(["oauth2-srv"]) == set()
assert rest_endpoints._v1_resolved_oauth2_server_ids(["oauth2-srv", "delegate-srv"]) == {"delegate-srv"}