This commit is contained in:
devin-ai-integration[bot] 2026-10-03 18:02:49 +00:00 • committed by GitHub
commit ca3db98365
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
26 changed files with 4615 additions and 545 deletions

View file

@ -0,0 +1,196 @@
"""Caller-side sign-in requirements for MCP connects.
A guardrail that evaluates tool calls in the caller's own identity (an On-Behalf-Of exchange of the caller's
bearer) needs the caller signed in with its issuer before the first tool call, and a tool call's JSON-RPC
error cannot carry ``WWW-Authenticate``. ``token_exchange`` (OBO) servers have the same need: the caller must
present a subject token the gateway can exchange. Both cases share one contract: a connect that carries no
usable subject answers 401 with the RFC 9728 challenge, and the protected-resource metadata advertises the
issuers and scopes the caller signs in for.
Guardrails implement :class:`CallerSignInProvider`; :func:`caller_sign_in_for` merges the OBO server's own
requirement with every registered provider's so the challenge and the metadata always agree. The MCP package
never imports a concrete provider.
"""
from __future__ import annotations
import itertools
from collections.abc import Awaitable, Callable, Mapping
from dataclasses import dataclass
from typing import TYPE_CHECKING, Final, Protocol, runtime_checkable
from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError
from typing_extensions import assert_never
import litellm
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.types.mcp import MCPAuth
if TYPE_CHECKING:
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.mcp_server.mcp_server_manager import MCPServer
@dataclass(frozen=True, slots=True)
class CallerSignIn:
"""The issuers a caller signs in with and the scopes it requests before calling a server."""
issuers: tuple[str, ...]
scopes: tuple[str, ...]
@dataclass(frozen=True, slots=True)
class SignedIn:
"""The subject token the caller presented satisfies this provider's sign-in."""
@dataclass(frozen=True, slots=True)
class Rejected:
"""The caller's identity provider rejected the presented subject token."""
detail: str
claims: str | None = None
@dataclass(frozen=True, slots=True)
class Unavailable:
"""The provider could not reach a verdict; ``fail_open`` is the provider's own fallback policy."""
detail: str
fail_open: bool
CallerSignInPreflight = SignedIn | Rejected | Unavailable
@runtime_checkable
class CallerSignInProvider(Protocol):
def caller_sign_in(self, server: MCPServer, user_api_key_auth: UserAPIKeyAuth | None) -> CallerSignIn | None:
"""The sign-in this provider requires of callers hitting ``server``; ``None`` when it does not gate
the server for this caller (``user_api_key_auth=None`` is the anonymous metadata fetch that follows
a challenge)."""
...
async def preflight_caller_sign_in(
self, server: MCPServer, user_api_key_auth: UserAPIKeyAuth | None, subject_token: str
) -> CallerSignInPreflight:
"""Validate ``subject_token`` against this provider at connect time, where a challenge's
``WWW-Authenticate`` still reaches the client."""
...
class _JwtIssuerEntry(BaseModel):
model_config = ConfigDict(extra="ignore")
issuer: str | None = None
class _JwtAuthConfig(BaseModel):
model_config = ConfigDict(extra="ignore")
issuers: list[_JwtIssuerEntry] = [] # mutable-ok: pydantic copies the default per instance
_JWT_AUTH_ADAPTER: Final = TypeAdapter(_JwtAuthConfig)
def _jwt_auth_issuer_entries(jwtauth: object) -> tuple[_JwtIssuerEntry, ...]:
try:
return tuple(_JWT_AUTH_ADAPTER.validate_python(jwtauth, from_attributes=True).issuers)
except ValidationError:
return ()
def _providers() -> tuple[CallerSignInProvider, ...]:
return tuple(
callback
for callback in litellm.logging_callback_manager.get_custom_loggers_for_type(CustomGuardrail)
if isinstance(callback, CallerSignInProvider)
)
def jwt_auth_issuers() -> tuple[str, ...]:
"""The OAuth issuer identifier(s) LiteLLM's JWT auth trusts, for the OBO PRM authorization_servers.
In token_exchange the IdP that issues the subject JWT is the same one LiteLLM validates it
against, so OBO discovery points clients at the JWT-auth issuer to obtain a subject token.
Sourced from ``JWT_ISSUER`` and any configured ``litellm_jwtauth.issuers``.
"""
import os # noqa: PLC0415
from litellm.proxy.proxy_server import ( # noqa: PLC0415 # lazy: proxy_server pulls the whole proxy graph
general_settings, # pyright: ignore[reportUnknownVariableType] # proxy_server.general_settings is a raw untyped dict
)
env_issuer: Final = os.getenv("JWT_ISSUER")
env: Final[tuple[str, ...]] = (env_issuer,) if env_issuer else ()
settings: Final[Mapping[str, object]] = general_settings if isinstance(general_settings, Mapping) else {}
configured: Final = tuple(
entry.issuer for entry in _jwt_auth_issuer_entries(settings.get("litellm_jwtauth")) if entry.issuer
)
return tuple(dict.fromkeys((*env, *configured)))
def caller_sign_in_for(server: MCPServer, user_api_key_auth: UserAPIKeyAuth | None) -> CallerSignIn | None:
"""The merged sign-in requirement for ``server``: the OBO server's own issuer/scopes plus every
registered provider's contribution. ``None`` when nothing requires sign-in, which is also the gate the
connect-time challenge branches on."""
contributions: Final = tuple(
contribution
for contribution in (
*(
(CallerSignIn(issuers=jwt_auth_issuers(), scopes=tuple(server.scopes or ())),)
if server.auth_type == MCPAuth.oauth2_token_exchange
else ()
),
*(provider.caller_sign_in(server, user_api_key_auth) for provider in _providers()),
)
if contribution is not None
)
if not contributions:
return None
issuers: Final = tuple(dict.fromkeys(itertools.chain.from_iterable(c.issuers for c in contributions)))
scopes: Final = tuple(dict.fromkeys(itertools.chain.from_iterable(c.scopes for c in contributions)))
return CallerSignIn(issuers=issuers, scopes=scopes)
async def preflight_caller_sign_in(
server: MCPServer,
user_api_key_auth: UserAPIKeyAuth | None,
subject_token: str,
*,
root_path: str,
resource_metadata: str | None,
connecting: Callable[[], Awaitable[bool]],
) -> None:
"""Run every provider's connect-time check against the subject token while ``connecting``, so a bearer
the IdP will reject surfaces as a challenge here rather than a JSON-RPC error at the first tool call, and
a fail-closed provider outage is the connect's 503. On an open session nothing is exchanged here: the
tool-call hook runs the one exchange and answers inside the JSON-RPC envelope, with its guardrail Logs
row."""
from fastapi import HTTPException # noqa: PLC0415 # lazy: fastapi import stays off the cold path
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import ( # noqa: PLC0415 # lazy: adapter pulls MCP subgraph
raise_token_exchange_challenge,
)
gating: Final = tuple(
provider for provider in _providers() if provider.caller_sign_in(server, user_api_key_auth) is not None
)
if not gating or not await connecting():
return
for provider in gating:
match await provider.preflight_caller_sign_in(server, user_api_key_auth, subject_token):
case SignedIn():
continue
case Rejected(detail=_, claims=claims):
raise_token_exchange_challenge(
server, root_path=root_path, claims=claims, resource_metadata=resource_metadata
)
case Unavailable(fail_open=True):
continue
case Unavailable(detail=detail, fail_open=False):
raise HTTPException(status_code=503, detail=detail)
case _ as verdict:
assert_never(verdict)

View file

@ -37,6 +37,7 @@ from litellm.proxy._experimental.mcp_server.bridge_token_flow import (
can_store_oauth_credential,
oauth_authorization_uses_gateway_credential,
)
from litellm.proxy._experimental.mcp_server.caller_sign_in import caller_sign_in_for
from litellm.proxy._experimental.mcp_server.faults import (
CallerRejected,
CredentialSource,
@ -527,10 +528,7 @@ def _resolve_mcp_server_by_name_or_id(lookup: str, client_ip: str | None) -> MCP
global_mcp_server_manager,
)
by_name: Final = global_mcp_server_manager.get_mcp_server_by_name(lookup, client_ip=client_ip)
if by_name is not None:
return by_name
return global_mcp_server_manager.get_mcp_server_by_id(lookup, client_ip=client_ip)
return global_mcp_server_manager.get_mcp_server_answering_to(lookup, client_ip=client_ip)
def _resolve_oauth2_server_for_root_endpoints(
@ -2580,9 +2578,9 @@ async def _build_oauth_protected_resource_response(
detail=(f"Upstream oauth-protected-resource metadata unavailable for MCP server {mcp_server.name!r}"),
)
obo_response: Final = _obo_protected_resource_response(mcp_server, resource_url)
if obo_response is not None:
return obo_response
sign_in_response: Final = _caller_sign_in_protected_resource_response(mcp_server, resource_url)
if sign_in_response is not None:
return sign_in_response
if mcp_server is not None and mcp_server.advertises_gateway_authorization_server:
return {
@ -2603,51 +2601,30 @@ async def _build_oauth_protected_resource_response(
}
def _obo_protected_resource_response(mcp_server: MCPServer | None, resource_url: str) -> dict | None:
"""The OBO (token_exchange) PRM, or None when this server is not OBO / no issuer is configured.
def _caller_sign_in_protected_resource_response(
mcp_server: MCPServer | None, resource_url: str
) -> dict[str, object] | None:
"""The caller sign-in PRM: the OBO issuer(s) LiteLLM trusts merged with every registered
``CallerSignInProvider``'s contribution, or None when no sign-in gates this server.
The client SSOs with the IdP to obtain a subject token, which LiteLLM then exchanges, so discovery
points at the JWT-auth issuer(s) LiteLLM trusts (the same IdP that issues and validates the
subject), not the gateway. None falls the caller back to the gateway default so discovery still
returns metadata; it just can't name the IdP.
The client SSOs with the IdP to obtain a subject token, which LiteLLM then exchanges (or a
guardrail consumes directly), so discovery points at the issuer(s), not the gateway. None falls
the caller back to the gateway default so discovery still returns metadata; it just can't name
the IdP. The anonymous metadata fetch passes ``user_api_key_auth=None`` because it cannot see
which key selected a provider.
"""
if mcp_server is None or mcp_server.auth_type != MCPAuth.oauth2_token_exchange:
if mcp_server is None:
return None
issuers: Final = _jwt_auth_issuers()
if not issuers:
sign_in: Final = caller_sign_in_for(mcp_server, None)
if sign_in is None or not sign_in.issuers:
return None
return {
"authorization_servers": issuers,
"authorization_servers": sign_in.issuers,
"resource": resource_url,
"scopes_supported": (mcp_server.scopes if mcp_server.scopes else []),
"scopes_supported": sign_in.scopes,
}
def _jwt_auth_issuers() -> list:
"""The OAuth issuer identifier(s) LiteLLM's JWT auth trusts, for the OBO PRM authorization_servers.
In token_exchange the IdP that issues the subject JWT is the same one LiteLLM validates it
against, so OBO discovery points clients at the JWT-auth issuer to obtain a subject token.
Sourced from ``JWT_ISSUER`` and any configured ``litellm_jwtauth.issuers``.
"""
import os # noqa: PLC0415
from litellm.proxy.proxy_server import general_settings # noqa: PLC0415
issuers: Final[list] = []
env_issuer: Final = os.getenv("JWT_ISSUER")
if env_issuer:
issuers.append(env_issuer)
jwtauth: Final = general_settings.get("litellm_jwtauth") if isinstance(general_settings, Mapping) else None
raw_issuers: Final = jwtauth.get("issuers") if isinstance(jwtauth, dict) else getattr(jwtauth, "issuers", None)
for cfg in raw_issuers or []:
issuer = cfg.get("issuer") if isinstance(cfg, dict) else getattr(cfg, "issuer", None)
if issuer and issuer not in issuers:
issuers.append(issuer)
return issuers
@router.get("/.well-known/oauth-protected-resource")
def oauth_protected_resource_root(request: Request) -> dict[str, str | tuple[str, ...]]:
request_base_url: Final = get_request_base_url(request)

View file

@ -13,6 +13,7 @@ import json
import math
import os
import re
import secrets
import time
from collections.abc import (
AsyncIterator,
@ -169,6 +170,7 @@ from litellm.proxy._experimental.mcp_server.utils import (
normalize_server_name,
openapi_tool_name,
parse_admin_env_vars,
server_answers_to_name,
strip_known_server_prefix,
validate_mcp_server_name,
)
@ -1163,6 +1165,10 @@ def _raw_header_value(raw_headers: Mapping[str, str] | None, name: str) -> str |
return next((v for k, v in (raw_headers or {}).items() if isinstance(k, str) and k.lower() == name), None)
def _is_master_key(bearer: str, master_key: str | None) -> bool:
return bool(master_key) and secrets.compare_digest(bearer.encode(), (master_key or "").encode())
def _has_explicit_litellm_admission_header(raw_headers: Mapping[str, str] | None) -> bool:
"""Admission only consumes a non-empty ``x-litellm-api-key``; an empty one falls back to ``Authorization``."""
return bool(_raw_header_value(raw_headers, "x-litellm-api-key"))
@ -3897,6 +3903,31 @@ class MCPServerManager:
return None
return bearer
@staticmethod
def _caller_sign_in_subject_token(
oauth2_headers: Mapping[str, str] | None,
raw_headers: Mapping[str, str] | None,
) -> str | None:
"""The ``Bearer`` credential a caller sign-in provider validates. An admission that consumed
``Authorization`` (custom auth, built-in OAuth2, JWT) did so on the caller's own IdP token, so that token
is the subject; any other scheme, a LiteLLM key (virtual or master) and a bearer repeating
``x-litellm-api-key`` are withheld."""
from litellm.proxy.proxy_server import master_key # noqa: PLC0415 # circular import
authorization: Final = (oauth2_headers or {}).get("Authorization") or _raw_header_value(
raw_headers, "authorization"
)
scheme_and_credential: Final = (authorization or "").split(None, 1)
if len(scheme_and_credential) != 2 or scheme_and_credential[0].lower() != "bearer":
return None
bearer: Final = scheme_and_credential[1]
if bearer.startswith(LITELLM_VIRTUAL_KEY_PREFIX) or _is_master_key(bearer, master_key):
return None
admission_header: Final = _raw_header_value(raw_headers, "x-litellm-api-key")
if admission_header and strip_auth_scheme(admission_header, "Bearer") == bearer:
return None
return bearer
def _obo_subject_token(
self,
server: MCPServer,
@ -4109,6 +4140,7 @@ class MCPServerManager:
oauth2_headers: dict[str, str] | None,
user_api_key_auth: UserAPIKeyAuth | None,
raw_headers: Mapping[str, str] | None = None,
resource_metadata: str | None = None,
) -> None:
"""Mint an exchange-backed server's upstream credential at the transport edge.
@ -4140,13 +4172,24 @@ class MCPServerManager:
if subject_token is not None:
return
case _:
from litellm.proxy._experimental.mcp_server.caller_sign_in import ( # noqa: PLC0415 # lazy: provider discovery pulls the guardrail registry
caller_sign_in_for,
)
sign_in_subject: Final = self._caller_sign_in_subject_token(oauth2_headers, raw_headers)
if sign_in_subject is None and caller_sign_in_for(server, user_api_key_auth) is not None:
raise_token_exchange_challenge(
server, root_path=get_request_root_path(), resource_metadata=resource_metadata
)
return
resolved_server: Final = await self.ensure_oauth_metadata_discovered(server)
spec: Final = to_server_spec_fail_closed(resolved_server)
if spec is None or not isinstance(spec.config, (TokenExchangeConfig, IdJagConfig)):
return
if subject_token is None and isinstance(spec.config, TokenExchangeConfig):
raise_token_exchange_challenge(resolved_server, root_path=get_request_root_path())
raise_token_exchange_challenge(
resolved_server, root_path=get_request_root_path(), resource_metadata=resource_metadata
)
match await self._cred_provider.resolve_credentials(to_subject(user_api_key_auth, subject_token), spec):
case Ok(_):
return
@ -4156,6 +4199,7 @@ class MCPServerManager:
resolved_server,
root_path=get_request_root_path(),
claims=err.unauthorized.claims,
resource_metadata=resource_metadata,
)
raise_public(err)
@ -5747,13 +5791,11 @@ class MCPServerManager:
if proxy_logging_obj is None:
return hook_result
# Extract incoming Bearer token from raw request headers so
# guardrails like MCPJWTSigner can verify + re-sign it (FR-5).
normalized_raw: Final = {k.lower(): v for k, v in (raw_headers or {}).items()}
incoming_bearer_token: str | None = None
auth_hdr: Final = normalized_raw.get("authorization", "")
if auth_hdr.lower().startswith("bearer "):
incoming_bearer_token = auth_hdr[len("bearer ") :]
inbound_authorization: Final = _raw_header_value(raw_headers, "authorization") or ""
incoming_bearer_token: Final = (
inbound_authorization[len("bearer ") :] if inbound_authorization.lower().startswith("bearer ") else None
)
incoming_subject_token: Final = self._caller_sign_in_subject_token(None, raw_headers)
pre_hook_kwargs: Final = {
"guardrail_context": guardrail_context,
@ -5769,6 +5811,7 @@ class MCPServerManager:
),
"user_api_key_hash": (getattr(user_api_key_auth, "api_key_hash", None) if user_api_key_auth else None),
"incoming_bearer_token": incoming_bearer_token,
"incoming_subject_token": incoming_subject_token,
"headers": logging_safe_mcp_headers(raw_headers),
"tool_description": tool.description if tool is not None else None,
"tool_input_schema": tool.input_schema if tool is not None else None,
@ -6981,6 +7024,70 @@ class MCPServerManager:
return server
return None
def get_mcp_server_answering_to(
self, name: str, client_ip: str | None = None, *, among: Sequence[MCPServer] | None = None
) -> MCPServer | None:
"""The one server a ``/mcp/{name}`` segment denotes, shared by the connect preflight, the scoped
router, and RFC 9728 discovery so all three name the same server: the exact ``get_mcp_server_by_name``
priority first, then the exact ``server_id``, then the name priority case-insensitively, then any prefix
form routing accepts. A name that denotes a server hidden from ``client_ip`` resolves to ``None`` at the
pass that found it: it never falls through to a looser pass that could name another server. ``among``
runs the same passes over those servers alone instead of the registry, which is how the scoped router
picks the caller's granted server answering to ``name``."""
if among is not None:
return self._server_among_answering_to(name, tuple(among), client_ip)
exact: Final = self.get_mcp_server_by_name(name)
if exact is not None:
return exact if self._is_server_accessible_from_ip(exact, client_ip) else None
by_id: Final = self.get_mcp_server_by_id(name)
if by_id is not None:
return by_id if self._is_server_accessible_from_ip(by_id, client_ip) else None
requested: Final = name.lower()
servers: Final = tuple(self.get_registry().values())
identifiers: Final[tuple[Callable[[MCPServer], str | None], ...]] = (
lambda server: server.alias,
lambda server: server.server_name,
lambda server: server.name,
)
for identifier in identifiers:
if (found := next((s for s in servers if (identifier(s) or "").lower() == requested), None)) is not None:
return found if self._is_server_accessible_from_ip(found, client_ip) else None
return next(
(
server
for server in self.get_filtered_registry(client_ip).values()
if server_answers_to_name(server, name)
),
None,
)
def _server_among_answering_to(
self, name: str, servers: Sequence[MCPServer], client_ip: str | None
) -> MCPServer | None:
"""``get_mcp_server_answering_to`` over ``servers`` instead of the registry: the same passes in the
same order, with a server hidden from ``client_ip`` resolving to ``None`` at the pass that found it."""
requested: Final = name.lower()
passes: Final[tuple[Callable[[MCPServer], bool], ...]] = (
lambda server: server.alias == name,
lambda server: server.server_name == name,
lambda server: server.name == name,
lambda server: server.server_id == name,
lambda server: (server.alias or "").lower() == requested,
lambda server: (server.server_name or "").lower() == requested,
lambda server: (server.name or "").lower() == requested,
)
for matches in passes:
if (found := next((server for server in servers if matches(server)), None)) is not None:
return found if self._is_server_accessible_from_ip(found, client_ip) else None
return next(
(
server
for server in servers
if self._is_server_accessible_from_ip(server, client_ip) and server_answers_to_name(server, name)
),
None,
)
def get_filtered_registry(self, client_ip: str | None = None) -> dict[str, MCPServer]:
"""
Get registry filtered by client IP access control.

View file

@ -120,6 +120,7 @@ from litellm.proxy._experimental.mcp_server.utils import (
logging_safe_mcp_headers,
match_known_tool_name,
normalize_server_name,
server_answers_to_name,
split_server_prefix_from_name,
strip_known_server_prefix,
)
@ -435,6 +436,7 @@ async def _dispatch_virtual_mcp_tool(
async def _get_allowed_mcp_servers_from_mcp_server_names(
mcp_servers: Sequence[str] | None,
allowed_mcp_servers: list[MCPServer],
client_ip: str | None = None,
) -> list[MCPServer]:
"""
Get the filtered MCP servers from the MCP server names.
@ -451,15 +453,10 @@ async def _get_allowed_mcp_servers_from_mcp_server_names(
# Filter servers based on mcp_servers parameter if provided
if mcp_servers is not None:
for server_or_group in mcp_servers:
server_name_matched = False
for server in allowed_mcp_servers:
if server and _server_answers_to(server, server_or_group):
filtered_server[server.server_id] = server
server_name_matched = True
break
if not server_name_matched:
scoped = _scoped_server(server_or_group, allowed_mcp_servers, client_ip)
if scoped is not None:
filtered_server[scoped.server_id] = scoped
else:
try:
access_group_server_ids = await MCPRequestHandler._get_mcp_servers_from_access_groups(
[server_or_group]
@ -493,8 +490,20 @@ def _http_detail_message(detail: object) -> str:
def _server_answers_to(server: MCPServer, name: str) -> bool:
requested: Final = name.lower()
return any(requested == known.lower() for known in iter_known_server_prefixes(server) if known)
return server_answers_to_name(server, name)
def _scoped_server(name: str, allowed_mcp_servers: Sequence[MCPServer], client_ip: str | None) -> MCPServer | None:
"""The granted server a scoped ``name`` selects for the caller: the registry's own pass order run over
``allowed_mcp_servers`` alone, so a granted server wins over an ungranted alias or case variant the registry
would pick. ``None`` when no granted server answers, or when the registry's own pick for ``name`` is a server
hidden from ``client_ip``; the caller then retries the name as an access group it holds."""
if (
global_mcp_server_manager.get_mcp_server_answering_to(name, client_ip=client_ip) is None
and global_mcp_server_manager.get_mcp_server_answering_to(name) is not None
):
return None
return global_mcp_server_manager.get_mcp_server_answering_to(name, client_ip=client_ip, among=allowed_mcp_servers)
async def raise_denied_scoped_mcp_access(
@ -675,6 +684,7 @@ async def _get_allowed_mcp_servers(
allowed_mcp_servers = await _get_allowed_mcp_servers_from_mcp_server_names(
mcp_servers=mcp_servers,
allowed_mcp_servers=allowed_mcp_servers,
client_ip=client_ip,
)
return allowed_mcp_servers
@ -1776,7 +1786,9 @@ def _challenge_missing_token_exchange_subject(
The listing that fills a cold catalog absorbs the upstream 401 by design, so without this
check a missing subject surfaces as an unknown-tool error instead of the challenge the
warm path already raises. Gated to servers the key may reach so an unauthorized caller
learns nothing about the catalog.
learns nothing about the catalog. Guardrail-only sign-in is challenged at connect instead: a
tool call's JSON-RPC error drops ``WWW-Authenticate``, so the guardrail's own rejection is
the more useful answer there.
"""
if server is None or server.auth_type != MCPAuth.oauth2_token_exchange:
return
@ -2439,6 +2451,7 @@ async def call_mcp_tool(
allowed_mcp_servers = await _get_allowed_mcp_servers_from_mcp_server_names(
mcp_servers=mcp_servers,
allowed_mcp_servers=allowed_mcp_servers,
client_ip=client_ip,
)
if mcp_servers and not allowed_mcp_servers:
await raise_denied_scoped_mcp_access(

View file

@ -348,6 +348,7 @@ def raise_token_exchange_challenge(
*,
root_path: str,
claims: str | None = None,
resource_metadata: str | None = None,
) -> NoReturn:
"""Raise the RFC 9728 / RFC 6750 challenge an OBO (``token_exchange``) server returns when the
caller's subject token is missing or the IdP rejected it.
@ -365,8 +366,12 @@ def raise_token_exchange_challenge(
``error="invalid_token"`` and is byte-identical to the static one. Both the error value (one of
two literals) and the base64 claims draw from a fixed alphabet, so nothing from the IdP body
reaches the header unescaped.
``resource_metadata`` is the absolute metadata URL of the route the client connected on (RFC 9728
5.1 names the parameter a URL, and the MCP SDK fetches it verbatim), supplied by the connect gate
that still holds the request; without it the challenge falls back to the alias's relative path.
"""
resource_metadata: Final = oauth_protected_resource_path(root_path, server)
metadata_url: Final = resource_metadata or oauth_protected_resource_path(root_path, server)
encoded_claims: Final = base64.b64encode(claims.encode()).decode() if claims else None
error: Final = "insufficient_claims" if encoded_claims else "invalid_token"
error_description: Final = (
@ -376,7 +381,7 @@ def raise_token_exchange_challenge(
)
www_authenticate: Final = ", ".join(
(
f'Bearer resource_metadata="{resource_metadata}"',
f'Bearer resource_metadata="{metadata_url}"',
f'error="{error}"',
f'error_description="{error_description}"',
*((f'claims="{encoded_claims}"',) if encoded_claims else ()),

View file

@ -9,6 +9,7 @@ call), so it needs no lazy wrapper.
from __future__ import annotations
from dataclasses import dataclass
from typing import Final
import httpx
@ -24,6 +25,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_sto
InMemoryTokenCacheBackend,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchanger import (
ExchangeHttpPost,
OboTokenExchanger,
SubjectTokenRejected,
TokenExchangeClientError,
@ -34,11 +36,27 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchanger
_GATEWAY_FAULT_OAUTH_ERRORS: Final = frozenset(
{"invalid_client", "unauthorized_client", "unsupported_grant_type", "invalid_target", "invalid_scope"}
)
_INVALID_ASSERTION_AADSTS_PREFIX: Final = "50027"
def _oauth_error_fields(response: httpx.Response) -> tuple[str | None, str | None]:
"""Read the RFC 6749 5.2 ``error`` code and the IdP's step-up ``claims`` blob from a
token-endpoint error body, as ``(error, claims)`` with None for whatever is absent.
@dataclass(frozen=True, slots=True)
class OAuthErrorBody:
error: str | None
claims: str | None
error_codes: tuple[str, ...]
@property
def gateway_fault(self) -> str | None:
if self.error is None or self.error not in _GATEWAY_FAULT_OAUTH_ERRORS:
return None
if any(code.startswith(_INVALID_ASSERTION_AADSTS_PREFIX) for code in self.error_codes):
return None
return self.error
def oauth_error_fields(response: httpx.Response) -> OAuthErrorBody:
"""Read the RFC 6749 5.2 ``error`` code, the IdP's step-up ``claims`` blob and Entra's
``error_codes`` sub-codes from a token-endpoint error body, None or empty for whatever is absent.
``claims`` is the Entra Conditional Access / CAE challenge (a JSON string the client must
replay to the IdP to satisfy the step-up); it is the caller's own requirement, not an IdP
@ -48,14 +66,18 @@ def _oauth_error_fields(response: httpx.Response) -> tuple[str | None, str | Non
try:
body: Final[object] = response.json()
except Exception: # noqa: BLE001
return None, None
return OAuthErrorBody(error=None, claims=None, error_codes=())
if not isinstance(body, dict):
return None, None
return OAuthErrorBody(error=None, claims=None, error_codes=())
code: Final = body.get("error")
claims: Final = body.get("claims")
return (
code if isinstance(code, str) else None,
claims if isinstance(claims, str) and claims else None,
raw_codes: Final = body.get("error_codes")
return OAuthErrorBody(
error=code if isinstance(code, str) else None,
claims=claims if isinstance(claims, str) and claims else None,
error_codes=tuple(str(c) for c in raw_codes if isinstance(c, (int, str)))
if isinstance(raw_codes, list)
else (),
)
@ -74,24 +96,32 @@ async def _post_exchange_endpoint(
headers: Final = {"Accept": "application/json", **client_auth_headers}
try:
client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.MCP) # pyright: ignore
response: Final = await client.post(url, headers=headers, data=form) # pyright: ignore
response: Final = await client.post( # pyright: ignore[reportUnknownMemberType] # untyped handler
url, headers=headers, data=form
)
response.raise_for_status() # pyright: ignore
parsed: Final[object] = response.json() # pyright: ignore
except httpx.HTTPStatusError as status_err:
status_code: Final = status_err.response.status_code
if status_code in (408, 429):
# Retry hints, not subject rejections: the IdP is shedding load, so a 401 would tell the
# caller to sign in again for nothing; surface it like a transport failure.
verbose_logger.warning("MCP token exchange throttled or timed out (HTTP %d)", status_code)
return None
if 400 <= status_code < 500:
oauth_error, claims = _oauth_error_fields(status_err.response)
if oauth_error in _GATEWAY_FAULT_OAUTH_ERRORS:
oauth_error: Final = oauth_error_fields(status_err.response)
gateway_fault: Final = oauth_error.gateway_fault
if gateway_fault is not None:
verbose_logger.warning(
"MCP token exchange rejected as %s (HTTP %d); check the gateway client credentials, "
"audience, and scope for this server",
oauth_error,
gateway_fault,
status_code,
)
raise TokenExchangeClientError(oauth_error) from status_err
raise TokenExchangeClientError(gateway_fault) from status_err
raise SubjectTokenRejected(
f"IdP rejected the subject token (HTTP {status_code})",
claims=claims,
claims=oauth_error.claims,
) from status_err
verbose_logger.warning("MCP token exchange request failed: %s", status_err)
return None
@ -106,9 +136,9 @@ async def _post_exchange_endpoint(
return parsed # pyright: ignore
def build_token_exchanger() -> OboTokenExchanger:
def build_token_exchanger(*, post: ExchangeHttpPost = _post_exchange_endpoint) -> OboTokenExchanger:
return OboTokenExchanger(
_post_exchange_endpoint,
post,
cache=InMemoryTokenCacheBackend(max_size=MCP_TOKEN_EXCHANGE_CACHE_MAX_SIZE),
default_ttl_seconds=MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL,
min_ttl_seconds=MCP_OAUTH2_TOKEN_CACHE_MIN_TTL,

View file

@ -8,12 +8,13 @@ import asyncio
import contextlib
import contextvars
import hashlib
import itertools
import json
import os
import time
import types
from collections import Counter
from collections.abc import AsyncGenerator, AsyncIterator, Callable, Iterable, Mapping, Sequence
from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Callable, Iterable, Iterator, Mapping, Sequence
from typing import TYPE_CHECKING, Final, NoReturn, Protocol
import httpx
@ -36,6 +37,7 @@ from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
MCPRequestHandler,
_is_mcp_admitted_user_subject,
)
from litellm.proxy._experimental.mcp_server.caller_sign_in import caller_sign_in_for
from litellm.proxy._experimental.mcp_server.client_allowlist import (
MCPClientAllowlist,
check_mcp_client_allowed,
@ -62,6 +64,7 @@ from litellm.proxy._experimental.mcp_server.mcp_debug import (
)
from litellm.proxy._experimental.mcp_server.oauth_utils import (
_redact_mcp_resource_url,
get_passthrough_resource_metadata_url,
get_passthrough_www_authenticate,
get_route_relative_request_path,
well_known_root_suffix,
@ -1424,6 +1427,41 @@ if MCP_AVAILABLE:
return consumed_messages, b"".join(body_chunks)
class _ConnectBodyPeek:
"""Reads a session-less ``POST`` body only once a gate asks whether it is ``initialize``, so a challenge
that needs no body still answers before the body arrives; consumed messages replay through ``receive``."""
def __init__(self, receive: Receive, peekable: bool) -> None:
self._receive: Final = receive
self._peekable: Final = peekable
self._body: bytes | None = None
self._replay: Iterator[Message] = iter(())
async def read(self) -> bytes:
messages, body = await _read_request_body_for_routing(self._receive)
self._replay = itertools.chain(self._replay, messages)
return body
async def body(self) -> bytes:
if not self._peekable:
return b""
if self._body is None:
self._body = await self.read()
return self._body
async def connecting(self) -> bool:
return _is_initialize_request(await self.body())
async def receive(self) -> Message:
replayed: Final = next(self._replay, None)
return replayed if replayed is not None else await self._receive()
def _known_connecting(value: bool) -> Callable[[], Awaitable[bool]]:
async def answer() -> bool:
return value
return answer
async def _handle_stale_mcp_session(
scope: Scope,
receive: Receive,
@ -1606,6 +1644,7 @@ if MCP_AVAILABLE:
mcp_server_auth_headers: dict[str, dict[str, str]] | None,
user_api_key_auth: UserAPIKeyAuth | None,
client_ip: str | None,
connecting: Callable[[], Awaitable[bool]],
allowed_server_ids: set[str] | None = None,
raw_headers: Mapping[str, str] | None = None,
) -> None:
@ -1621,7 +1660,28 @@ if MCP_AVAILABLE:
a server it will be 403'd on immediately after authentication.
"""
for server_name in mcp_servers or []:
server = operations.global_mcp_server_manager.get_mcp_server_by_name(server_name, client_ip=client_ip)
registry_pick = operations.global_mcp_server_manager.get_mcp_server_answering_to(
server_name, client_ip=client_ip
)
allowed_single = (
await operations._get_allowed_mcp_servers(
user_api_key_auth=user_api_key_auth, mcp_servers=mcp_servers, client_ip=client_ip
)
if registry_pick and mcp_servers is not None and len(mcp_servers) == 1
else ()
)
granted = (
operations.global_mcp_server_manager.get_mcp_server_answering_to(
server_name, client_ip=client_ip, among=allowed_single
)
if allowed_single
else None
)
server = granted if granted is not None else registry_pick
granted_single = granted is not None
obo_without_subject = (
server is not None and server.auth_type == MCPAuth.oauth2_token_exchange and not oauth2_headers
)
if server is not None and allowed_server_ids is not None and server.server_id not in allowed_server_ids:
# Caller's narrowed scope excludes this server — skip the
# preemptive challenge and let downstream authorization
@ -1717,12 +1777,21 @@ if MCP_AVAILABLE:
# reaches the token_exchange / pass-through blocks below.
continue
# token_exchange (OBO): the caller supplied no subject token. Challenge at connect
# (transport level, where WWW-Authenticate survives) with the RFC 9728 resource_metadata
# so the client discovers the IdP, SSOs, and retries with a subject token, which LiteLLM
# then exchanges. A tool-call-time 401 would be wrapped into a JSON-RPC error and the
# header lost, so the discovery flow needs this pre-emptive challenge.
if server and server.auth_type == MCPAuth.oauth2_token_exchange and not oauth2_headers:
# Caller sign-in: challenge at connect because a tool-call-time 401 is wrapped into a
# JSON-RPC error and the WWW-Authenticate header is lost. OBO keeps its connect gate;
# guardrail-only gates fire only on a single-server connect the key's grant admits, so a
# key without access gets the grant's 403 instead of a sign-in it could not use. The one
# admission lookup above serves the challenge, the sign-in preflight and the exchange.
sign_in = caller_sign_in_for(server, user_api_key_auth) if server is not None else None
resource_metadata = get_passthrough_resource_metadata_url(scope, server_name)
subject_token = (
operations.global_mcp_server_manager._caller_sign_in_subject_token( # pyright: ignore[reportPrivateUsage] # the manager owns the subject/admission filter shared with the preflight
oauth2_headers, raw_headers
)
if server is not None
else None
)
if server and sign_in is not None and subject_token is None and (obo_without_subject or granted_single):
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import ( # noqa: PLC0415 # lazy: adapter pulls MCP subgraph
raise_token_exchange_challenge,
)
@ -1730,7 +1799,25 @@ if MCP_AVAILABLE:
get_request_root_path,
)
raise_token_exchange_challenge(server, root_path=get_request_root_path())
raise_token_exchange_challenge(
server, root_path=get_request_root_path(), resource_metadata=resource_metadata
)
if server and sign_in is not None and subject_token is not None and granted_single:
from litellm.proxy._experimental.mcp_server.caller_sign_in import ( # noqa: PLC0415 # lazy: provider discovery pulls the guardrail registry
preflight_caller_sign_in,
)
from litellm.proxy.middleware.per_request_root_path_middleware import ( # noqa: PLC0415 # lazy: middleware imports proxy utils
get_request_root_path,
)
await preflight_caller_sign_in(
server,
user_api_key_auth,
subject_token,
root_path=get_request_root_path(),
resource_metadata=resource_metadata,
connecting=connecting,
)
# Exchange-backed modes (token_exchange's OBO mint, id_jag's stored-assertion mint): run
# the exchange here at the transport edge, so a rejected subject raises the RFC 9728
@ -1739,22 +1826,13 @@ if MCP_AVAILABLE:
# and what each mints from. Gated to single-server routes the key may reach; the
# multi-server aggregate keeps absorbing per-server auth failures so one bad server
# cannot 401 the whole connect.
if (
server
and len(mcp_servers or []) == 1
and server.server_id
in frozenset(
allowed.server_id
for allowed in await operations._get_allowed_mcp_servers(
user_api_key_auth=user_api_key_auth, mcp_servers=mcp_servers, client_ip=client_ip
)
)
):
if server and granted_single:
await operations.global_mcp_server_manager.preflight_token_exchange(
server=server,
oauth2_headers=oauth2_headers,
user_api_key_auth=user_api_key_auth,
raw_headers=raw_headers,
resource_metadata=resource_metadata,
)
# Pass-through OAuth: when the admin has opted a server into
@ -2030,6 +2108,20 @@ if MCP_AVAILABLE:
user_api_key_auth = await _apply_toolset_scope(user_api_key_auth, active_toolset_id)
toolset_allowed_server_ids = await _toolset_server_ids(active_toolset_id)
named_session_id: Final = _get_session_id_from_scope(scope)
names_live_session: Final = (
named_session_id is not None and named_session_id in _stateful_server_instances()
)
request_owner: Final = _owner_fingerprint_for(user_api_key_auth, oauth2_headers, _client_ip)
expected_owner: Final = (
_stateful_session_owners.get(named_session_id) if named_session_id is not None else None
)
owner_mismatch: Final = expected_owner is not None and expected_owner != request_owner
connect_peek: Final = _ConnectBodyPeek(
receive, peekable=scope.get("method") == "POST" and not names_live_session and not owner_mismatch
)
receive = connect_peek.receive
# https://datatracker.ietf.org/doc/html/rfc9728#name-www-authenticate-response
# Must run after toolset scoping so the challenge set is derived
# from the fully-authorized server set: a passthrough server that
@ -2042,6 +2134,7 @@ if MCP_AVAILABLE:
mcp_server_auth_headers=mcp_server_auth_headers,
user_api_key_auth=user_api_key_auth,
client_ip=_client_ip,
connecting=connect_peek.connecting,
allowed_server_ids=toolset_allowed_server_ids,
raw_headers=raw_headers,
)
@ -2077,8 +2170,6 @@ if MCP_AVAILABLE:
# - No session ID + initialize → stateful (so client gets mcp-session-id)
# - No session ID + other → stateless (curl, Inspector, Notion)
session_id = _get_session_id_from_scope(scope)
is_initialize = False
consumed_messages: list[Message] = []
# Owner-binding: a live stateful session may only be driven by the
# caller that created it. Reject mismatches with 403 so a leaked
@ -2086,12 +2177,9 @@ if MCP_AVAILABLE:
#
# Run before ``_handle_stale_mcp_session`` so a non-owner cannot
# force-clean another caller's residual tracking entries via a
# stale DELETE, and before peeking the request body so the 403
# response sees a pristine ``receive`` channel.
# stale DELETE.
if session_id:
expected_owner: Final = _stateful_session_owners.get(session_id)
request_owner = _owner_fingerprint_for(user_api_key_auth, oauth2_headers, _client_ip)
if expected_owner is not None and expected_owner != request_owner:
if owner_mismatch:
verbose_logger.warning(
"Rejecting MCP request: session '%s' owner mismatch.",
session_id,
@ -2116,10 +2204,11 @@ if MCP_AVAILABLE:
return
session_id = _get_session_id_from_scope(scope)
body = b""
if scope.get("method") == "POST":
consumed_messages, body = await _read_request_body_for_routing(receive)
is_initialize = _is_initialize_request(body)
session_body: Final = (
await connect_peek.read() if scope.get("method") == "POST" and names_live_session else b""
)
body: Final = await connect_peek.body() or session_body
is_initialize: Final = _is_initialize_request(body)
use_stateful: Final = bool(session_id or is_initialize)
target_manager: Final = session_manager_stateful if use_stateful else session_manager_stateless
@ -2134,7 +2223,6 @@ if MCP_AVAILABLE:
# session. Cap how many a single caller can hold so an authenticated
# client cannot spam `initialize` and exhaust memory.
if is_initialize and not session_id:
request_owner = _owner_fingerprint_for(user_api_key_auth, oauth2_headers, _client_ip)
if not await _enforce_stateful_session_cap_for_owner(request_owner):
verbose_logger.warning(
"Rejecting MCP initialize: caller already holds the maximum number of active stateful sessions."
@ -2149,17 +2237,6 @@ if MCP_AVAILABLE:
await too_many_response(scope, receive, send)
return
# Replay body messages if we consumed them for peeking
original_receive: Final = receive
if consumed_messages:
async def wrapped_receive():
if consumed_messages:
return consumed_messages.pop(0)
return await original_receive()
receive = wrapped_receive
# Serialize requests on the same stateful session so concurrent
# callers don't clobber each other's auth context mid-flight.
#
@ -2389,6 +2466,7 @@ if MCP_AVAILABLE:
mcp_server_auth_headers=mcp_server_auth_headers,
user_api_key_auth=user_api_key_auth,
client_ip=_sse_client_ip,
connecting=_known_connecting(scope["method"] == "GET"),
allowed_server_ids=toolset_allowed_server_ids,
raw_headers=raw_headers,
)

View file

@ -362,6 +362,14 @@ def iter_known_server_prefixes(server: _McpServerLike) -> Iterator[str]:
yield from _emit(server_id)
def server_answers_to_name(server: _McpServerLike, name: str) -> bool:
"""Whether a scoped ``/mcp/{name}`` connect resolves to ``server``: case-insensitive over every prefix
form routing accepts (alias, server_name, server_id, short prefix), the same match
``_server_answers_to`` applies when the router scopes a request."""
requested: Final = name.lower()
return any(requested == known.lower() for known in iter_known_server_prefixes(server) if known)
def iter_known_tool_name_spellings(tool_name: str, server: MCPServer) -> Iterator[str]:
"""Yield every name that denotes the bare ``tool_name`` on ``server``: the bare name,
then its wire spelling under each prefix ``iter_known_server_prefixes`` accepts.

View file

@ -7,6 +7,7 @@ from .agent_365 import Agent365Guardrail
if TYPE_CHECKING:
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchanger import TokenExchanger
from litellm.types.guardrails import Guardrail, LitellmParams
@ -15,6 +16,7 @@ def initialize_guardrail(
guardrail: "Guardrail",
*,
async_handler: "AsyncHTTPHandler | None" = None,
token_exchanger: "TokenExchanger | None" = None,
) -> Agent365Guardrail:
import litellm
from litellm.secret_managers.main import get_secret_str
@ -64,6 +66,7 @@ def initialize_guardrail(
request_timeout=litellm_params.timeout if litellm_params.timeout is not None else 10.0,
unreachable_fallback=litellm_params.unreachable_fallback,
async_handler=async_handler,
token_exchanger=token_exchanger,
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
)

View file

@ -9,17 +9,14 @@ exchanged for a delegated Agent 365 token, so Defender evaluates and audits
as the signed-in user.
"""
import hashlib
import threading
import time
import uuid
from collections import OrderedDict
from collections.abc import Mapping
from typing import TYPE_CHECKING, ClassVar, Final, Literal, NoReturn
import httpx
from fastapi import HTTPException
from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError
from pydantic import BaseModel, ConfigDict, Field, SecretStr, TypeAdapter, ValidationError
from typing_extensions import ReadOnly, TypedDict
from litellm._logging import verbose_proxy_logger
@ -34,7 +31,30 @@ from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
)
from litellm.proxy._experimental.mcp_server.caller_sign_in import (
CallerSignIn,
CallerSignInPreflight,
Rejected,
SignedIn,
Unavailable,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import OAuthToken
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Error, Ok, Result
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchange_provider import (
build_token_exchanger,
oauth_error_fields,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchanger import (
SubjectTokenRejected,
TokenExchanger,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
CredError,
ServerSpec,
TokenExchangeConfig,
)
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.mcp_server.mcp_server_manager import MCPServer
from litellm.types.proxy.guardrails.guardrail_hooks.agent_365 import (
AGENT_365_PROD_API_BASE,
AGENT_365_PROD_RESOURCE_APP_ID,
@ -49,43 +69,19 @@ if TYPE_CHECKING:
from litellm.types.utils import GuardrailStatus
TOKEN_ENDPOINT_TEMPLATE: Final = "https://login.microsoftonline.com/{tenant_id}/oauth2/v2.0/token"
ENTRA_ISSUER_TEMPLATE: Final = "https://login.microsoftonline.com/{tenant_id}/v2.0"
EVALUATE_URL: Final = f"{AGENT_365_PROD_API_BASE}/agents/tool-evaluation/evaluate"
OBO_SCOPE: Final = f"{AGENT_365_PROD_RESOURCE_APP_ID}/{AGENT_365_SCOPE_NAME}"
MCP_SESSION_ID_HEADER: Final = "mcp-session-id"
DEFENDER_STATUS_EVALUATED: Final = "Evaluated"
_GATEWAY_OWNED_TOKEN_ERRORS: Final = frozenset(
{"invalid_client", "unauthorized_client", "invalid_scope", "invalid_resource"}
)
# Entra reports a malformed or unverifiable assertion as ``invalid_client`` too; only its AADSTS50027xx
# (InvalidJwtToken) sub-codes tell that apart from a bad gateway secret.
_INVALID_ASSERTION_AADSTS_PREFIX: Final = "50027"
_AADSTS_CODES_ADAPTER: Final = TypeAdapter(tuple[int, ...])
GATEWAY_SCOPE_TEMPLATE: Final = "api://{client_id}/access_as_user"
_MCP_CALL_TYPES: Final[tuple[str, ...]] = ("mcp_call", "call_mcp_tool")
_TOOL_INPUT_SCHEMA_ADAPTER: Final = TypeAdapter(dict[str, object])
_OBO_CACHE_MAX_ENTRIES: Final = 1000
_DEFAULT_TOKEN_TTL_SECONDS: Final = 3599.0
_TOKEN_EXPIRY_SLACK_SECONDS: Final = 60.0
def _parse_expires_in(raw: object) -> float:
if not isinstance(raw, (int, float, str)):
return _DEFAULT_TOKEN_TTL_SECONDS
try:
return float(raw)
except ValueError:
return _DEFAULT_TOKEN_TTL_SECONDS
def _parse_aadsts_codes(raw: object) -> tuple[int, ...]:
try:
return _AADSTS_CODES_ADAPTER.validate_python(raw)
except ValidationError:
return ()
_JSON_OBJECT_ADAPTER: Final = TypeAdapter(dict[str, object])
def _parse_tool_input_schema(raw: object) -> Mapping[str, object] | None:
try:
return _TOOL_INPUT_SCHEMA_ADAPTER.validate_python(raw)
return _JSON_OBJECT_ADAPTER.validate_python(raw)
except ValidationError:
return None
@ -108,6 +104,15 @@ class _EvaluateResponse(TypedDict, total=False):
correlationId: ReadOnly[str]
class _AdmissionMetadata(TypedDict):
user_api_key_metadata: ReadOnly[dict | None]
user_api_key_team_metadata: ReadOnly[dict | None]
class _AdmissionProbe(TypedDict):
metadata: ReadOnly[_AdmissionMetadata]
class _ToolReference(BaseModel):
model_config = ConfigDict(frozen=True)
@ -130,20 +135,12 @@ class _BlockedDetail(TypedDict):
class Agent365TokenExchangeError(Exception):
def __init__(self, status_code: int, error_code: str, description: str, aadsts_codes: tuple[int, ...] = ()) -> None:
super().__init__(f"{error_code}: {description}")
self.status_code = status_code
self.error_code = error_code
self.description = description
self.aadsts_codes = aadsts_codes
"""Entra refused the gateway's own client credentials, scope or resource; the caller cannot fix that by
signing in again, so it is the gateway's outage, never a 401."""
@property
def gateway_owned(self) -> bool:
"""Whether the gateway's own client credentials, scope or resource were refused, as opposed to the
caller's assertion. The caller cannot fix a gateway-owned rejection by signing in again."""
if self.error_code not in _GATEWAY_OWNED_TOKEN_ERRORS:
return False
return not any(str(code).startswith(_INVALID_ASSERTION_AADSTS_PREFIX) for code in self.aadsts_codes)
def __init__(self, error_code: str) -> None:
super().__init__(error_code)
self.error_code = error_code
class Agent365MalformedResponseError(Exception):
@ -156,6 +153,13 @@ class Agent365ThrottledError(Exception):
self.status_code = status_code
def _gateway_fault_reason(error_code: str) -> str:
return (
f"Entra rejected the gateway's own Agent 365 credentials ({error_code}); "
"check the guardrail's client_id and client_secret"
)
class Agent365Guardrail(CustomGuardrail):
"""Pre-MCP-call guardrail enforcing Microsoft Agent 365 tool-evaluation verdicts.
@ -173,6 +177,7 @@ class Agent365Guardrail(CustomGuardrail):
request_timeout: float = 10.0,
unreachable_fallback: Literal["fail_closed", "fail_open"] = "fail_closed",
async_handler: AsyncHTTPHandler | None = None,
token_exchanger: TokenExchanger | None = None,
**kwargs, # noqa: ANN003 # kwargs-ok: forwarded verbatim to CustomGuardrail (event_hook, default_on)
) -> None:
super().__init__(
@ -192,8 +197,23 @@ class Agent365Guardrail(CustomGuardrail):
self.async_handler = async_handler or get_async_httpx_client(
llm_provider=httpxSpecialProvider.GuardrailCallback
)
self._obo_token_cache: OrderedDict[str, tuple[str, float]] = OrderedDict() # mutable-ok: lock-guarded LRU
self._obo_cache_lock = threading.Lock()
self._exchange_config: Final = TokenExchangeConfig(
profile="entra_obo",
token_exchange_endpoint=TOKEN_ENDPOINT_TEMPLATE.format(tenant_id=tenant_id),
client_id=client_id,
client_secret=SecretStr(client_secret),
scopes=(OBO_SCOPE,),
)
self._exchange_server: Final = ServerSpec(
server_id=f"agent-365:{tenant_id}",
resource=AGENT_365_PROD_API_BASE,
config=self._exchange_config,
)
self._token_exchanger: Final = (
token_exchanger
if token_exchanger is not None
else build_token_exchanger(post=self._post_entra_token_endpoint)
)
verbose_proxy_logger.info("Initialized Microsoft Agent 365 guardrail: %s", guardrail_name)
@staticmethod
@ -220,7 +240,7 @@ class Agent365Guardrail(CustomGuardrail):
return data
tool_name: Final = str(data.get("mcp_tool_name") or "")
assertion: Final = entra_assertion(data.get("incoming_bearer_token"))
assertion: Final = entra_assertion(data.get("incoming_subject_token"))
if assertion is None:
self._handle_caller_fault(
data=data,
@ -233,22 +253,10 @@ class Agent365Guardrail(CustomGuardrail):
)
try:
obo_token: Final = await self._get_obo_token(assertion)
exchange_result: Final = await self._exchange_caller_assertion(assertion)
except Agent365TokenExchangeError as exc:
if exc.gateway_owned:
return self._handle_unavailable(
data=data,
tool_name=tool_name,
reason=(
f"Entra rejected the gateway's own Agent 365 credentials ({exc.error_code}); "
"check the guardrail's client_id and client_secret"
),
)
self._handle_caller_fault(
data=data,
tool_name=tool_name,
status_code=401,
reason=f"the Entra On-Behalf-Of token exchange was rejected ({exc.error_code})",
return self._handle_unavailable(
data=data, tool_name=tool_name, reason=_gateway_fault_reason(exc.error_code)
)
except Agent365ThrottledError as exc:
self._handle_throttled(
@ -264,11 +272,25 @@ class Agent365Guardrail(CustomGuardrail):
reason=f"the Entra token endpoint could not be reached ({type(exc).__name__})",
)
except Agent365MalformedResponseError as exc:
return self._handle_unavailable(
data=data,
tool_name=tool_name,
reason=str(exc),
)
return self._handle_unavailable(data=data, tool_name=tool_name, reason=str(exc))
match exchange_result:
case Ok(token):
obo_token: Final = token.access_token
case Error(error):
match error.tag:
case "unauthorized":
self._handle_caller_fault(
data=data,
tool_name=tool_name,
status_code=401,
reason=f"the Entra On-Behalf-Of token exchange was rejected ({error.unauthorized.detail})",
)
case _:
return self._handle_unavailable(
data=data,
tool_name=tool_name,
reason=f"the Entra token exchange failed ({error.summary})",
)
start: Final = time.perf_counter()
try:
@ -284,14 +306,14 @@ class Agent365Guardrail(CustomGuardrail):
reason=f"the Agent 365 endpoint could not be reached ({type(exc).__name__})",
)
latency_ms: Final = (time.perf_counter() - start) * 1000.0
fallback: Final = self._handle_evaluate_error(
fallback: Final = await self._handle_evaluate_error(
data=data, tool_name=tool_name, assertion=assertion, response=response, latency_ms=latency_ms
)
if fallback is not None:
return fallback
return self._enforce_verdict(data=data, tool_name=tool_name, response=response, latency_ms=latency_ms)
def _handle_evaluate_error(
async def _handle_evaluate_error(
self,
data: dict, # mutable-ok: guardrail logging appends into the request metadata in place
tool_name: str,
@ -308,7 +330,9 @@ class Agent365Guardrail(CustomGuardrail):
)
if 400 <= response.status_code < 500:
if response.status_code == 401:
self._evict_obo_token(assertion)
await self._token_exchanger.invalidate(
assertion, self._exchange_server, self._exchange_config, tenant_id=self.tenant_id
)
self._record_verdict(
data=data,
verdict="Rejected",
@ -455,26 +479,54 @@ class Agent365Guardrail(CustomGuardrail):
return call_id
return str(uuid.uuid4())
async def _get_obo_token(self, assertion: str) -> str:
cache_key: Final = hashlib.sha256(assertion.encode("utf-8")).hexdigest()
now: Final = time.time()
with self._obo_cache_lock:
cached: Final = self._obo_token_cache.get(cache_key)
if cached and cached[1] > now + _TOKEN_EXPIRY_SLACK_SECONDS:
self._obo_token_cache.move_to_end(cache_key)
return cached[0]
def caller_sign_in(self, server: MCPServer, user_api_key_auth: "UserAPIKeyAuth | None") -> CallerSignIn | None:
"""The Entra sign-in this guardrail requires of callers: only a ``default_on`` guardrail whose mode gates a
tagless MCP connect and that the caller's key or team has not opted out of, because the anonymous
metadata fetch that follows a challenge cannot see which key selected a guardrail and would advertise
the wrong issuer. Only servers that leave the caller's top-level ``Authorization`` with the gateway
qualify: a forwarded API-key header travels upstream in its own slot and does not displace the Entra
assertion."""
if not (self.default_on and server.keeps_caller_authorization):
return None
probe: Final[_AdmissionProbe] = {
"metadata": {
"user_api_key_metadata": user_api_key_auth.metadata if user_api_key_auth else None, # pyright: ignore[reportUnknownMemberType] # UserAPIKeyAuth.metadata is a raw dict
"user_api_key_team_metadata": user_api_key_auth.team_metadata if user_api_key_auth else None, # pyright: ignore[reportUnknownMemberType] # UserAPIKeyAuth.team_metadata is a raw dict
}
}
if (
self.should_run_guardrail( # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # should_run_guardrail takes an untyped data dict
data=probe, event_type=GuardrailEventHooks.pre_mcp_call
)
is not True
):
return None
return CallerSignIn(
issuers=(ENTRA_ISSUER_TEMPLATE.format(tenant_id=self.tenant_id),),
scopes=tuple(server.scopes)
if server.scopes
else (GATEWAY_SCOPE_TEMPLATE.format(client_id=self.client_id),),
)
async def _exchange_caller_assertion(self, assertion: str) -> Result[OAuthToken, CredError]:
"""The one Entra OBO exchange call both the tool-call path and the connect preflight run; a
successful result is cached by the exchanger, so the session reuses what the preflight minted."""
return await self._token_exchanger.exchange(
assertion, self._exchange_server, self._exchange_config, tenant_id=self.tenant_id
)
async def _post_entra_token_endpoint(
self,
url: str,
form: dict[str, str], # mutable-ok: ExchangeHttpPost contract
client_auth_headers: dict[str, str], # mutable-ok: ExchangeHttpPost contract
) -> dict[str, object] | None:
"""The exchanger's HTTP edge for this guardrail: every way Entra can fail keeps its own exception, so
the verdict reason the Logs row carries names the OAuth error code or the transport fault."""
response: Final = await self._post_allowing_error_status(
url=TOKEN_ENDPOINT_TEMPLATE.format(tenant_id=self.tenant_id),
data={
"grant_type": "urn:ietf:params:oauth:grant-type:jwt-bearer",
"client_id": self.client_id,
"client_secret": self.client_secret,
"assertion": assertion,
"scope": OBO_SCOPE,
"requested_token_use": "on_behalf_of",
},
headers={"Content-Type": "application/x-www-form-urlencoded"},
url=url,
data=form,
headers={"Content-Type": "application/x-www-form-urlencoded", **client_auth_headers},
)
if response.status_code in (408, 429):
raise Agent365ThrottledError(status_code=response.status_code)
@ -485,32 +537,54 @@ class Agent365Guardrail(CustomGuardrail):
response=response,
)
try:
parsed_body: Final = response.json()
parsed_body: Final[object] = response.json()
except ValueError as exc:
raise Agent365MalformedResponseError("the Entra token endpoint returned a non-JSON body") from exc
if not isinstance(parsed_body, dict):
raise Agent365MalformedResponseError("the Entra token endpoint returned a non-object JSON body")
body: Final = parsed_body
try:
body: Final = _JSON_OBJECT_ADAPTER.validate_python(parsed_body)
except ValidationError as exc:
raise Agent365MalformedResponseError("the Entra token endpoint returned a non-object JSON body") from exc
if response.status_code >= 400:
raise Agent365TokenExchangeError(
status_code=response.status_code,
error_code=str(body.get("error", "invalid_grant")),
description=str(body.get("error_description", ""))[:512],
aadsts_codes=_parse_aadsts_codes(body.get("error_codes")),
)
oauth_error: Final = oauth_error_fields(response)
gateway_fault: Final = oauth_error.gateway_fault
if gateway_fault is not None:
raise Agent365TokenExchangeError(error_code=gateway_fault)
raise SubjectTokenRejected(oauth_error.error or "invalid_grant", claims=oauth_error.claims)
if "access_token" not in body:
raise Agent365MalformedResponseError("the Entra token endpoint returned no access_token")
raw_access_token: Final = body.get("access_token")
if not isinstance(raw_access_token, str) or not raw_access_token:
access_token: Final = body["access_token"]
if not isinstance(access_token, str) or not access_token:
raise Agent365MalformedResponseError("the Entra token endpoint returned a non-string access_token")
access_token: Final = raw_access_token
expires_at: Final = time.time() + _parse_expires_in(body.get("expires_in", 3599))
with self._obo_cache_lock:
self._obo_token_cache[cache_key] = (access_token, expires_at)
self._obo_token_cache.move_to_end(cache_key)
while len(self._obo_token_cache) > _OBO_CACHE_MAX_ENTRIES:
self._obo_token_cache.popitem(last=False)
return access_token
return body
async def preflight_caller_sign_in(
self, server: MCPServer, user_api_key_auth: "UserAPIKeyAuth | None", subject_token: str
) -> CallerSignInPreflight:
"""The connect-time check the preemptive gate runs: a bearer Entra rejects, or one it could never
accept, gets the sign-in challenge here, where ``WWW-Authenticate`` still reaches the client, instead of
surfacing as a JSON-RPC error on every tools/call. ``subject_token=None`` stays the challenge gate's job."""
assertion: Final = entra_assertion(subject_token)
if assertion is None:
return Rejected(detail="the caller's bearer is not an Entra token; sign in with Entra and retry")
fail_open: Final = self.unreachable_fallback == "fail_open"
try:
exchange_result: Final = await self._exchange_caller_assertion(assertion)
except Agent365TokenExchangeError as exc:
return Unavailable(detail=_gateway_fault_reason(exc.error_code), fail_open=fail_open)
except Agent365ThrottledError as exc:
return Unavailable(detail=f"the Entra token endpoint returned HTTP {exc.status_code}", fail_open=False)
except (httpx.HTTPError, LitellmTimeout, TimeoutError) as exc:
return Unavailable(
detail=f"the Entra token endpoint could not be reached ({type(exc).__name__})", fail_open=fail_open
)
except Agent365MalformedResponseError as exc:
return Unavailable(detail=str(exc), fail_open=fail_open)
if isinstance(exchange_result, Ok):
return SignedIn()
error: Final = exchange_result.error
if error.tag == "unauthorized":
return Rejected(detail=error.unauthorized.detail, claims=error.unauthorized.claims)
return Unavailable(detail=f"the Entra token exchange failed ({error.summary})", fail_open=fail_open)
async def _post_allowing_error_status(
self,
@ -577,11 +651,6 @@ class Agent365Guardrail(CustomGuardrail):
}
raise HTTPException(status_code=503, detail=throttled_detail)
def _evict_obo_token(self, assertion: str) -> None:
cache_key: Final = hashlib.sha256(assertion.encode("utf-8")).hexdigest()
with self._obo_cache_lock:
self._obo_token_cache.pop(cache_key, None)
def _handle_unavailable(
self,
data: dict, # mutable-ok: guardrail logging appends into the request metadata in place

View file

@ -1518,6 +1518,7 @@ class ProxyLogging:
# (e.g. MCPJWTSigner) to independently verify the caller's identity
# before re-signing an outbound token (FR-5 verify+re-sign).
"incoming_bearer_token": kwargs.get("incoming_bearer_token"),
"incoming_subject_token": kwargs.get("incoming_subject_token"),
"metadata": synthetic_metadata,
}
user_api_key_auth: Final = kwargs.get("user_api_key_auth")

View file

@ -315,10 +315,10 @@ class MCPServer(BaseModel):
return self.per_server_oauth_discovery and self.auth_type == MCPAuth.oauth2 and not self.has_client_credentials
@property
def advertises_gateway_authorization_server(self) -> bool:
"""Whether named discovery should advertise the aggregate gateway authorization server."""
if self.auth_type == MCPAuth.oauth2:
return self.is_gateway_managed_oauth2 and not self.uses_per_server_oauth_relay
def keeps_caller_authorization(self) -> bool:
"""Whether the caller's top-level ``Authorization`` stays with the gateway: the server neither relays
it upstream nor runs an OAuth mode that fills that slot itself, so a gateway guardrail may consume it
as the caller's own assertion. Forwarding a separate API-key header leaves the slot untouched."""
if self.auth_type not in (
None,
MCPAuth.none,
@ -328,11 +328,20 @@ class MCPServer(BaseModel):
MCPAuth.authorization,
MCPAuth.token,
MCPAuth.aws_sigv4,
MCPAuth.oauth2_token_exchange,
):
return False
return not any(
header.lower() in ("authorization", "x-api-key", "api-key", "apikey")
for header in (self.extra_headers or ())
return not any(header.lower() == "authorization" for header in (self.extra_headers or ()))
@property
def advertises_gateway_authorization_server(self) -> bool:
"""Whether named discovery should advertise the aggregate gateway authorization server."""
if self.auth_type == MCPAuth.oauth2:
return self.is_gateway_managed_oauth2 and not self.uses_per_server_oauth_relay
if self.auth_type == MCPAuth.oauth2_token_exchange:
return False
return self.keeps_caller_authorization and not any(
header.lower() in ("x-api-key", "api-key", "apikey") for header in (self.extra_headers or ())
)
@property

View file

@ -1,3 +1,4 @@
import json
import uuid
from typing import Final
@ -92,6 +93,36 @@ def test_subject_grant_lists_only_reachable_tools_and_denies_the_rest(
assert not any(name.startswith(denied_alias) for name in denied_listed.tools), denied_listed.tools
@pytest.mark.parametrize("entry", ("mcp", "server_mcp"))
def test_access_group_named_like_an_ungranted_server_still_routes_the_groups_servers(
gateway: Gateway, entry: EntryPoint
) -> None:
with peer_of("http") as shadow, peer_of("http") as member, gateway.scenario() as scenario:
group: Final = "docs" + uuid.uuid4().hex[:8]
member_alias: Final = "mem" + uuid.uuid4().hex[:8]
register_mcp(scenario, shadow, group)
register_mcp(scenario, member, member_alias, mcp_access_groups=[group])
key: Final = scenario.key(object_permission={"mcp_access_groups": [group]})
unmatched: Final = "none" + uuid.uuid4().hex[:8]
denied: Final = McpCaller(
gateway, key, entry, unmatched, {"x-mcp-servers": unmatched} if entry == "mcp" else {}
).list_tools()
assert (denied.status, json.loads(denied.error or "null"), denied.tools) == (
(
200,
{"code": -32600, "message": f"The key is not allowed to access the requested MCP servers: {unmatched}"},
(),
)
if entry == "mcp"
else (404, {"detail": f"MCP server, toolset, or access group '{unmatched}' not found"}, ())
), denied.raw
selection: Final = {"x-mcp-servers": group} if entry == "mcp" else {}
listed: Final = McpCaller(gateway, key, entry, group, selection).list_tools()
assert listed.ok, listed.raw
assert set(listed.tools) == {f"{member_alias}-{tool}" for tool in ("add", "multiply", "fail")}, listed.tools
assert tool_calls(shadow.drain()) == ()
@pytest.mark.parametrize("entry", ("mcp", "server_mcp", "rest"))
def test_key_without_any_grant_sees_no_scoped_server(gateway: Gateway, entry: EntryPoint) -> None:
with peer_of("http") as peer, gateway.scenario() as scenario:

View file

@ -16,6 +16,7 @@ from integration._support.client import Gateway, eventually
from integration._support.database import read_rows
from integration._support.mcp import (
ENTRY_POINTS,
INITIALIZE,
EntryPoint,
McpCaller,
McpPeer,
@ -38,6 +39,10 @@ GUARDRAIL_ROWS: Final = (
FALLBACKS: Final = (None, "fail_open", "fail_closed")
def _origin(gateway: Gateway) -> str:
return str(gateway.client.base_url).rstrip("/")
def _generic_guardrail_outage(request: Request) -> Reply:
assert request.target == "/beta/litellm_basic_guardrail_api", request.target
return Reply(status=503, body=json.dumps({"error": "synthetic sibling guardrail outage"}).encode())
@ -136,15 +141,26 @@ def test_a_missing_or_malformed_caller_bearer_blocks_on_every_entry_point_whatev
) -> None:
with _rig(gateway, tmp_path, fallback) as rig:
assert f"{rig.alias}-add" in rig.caller().list_tools().tools, "the catalog needs only the virtual key"
metadata_url: Final = f"{_origin(rig.candidate)}/.well-known/oauth-protected-resource/{rig.alias}/mcp"
for entry in ENTRY_POINTS:
missing: Final = rig.caller(entry).call(f"{rig.alias}-add", {"entry": entry}, server_id=rig.server_id)
assert missing.error is not None and REJECTED in missing.raw, f"{entry} without a bearer: {missing.raw}"
malformed: Final = rig.caller(entry, "not-a-jws").call(
f"{rig.alias}-add", {"entry": entry}, server_id=rig.server_id
)
assert malformed.error is not None and REJECTED in malformed.raw, f"{entry} opaque bearer: {malformed.raw}"
for label, bearer in (("without a bearer", None), ("opaque bearer", "not-a-jws")):
caller: Final = rig.caller(entry, bearer)
if entry == "server_mcp":
challenged: Final = caller.rpc("initialize", INITIALIZE)
assert challenged.status_code == 401, f"{entry} {label}: {challenged.text}"
authenticate: Final = challenged.headers.get("www-authenticate", "")
assert f'resource_metadata="{metadata_url}"' in authenticate, f"{entry} {label}: {authenticate!r}"
assert 'error="invalid_token"' in authenticate, f"{entry} {label}: {authenticate!r}"
if bearer is None:
unsigned: Final = caller.rpc(
"tools/call", {"name": f"{rig.alias}-add", "arguments": {"entry": entry}}
)
assert unsigned.status_code == 401, f"{entry} {label}: {unsigned.text}"
continue
outcome: Final = caller.call(f"{rig.alias}-add", {"entry": entry}, server_id=rig.server_id)
assert outcome.error is not None and REJECTED in outcome.raw, f"{entry} {label}: {outcome.raw}"
assert rig.upstream_tool_names() == ()
expected: Final = 2 * len(ENTRY_POINTS)
expected: Final = 2 * len(ENTRY_POINTS) - 1
assert rig.guardrail_statuses("call_mcp_tool", expected) == ["guardrail_intervened"] * expected

View file

@ -0,0 +1,463 @@
import json
import uuid
from collections.abc import Mapping
from pathlib import Path
from types import MappingProxyType
from typing import Final
import httpx
import yaml
from cryptography.hazmat.primitives.asymmetric import rsa
from integration._support.client import Gateway
from integration._support.mcp import (
INITIALIZE,
McpCaller,
forget_mcp,
mcp_peer,
register_mcp,
tool_calls,
)
from integration._support.process import owned_proxy
from integration._support.wire import Reply, Request, wire_server
from jwt import algorithms as jwt_algorithms
ADD: Final = {"a": 2, "b": 3}
ACCEPT: Final = {"Accept": "application/json, text/event-stream"}
def _rpc(
gateway: Gateway,
path: str,
key: str,
headers: dict[str, str],
method: str = "initialize",
params: Mapping[str, object] = INITIALIZE,
) -> httpx.Response:
return gateway.client.post(
path,
json={"jsonrpc": "2.0", "id": 1, "method": method, "params": dict(params)},
headers={"x-litellm-api-key": key, **ACCEPT, **headers},
)
def _sse_data(response: httpx.Response) -> str:
return next(line[5:].strip() for line in response.text.splitlines() if line.startswith("data:"))
def _advertised(gateway: Gateway, segment: str) -> tuple[int, tuple[str, ...], object]:
response: Final = gateway.client.get(f"/.well-known/oauth-protected-resource/mcp/{segment}")
document: Final = response.json()
issuers: Final = tuple(
str(issuer).removesuffix(f"/{segment}") for issuer in document.get("authorization_servers", ())
)
return response.status_code, issuers, document.get("scopes_supported")
def _origin(gateway: Gateway) -> str:
return str(gateway.client.base_url).rstrip("/")
def _sign_in_config(
guardrail_params: dict[str, object], path: Path, general_settings: Mapping[str, object] = MappingProxyType({})
) -> Path:
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["guardrails"] = [{"guardrail_name": "signin" + uuid.uuid4().hex, "litellm_params": guardrail_params}]
config["general_settings"] = {**config.get("general_settings", {}), **general_settings}
path.write_text(yaml.safe_dump(config))
return path
def test_multi_server_connect_with_a_litellm_key_in_authorization_admits_an_obo_server(
gateway: Gateway,
) -> None:
with mcp_peer() as obo_peer, mcp_peer() as math_peer, gateway.scenario() as scenario:
obo_alias: Final = "obo" + uuid.uuid4().hex[:8]
math_alias: Final = "math" + uuid.uuid4().hex[:8]
obo_id: Final = register_mcp(
scenario,
obo_peer,
obo_alias,
auth_type="oauth2_token_exchange",
token_exchange_endpoint="http://127.0.0.1:9/token",
credentials={"client_id": "obo-client", "client_secret": "obo-secret"},
)
math_id: Final = register_mcp(scenario, math_peer, math_alias)
key: Final = scenario.key(object_permission={"mcp_servers": [obo_id, math_id]})
caller: Final = McpCaller(
gateway,
key,
"root",
headers={
"x-mcp-servers": f"{obo_alias},{math_alias}",
"Authorization": f"Bearer {key}",
},
)
init: Final = caller.initialize()
assert init.ok, init.raw
listed: Final = caller.list_tools()
assert listed.ok, listed.raw
assert any(name.endswith("add") for name in listed.tools), listed.tools
def test_alias_first_lookup_wins_over_a_server_whose_name_matches_the_alias(gateway: Gateway, tmp_path: Path) -> None:
with mcp_peer() as math_peer, mcp_peer() as obo_peer:
name: Final = "gh" + uuid.uuid4().hex[:6]
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["mcp_servers"] = {
name: {
"transport": "http",
"url": obo_peer.url,
"auth_type": "oauth2_token_exchange",
"token_exchange_endpoint": "http://127.0.0.1:9/token",
"credentials": {"client_id": "obo-client", "client_secret": "obo-secret"},
}
}
path: Final = tmp_path / "collision.yaml"
path.write_text(yaml.safe_dump(config))
with (
owned_proxy(gateway, tmp_path, {}, config=path) as candidate,
candidate.scenario() as scenario,
):
second: Final = candidate.request(
"POST",
"/v1/mcp/server",
{"server_name": name + "_public", "alias": name, **math_peer.registration()},
)
assert second.status_code == 201, second.text
second_id: Final = second.json()["server_id"]
scenario.cleanups.callback(forget_mcp, candidate, second_id)
key: Final = scenario.key(object_permission={"mcp_servers": [second_id]})
init: Final = _rpc(candidate, f"/mcp/{name}", key, {})
assert init.status_code == 200, init.text
listing: Final = candidate.client.post(
f"/mcp/{name}",
json={"jsonrpc": "2.0", "id": 2, "method": "tools/list", "params": {}},
headers={"x-litellm-api-key": key, **ACCEPT},
)
assert listing.status_code == 200, listing.text
body: Final = json.loads(
next(line[5:].strip() for line in listing.text.splitlines() if line.startswith("data:"))
)
names: Final = {tool["name"] for tool in body["result"]["tools"]}
assert any(name.endswith("add") for name in names), names
def test_exact_name_wins_over_a_case_folded_config_alias_for_connect_discovery_and_calls(
gateway: Gateway, tmp_path: Path
) -> None:
with mcp_peer() as obo_peer, mcp_peer() as math_peer:
stem: Final = "gh" + uuid.uuid4().hex[:6]
cased: Final = stem.capitalize()
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["mcp_servers"] = {
stem + "_obo": {
"alias": stem,
"transport": "http",
"url": obo_peer.url,
"auth_type": "oauth2_token_exchange",
"token_exchange_endpoint": "http://127.0.0.1:9/token",
"credentials": {"client_id": "obo-client", "client_secret": "obo-secret"},
}
}
path: Final = tmp_path / "case-collision.yaml"
path.write_text(yaml.safe_dump(config))
with owned_proxy(gateway, tmp_path, {}, config=path) as candidate, candidate.scenario() as scenario:
by_name: Final = register_mcp(scenario, math_peer, cased)
key: Final = scenario.key(object_permission={"mcp_servers": [by_name]})
init: Final = _rpc(candidate, f"/mcp/{cased}", key, {})
assert init.status_code == 200, init.text
listed: Final = _rpc(candidate, f"/mcp/{cased}", key, {}, method="tools/list")
assert listed.status_code == 200, listed.text
tools: Final = json.loads(_sse_data(listed))["result"]["tools"]
add: Final = next(tool["name"] for tool in tools if tool["name"].endswith("add"))
called: Final = _rpc(
candidate, f"/mcp/{cased}", key, {}, method="tools/call", params={"name": add, "arguments": ADD}
)
assert called.status_code == 200, called.text
assert json.loads(_sse_data(called))["result"]["content"][0]["text"] == "5", called.text
assert len(tool_calls(math_peer.drain())) == 1
assert tool_calls(obo_peer.drain()) == ()
assert _advertised(candidate, cased) == _advertised(candidate, by_name)
same_name: Final = _rpc(candidate, f"/mcp/{stem}", key, {})
assert same_name.status_code == 200, same_name.text
assert "www-authenticate" not in same_name.headers, same_name.headers
via_alias: Final = _rpc(
candidate, f"/mcp/{stem}", key, {}, method="tools/call", params={"name": add, "arguments": ADD}
)
assert json.loads(_sse_data(via_alias))["result"]["content"][0]["text"] == "5", via_alias.text
assert len(tool_calls(math_peer.drain())) == 1
assert tool_calls(obo_peer.drain()) == ()
challenged: Final = _rpc(candidate, f"/mcp/{stem}", scenario.key(), {})
assert challenged.status_code == 401, challenged.text
assert (
f'resource_metadata="{_origin(candidate)}/.well-known/oauth-protected-resource/mcp/{stem}"'
in challenged.headers.get("www-authenticate", "")
)
assert _advertised(candidate, stem) == _advertised(candidate, stem + "_obo")
assert _advertised(candidate, stem) != _advertised(candidate, cased)
def test_name_of_a_server_hidden_from_an_external_ip_does_not_reroute_to_a_case_variant(
gateway: Gateway, tmp_path: Path
) -> None:
stem: Final = "gh" + uuid.uuid4().hex[:6]
cased: Final = stem.capitalize()
external: Final = {"X-Forwarded-For": "203.0.113.7"}
with mcp_peer() as hidden_peer, mcp_peer() as public_peer:
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["general_settings"] = {
**config.get("general_settings", {}),
"use_x_forwarded_for": True,
"mcp_trusted_proxy_ranges": ["127.0.0.0/8"],
}
config["mcp_servers"] = {
stem: {"transport": "http", "url": hidden_peer.url, "available_on_public_internet": False}
}
path: Final = tmp_path / "external-ip.yaml"
path.write_text(yaml.safe_dump(config))
with owned_proxy(gateway, tmp_path, {}, config=path) as candidate, candidate.scenario() as scenario:
public_id: Final = register_mcp(scenario, public_peer, cased)
key: Final = scenario.key(object_permission={"mcp_servers": [public_id]})
hidden_name: Final = _rpc(candidate, f"/mcp/{stem}", key, external)
assert hidden_name.status_code == 403, hidden_name.text
assert "www-authenticate" not in hidden_name.headers, hidden_name.headers
assert (
candidate.client.get(f"/.well-known/oauth-protected-resource/mcp/{stem}", headers=external).status_code
== 404
)
assert _rpc(candidate, f"/mcp/{stem}", key, {}).status_code == 200
own_name: Final = _rpc(candidate, f"/mcp/{cased}", key, external)
assert own_name.status_code == 200, own_name.text
listed: Final = _rpc(candidate, f"/mcp/{cased}", key, external, method="tools/list")
add: Final = next(
tool["name"]
for tool in json.loads(_sse_data(listed))["result"]["tools"]
if tool["name"].endswith("add")
)
called: Final = _rpc(
candidate, f"/mcp/{cased}", key, external, method="tools/call", params={"name": add, "arguments": ADD}
)
assert called.status_code == 200, called.text
assert json.loads(_sse_data(called))["result"]["content"][0]["text"] == "5", called.text
assert len(tool_calls(public_peer.drain())) == 1
assert tool_calls(hidden_peer.drain()) == ()
def test_jwt_signer_verifies_the_bearer_that_admitted_the_call(gateway: Gateway, tmp_path: Path) -> None:
signer_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048)
jwk: Final = json.loads(jwt_algorithms.RSAAlgorithm.to_jwk(signer_key.public_key()))
jwk["kid"] = "idp"
holder: Final = []
def idp(request: Request) -> Reply:
issuer: Final = holder[0].url
if request.target.endswith("/.well-known/openid-configuration"):
return Reply(body=json.dumps({"issuer": issuer, "jwks_uri": issuer + "/jwks"}).encode())
return Reply(body=json.dumps({"keys": [jwk]}).encode())
def introspect(request: Request) -> Reply:
return Reply(body=b'{"active": true, "sub": "subject-1"}')
with (
mcp_peer() as peer,
wire_server(idp) as verify_idp,
wire_server(introspect) as introspect_stub,
gateway.scenario() as scenario,
):
holder.append(verify_idp)
alias: Final = "sig" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(scenario, peer, alias)
config: Final = _sign_in_config(
{
"guardrail": "mcp_jwt_signer",
"mode": "pre_mcp_call",
"default_on": True,
"access_token_discovery_uri": verify_idp.url + "/.well-known/openid-configuration",
"token_introspection_endpoint": introspect_stub.url,
"required_claims": ["sub"],
},
tmp_path / "signer.yaml",
)
with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as owned:
key: Final = owned.key(object_permission={"mcp_servers": [identity]})
def call(headers: Mapping[str, str]) -> httpx.Response:
return candidate.client.post(
"/mcp-rest/tools/call",
headers=dict(headers),
json={"name": "add", "arguments": ADD, "server_id": identity},
)
admitted: Final = call({"Authorization": f"Bearer {key}"})
assert admitted.status_code == 200, admitted.text
assert len(tool_calls(peer.drain())) == 1
probes: Final = tuple(item for item in introspect_stub.drain() if item.body)
assert any(key.encode() in probe.body for probe in probes), probes
split: Final = call({"x-litellm-api-key": key})
assert split.status_code == 403, split.text
assert introspect_stub.drain() == ()
AGENT_365_PARAMS: Final = {
"guardrail": "agent_365",
"mode": "pre_mcp_call",
"default_on": True,
"tenant_id": "00000000-0000-0000-0000-000000000000",
"client_id": "22222222-2222-2222-2222-222222222222",
"client_secret": "secret",
}
ENTRA_ISSUER: Final = "https://login.microsoftonline.com/00000000-0000-0000-0000-000000000000/v2.0"
GATEWAY_SCOPE: Final = "api://22222222-2222-2222-2222-222222222222/access_as_user"
def test_agent_365_gated_server_challenges_at_connect_and_advertises_entra(gateway: Gateway, tmp_path: Path) -> None:
config: Final = _sign_in_config(dict(AGENT_365_PARAMS), tmp_path / "agent365.yaml")
with (
owned_proxy(gateway, tmp_path, {}, config=config) as candidate,
mcp_peer() as peer,
candidate.scenario() as scenario,
):
alias: Final = "a365" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(scenario, peer, alias)
granted: Final = scenario.key(object_permission={"mcp_servers": [identity]})
denied: Final = scenario.key(object_permission={"mcp_servers": ["no-mcp-servers"]})
challenged: Final = _rpc(candidate, f"/mcp/{alias}", granted, {})
assert challenged.status_code == 401, challenged.text
authenticate: Final = challenged.headers.get("www-authenticate", "")
assert (
f'resource_metadata="{_origin(candidate)}/.well-known/oauth-protected-resource/mcp/{alias}"' in authenticate
)
assert 'error="invalid_token"' in authenticate
opaque: Final = _rpc(candidate, f"/mcp/{alias}", granted, {"Authorization": "Bearer not-a-jws"})
assert opaque.status_code == 401, opaque.text
assert (
f'resource_metadata="{_origin(candidate)}/.well-known/oauth-protected-resource/mcp/{alias}"'
in opaque.headers.get("www-authenticate", "")
)
discovery: Final = candidate.client.get(f"/.well-known/oauth-protected-resource/mcp/{alias}")
assert discovery.status_code == 200, discovery.text
document: Final = discovery.json()
assert document["authorization_servers"] == [ENTRA_ISSUER]
assert document["scopes_supported"] == [GATEWAY_SCOPE]
refused: Final = _rpc(candidate, f"/mcp/{alias}", denied, {})
assert refused.status_code == 403, refused.text
assert "www-authenticate" not in refused.headers
assert tool_calls(peer.drain()) == ()
def test_agent_365_prm_advertises_the_servers_configured_scopes(gateway: Gateway, tmp_path: Path) -> None:
config: Final = _sign_in_config(dict(AGENT_365_PARAMS), tmp_path / "agent365-scopes.yaml")
with (
owned_proxy(gateway, tmp_path, {}, config=config) as candidate,
mcp_peer() as peer,
candidate.scenario() as scenario,
):
scoped: Final = "a365" + uuid.uuid4().hex[:8]
register_mcp(
scenario,
peer,
scoped,
credentials={"scopes": ["https://example/mcp/scoped/access_as_user", "offline_access"]},
)
unscoped: Final = "a365" + uuid.uuid4().hex[:8]
register_mcp(scenario, peer, unscoped)
assert _advertised(candidate, scoped) == (
200,
(ENTRA_ISSUER,),
["https://example/mcp/scoped/access_as_user", "offline_access"],
)
assert _advertised(candidate, unscoped) == (200, (ENTRA_ISSUER,), [GATEWAY_SCOPE])
def test_challenge_names_the_server_first_route_the_client_connected_on(gateway: Gateway, tmp_path: Path) -> None:
config: Final = _sign_in_config(dict(AGENT_365_PARAMS), tmp_path / "agent365-route.yaml")
with (
owned_proxy(gateway, tmp_path, {}, config=config) as candidate,
mcp_peer() as peer,
candidate.scenario() as scenario,
):
alias: Final = "a365" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(scenario, peer, alias)
granted: Final = scenario.key(object_permission={"mcp_servers": [identity]})
metadata_url: Final = f"{_origin(candidate)}/.well-known/oauth-protected-resource/{alias}/mcp"
challenged: Final = _rpc(candidate, f"/{alias}/mcp", granted, {})
assert challenged.status_code == 401, challenged.text
assert f'resource_metadata="{metadata_url}"' in challenged.headers.get("www-authenticate", "")
document: Final = httpx.get(metadata_url, timeout=15).json()
assert document["resource"] == f"{_origin(candidate)}/{alias}/mcp", document
assert document["authorization_servers"] == [ENTRA_ISSUER]
assert tool_calls(peer.drain()) == ()
def test_challenge_names_the_forwarded_origin_only_from_a_trusted_proxy(gateway: Gateway, tmp_path: Path) -> None:
config: Final = _sign_in_config(
dict(AGENT_365_PARAMS),
tmp_path / "agent365-forwarded.yaml",
{"use_x_forwarded_for": True, "mcp_trusted_proxy_ranges": ["127.0.0.0/8"]},
)
with (
owned_proxy(gateway, tmp_path, {}, config=config) as candidate,
mcp_peer() as peer,
candidate.scenario() as scenario,
):
alias: Final = "a365" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(scenario, peer, alias)
granted: Final = scenario.key(object_permission={"mcp_servers": [identity]})
forwarded: Final = {"X-Forwarded-Proto": "https", "X-Forwarded-Host": "public.example"}
challenged: Final = _rpc(candidate, f"/mcp/{alias}", granted, forwarded)
assert challenged.status_code == 401, challenged.text
assert (
f'resource_metadata="https://public.example/.well-known/oauth-protected-resource/mcp/{alias}"'
in challenged.headers.get("www-authenticate", "")
)
plain: Final = _rpc(candidate, f"/mcp/{alias}", granted, {})
assert plain.status_code == 401, plain.text
assert (
f'resource_metadata="{_origin(candidate)}/.well-known/oauth-protected-resource/mcp/{alias}"'
in plain.headers.get("www-authenticate", "")
)
def test_challenge_and_prm_resolve_the_connected_case_variant(gateway: Gateway, tmp_path: Path) -> None:
config: Final = _sign_in_config(dict(AGENT_365_PARAMS), tmp_path / "agent365-case.yaml")
with (
owned_proxy(gateway, tmp_path, {}, config=config) as candidate,
mcp_peer() as peer,
candidate.scenario() as scenario,
):
alias: Final = "a365" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(scenario, peer, alias)
granted: Final = scenario.key(object_permission={"mcp_servers": [identity]})
connected_as: Final = alias.upper()
challenged: Final = _rpc(candidate, f"/mcp/{connected_as}", granted, {})
assert challenged.status_code == 401, challenged.text
authenticate: Final = challenged.headers.get("www-authenticate", "")
assert (
f'resource_metadata="{_origin(candidate)}/.well-known/oauth-protected-resource/mcp/{connected_as}"'
in authenticate
)
assert 'error="invalid_token"' in authenticate
discovery: Final = candidate.client.get(f"/.well-known/oauth-protected-resource/mcp/{connected_as}")
assert discovery.status_code == 200, discovery.text
document: Final = discovery.json()
assert document["authorization_servers"] == [ENTRA_ISSUER]
assert document["scopes_supported"] == [GATEWAY_SCOPE]
assert tool_calls(peer.drain()) == ()

View file

@ -1,6 +1,8 @@
import base64
import hashlib
import json
import secrets
import socket
import time
import uuid
from dataclasses import dataclass
@ -27,6 +29,7 @@ from integration._support.mcp import (
)
from integration._support.mcp_grants import create_toolset
from integration._support.oauth_server import AuthorizationServer, oauth_server
from integration._support.wire import Reply, Request, wire_server
ADD: Final = {"a": 2, "b": 3}
CLIENT_REDIRECT: Final = "http://127.0.0.1:9/cb"
@ -203,6 +206,105 @@ def test_token_exchange_without_a_subject_token_is_rejected_before_any_upstream_
_assert_subject_token_challenge(as_subject, alias)
@pytest.mark.parametrize("status", (429, 408), ids=("throttled", "timed-out"))
def test_a_throttled_token_exchange_is_an_outage_not_a_sign_in_challenge(gateway: Gateway, status: int) -> None:
def shedding_idp(request: Request) -> Reply:
assert request.method == "POST" and request.target == "/token", request
return Reply(status=status, body=json.dumps({"error": "temporarily_unavailable"}).encode())
with mcp_peer() as peer, wire_server(shedding_idp) as idp, gateway.scenario() as scenario:
alias: Final = "te" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(
scenario,
peer,
alias,
auth_type="oauth2_token_exchange",
token_exchange_endpoint=idp.url + "/token",
credentials={"client_id": "te-client", "client_secret": "te-secret"},
)
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
caller: Final = McpCaller(
gateway, key, "server_mcp", alias, headers={"Authorization": "Bearer subject-" + uuid.uuid4().hex}
)
peer.drain()
response: Final = caller.rpc("tools/call", {"name": f"{alias}-add", "arguments": ADD})
assert response.status_code == 503, (response.status_code, response.text, dict(response.headers))
assert "www-authenticate" not in response.headers, dict(response.headers)
assert len(idp.drain()) == 1
assert tool_calls(peer.drain()) == ()
def test_a_malformed_caller_assertion_is_a_sign_in_challenge_not_an_outage(gateway: Gateway) -> None:
def entra_like_idp(request: Request) -> Reply:
assert request.method == "POST" and request.target == "/token", request
body: Final = {"error": "invalid_client", "error_codes": [5002723], "error_description": "Invalid JWT token"}
return Reply(status=401, body=json.dumps(body).encode())
with mcp_peer() as peer, wire_server(entra_like_idp) as idp, gateway.scenario() as scenario:
alias: Final = "te" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(
scenario,
peer,
alias,
auth_type="oauth2_token_exchange",
token_exchange_endpoint=idp.url + "/token",
credentials={"client_id": "te-client", "client_secret": "te-secret"},
)
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
caller: Final = McpCaller(
gateway,
key,
"server_mcp",
alias,
headers={"Authorization": "Bearer eyJhbGciOiJSUzI1NiJ9.eyJhdWQiOiJ3cm9uZyJ9.c2ln"},
)
peer.drain()
response: Final = caller.rpc("tools/call", {"name": f"{alias}-add", "arguments": ADD})
assert response.status_code == 401, (response.status_code, response.text, dict(response.headers))
challenge: Final = response.headers["www-authenticate"]
origin: Final = str(gateway.client.base_url).rstrip("/")
assert f'resource_metadata="{origin}/.well-known/oauth-protected-resource/{alias}/mcp"' in challenge, challenge
assert 'error="invalid_token"' in challenge, challenge
assert len(idp.drain()) == 1
assert tool_calls(peer.drain()) == ()
def test_a_refused_token_exchange_answers_before_the_request_body_arrives(gateway: Gateway) -> None:
"""The exchange preflight runs on the headers alone, so a client whose body is still in flight gets the
challenge at once instead of the gateway waiting for bytes it will never use."""
def entra_like_idp(request: Request) -> Reply:
body: Final = {"error": "invalid_client", "error_codes": [5002723], "error_description": "Invalid JWT token"}
return Reply(status=401, body=json.dumps(body).encode())
with mcp_peer() as peer, wire_server(entra_like_idp) as idp, gateway.scenario() as scenario:
alias: Final = "te" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(
scenario,
peer,
alias,
auth_type="oauth2_token_exchange",
token_exchange_endpoint=idp.url + "/token",
credentials={"client_id": "te-client", "client_secret": "te-secret"},
)
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
origin: Final = urlsplit(str(gateway.client.base_url))
head: Final = (
f"POST /mcp/{alias} HTTP/1.1\r\nHost: {origin.netloc}\r\nx-litellm-api-key: {key}\r\n"
"Authorization: Bearer eyJhbGciOiJSUzI1NiJ9.eyJhdWQiOiJ3cm9uZyJ9.c2ln\r\n"
"Accept: application/json, text/event-stream\r\nContent-Type: application/json\r\n"
"Content-Length: 4096\r\n\r\n"
)
started: Final = time.monotonic()
with socket.create_connection((origin.hostname or "127.0.0.1", origin.port or 80), timeout=10) as raw:
raw.sendall(head.encode())
status_line: Final = raw.recv(4096).split(b"\r\n", 1)[0]
assert status_line == b"HTTP/1.1 401 Unauthorized", status_line
assert time.monotonic() - started < 5
assert len(idp.drain()) == 1
assert tool_calls(peer.drain()) == ()
def _assert_subject_token_challenge(response: httpx.Response, alias: str) -> None:
assert response.status_code == 401, response.text
challenge: Final = response.headers["www-authenticate"]

View file

@ -615,6 +615,21 @@ def test_raise_token_exchange_challenge_is_rfc9728_invalid_token():
assert "error_description=" in www
def test_raise_token_exchange_challenge_advertises_the_connected_route_metadata_url():
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import (
raise_token_exchange_challenge,
)
connected: Final = "https://gw.example/.well-known/oauth-protected-resource/obo-srv/mcp"
with pytest.raises(HTTPException) as exc_info:
raise_token_exchange_challenge(_server(alias="obo-srv"), root_path="/", resource_metadata=connected)
assert exc_info.value.headers["WWW-Authenticate"] == (
f'Bearer resource_metadata="{connected}", '
'error="invalid_token", '
'error_description="Missing or invalid subject token; authenticate with the IdP and retry"'
)
def test_raise_token_exchange_challenge_includes_server_root_path(monkeypatch):
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import (
raise_token_exchange_challenge,

View file

@ -7,7 +7,9 @@ the I/O edge that maps any transport/HTTP failure to None and parses a JSON body
from unittest.mock import patch
import pytest
from pydantic import SecretStr
from litellm.proxy._experimental.mcp_server.outbound_credentials import Error, Ok, ServerSpec
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchange_provider import (
_post_exchange_endpoint,
build_token_exchanger,
@ -17,24 +19,26 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchanger
SubjectTokenRejected,
TokenExchangeClientError,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import TokenExchangeConfig
_HTTP_CLIENT = "litellm.llms.custom_httpx.http_handler.get_async_httpx_client"
def _client_raising_4xx(body: object):
"""An httpx client whose POST returns a 4xx whose ``raise_for_status`` raises an HTTPStatusError
carrying ``body`` as its JSON, so the RFC 6749 error-code classification can be driven."""
def _client_raising_status(status: int, body: object):
"""An httpx client whose POST returns ``status`` whose ``raise_for_status`` raises an
HTTPStatusError carrying ``body`` as its JSON, so the RFC 6749 error-code classification can be
driven."""
import httpx
request = httpx.Request("POST", "https://idp/token")
response = httpx.Response(400, json=body, request=request)
response = httpx.Response(status, json=body, request=request)
class _Resp:
def raise_for_status(self) -> None:
raise httpx.HTTPStatusError("bad request", request=request, response=response)
class _Client:
async def post(self, url, headers, data):
async def post(self, url, headers, data, timeout=None):
return _Resp()
return _Client()
@ -49,6 +53,25 @@ def test_build_gives_each_caller_an_independent_cache():
assert build_token_exchanger() is not build_token_exchanger()
@pytest.mark.asyncio
async def test_build_token_exchanger_drives_the_injected_http_edge():
seen: list[tuple[str, dict[str, str]]] = []
async def post(url: str, form: dict[str, str], client_auth_headers: dict[str, str]) -> dict[str, object] | None:
seen.append((url, form))
return {"access_token": "x", "expires_in": 60}
config = TokenExchangeConfig(
token_exchange_endpoint="https://idp/token", client_id="cid", client_secret=SecretStr("csec")
)
server = ServerSpec(server_id="srv", resource="https://up.example.com", config=config)
result = await build_token_exchanger(post=post).exchange("jwt", server, config)
assert isinstance(result, Ok)
assert result.ok.access_token == "x"
assert [url for url, _ in seen] == ["https://idp/token"]
assert seen[0][1]["subject_token"] == "jwt"
@pytest.mark.asyncio
async def test_post_returns_none_on_transport_error():
with patch(_HTTP_CLIENT, side_effect=RuntimeError("boom")):
@ -66,7 +89,7 @@ async def test_post_parses_json_body_on_success():
return {"access_token": "x", "expires_in": 60}
class _Client:
async def post(self, url, headers, data):
async def post(self, url, headers, data, timeout=None):
return _Resp()
with patch(_HTTP_CLIENT, return_value=_Client()):
@ -80,7 +103,35 @@ async def test_post_parses_json_body_on_success():
)
async def test_post_maps_gateway_fault_4xx_to_client_error(code):
# RFC 6749 5.2 gateway-fault codes must raise TokenExchangeClientError (-> 500), not the caller 401.
with patch(_HTTP_CLIENT, return_value=_client_raising_4xx({"error": code})):
with patch(_HTTP_CLIENT, return_value=_client_raising_status(400, {"error": code})):
with pytest.raises(TokenExchangeClientError):
await _post_exchange_endpoint("https://idp/token", {"grant_type": "x"}, {})
@pytest.mark.asyncio
@pytest.mark.parametrize("aadsts_code", [5002723, "5002710"], ids=["invalid_jwt", "no_kid_as_string"])
async def test_post_maps_invalid_client_with_an_aadsts_50027xx_code_to_subject_rejected(aadsts_code):
# Entra reports a malformed or unverifiable caller assertion as invalid_client with an AADSTS50027xx
# sub-code (the same top-level error it uses for a bad gateway secret); that one is the caller's 401.
body = {
"error": "invalid_client",
"error_description": f"AADSTS{aadsts_code}: Invalid JWT token.",
"error_codes": [aadsts_code],
}
with patch(_HTTP_CLIENT, return_value=_client_raising_status(401, body)):
with pytest.raises(SubjectTokenRejected):
await _post_exchange_endpoint("https://idp/token", {"grant_type": "x"}, {})
@pytest.mark.asyncio
@pytest.mark.parametrize(
"error_codes",
[[7000215], [5002723.0], "5002723", None],
ids=["bad_secret_code", "float_code", "codes_not_a_list", "no_codes"],
)
async def test_post_keeps_invalid_client_without_an_assertion_code_as_client_error(error_codes):
body = {"error": "invalid_client", **({} if error_codes is None else {"error_codes": error_codes})}
with patch(_HTTP_CLIENT, return_value=_client_raising_status(401, body)):
with pytest.raises(TokenExchangeClientError):
await _post_exchange_endpoint("https://idp/token", {"grant_type": "x"}, {})
@ -93,7 +144,7 @@ async def test_post_maps_gateway_fault_4xx_to_client_error(code):
)
async def test_post_maps_subject_fault_4xx_to_subject_rejected(body):
# A subject-fault code (or an unparseable/absent error) is the caller's problem -> SubjectTokenRejected (401).
with patch(_HTTP_CLIENT, return_value=_client_raising_4xx(body)):
with patch(_HTTP_CLIENT, return_value=_client_raising_status(400, body)):
with pytest.raises(SubjectTokenRejected):
await _post_exchange_endpoint("https://idp/token", {"grant_type": "x"}, {})
@ -110,7 +161,7 @@ async def test_post_returns_none_on_non_object_json(payload):
return payload
class _Client:
async def post(self, url, headers, data):
async def post(self, url, headers, data, timeout=None):
return _Resp()
with patch(_HTTP_CLIENT, return_value=_Client()):
@ -128,7 +179,7 @@ async def test_post_threads_step_up_error_and_claims_into_subject_rejected():
"error_description": "AADSTS50079: the user must enroll MFA",
"claims": claims,
}
with patch(_HTTP_CLIENT, return_value=_client_raising_4xx(body)):
with patch(_HTTP_CLIENT, return_value=_client_raising_status(400, body)):
with pytest.raises(SubjectTokenRejected) as exc_info:
await _post_exchange_endpoint("https://idp/token", {"grant_type": "x"}, {})
assert exc_info.value.claims == claims
@ -137,7 +188,7 @@ async def test_post_threads_step_up_error_and_claims_into_subject_rejected():
@pytest.mark.asyncio
async def test_post_subject_rejection_without_claims_carries_none_claims():
with patch(_HTTP_CLIENT, return_value=_client_raising_4xx({"error": "invalid_grant"})):
with patch(_HTTP_CLIENT, return_value=_client_raising_status(400, {"error": "invalid_grant"})):
with pytest.raises(SubjectTokenRejected) as exc_info:
await _post_exchange_endpoint("https://idp/token", {"grant_type": "x"}, {})
assert exc_info.value.claims is None
@ -148,6 +199,26 @@ async def test_post_gateway_fault_still_wins_when_claims_are_present():
# A gateway-fault code stays a 500-class TokenExchangeClientError even if the body carries
# claims; the caller cannot fix invalid_client by stepping up.
body = {"error": "invalid_client", "claims": '{"access_token":{}}'}
with patch(_HTTP_CLIENT, return_value=_client_raising_4xx(body)):
with patch(_HTTP_CLIENT, return_value=_client_raising_status(400, body)):
with pytest.raises(TokenExchangeClientError):
await _post_exchange_endpoint("https://idp/token", {"grant_type": "x"}, {})
_CONFIG = TokenExchangeConfig(
token_exchange_endpoint="https://idp.example.com/token",
client_id="cid",
client_secret=SecretStr("csec"),
scopes=("s1",),
)
_SERVER = ServerSpec(server_id="srv", resource="https://up.example.com", config=_CONFIG)
@pytest.mark.asyncio
@pytest.mark.parametrize("status", [408, 429])
async def test_exchange_maps_throttled_or_timed_out_4xx_to_upstream_unavailable(status):
# 408/429 are the IdP shedding load, not the caller presenting a bad subject: the exchange must
# surface upstream_unavailable (503-class, retryable) and never tell the caller to sign in again.
with patch(_HTTP_CLIENT, return_value=_client_raising_status(status, {"error": "temporarily_unavailable"})):
result = await OboTokenExchanger(_post_exchange_endpoint).exchange("caller-jwt", _SERVER, _CONFIG)
assert isinstance(result, Error)
assert result.error.tag == "upstream_unavailable"

View file

@ -0,0 +1,179 @@
import asyncio
from collections.abc import Iterator, Mapping
from typing import Final
import pytest
import litellm
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.proxy._experimental.mcp_server.caller_sign_in import (
CallerSignIn,
CallerSignInPreflight,
CallerSignInProvider,
SignedIn,
caller_sign_in_for,
preflight_caller_sign_in,
)
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.mcp import MCPAuth, MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
class _SignInGuardrail(CustomGuardrail):
def __init__(self, issuer: str, scope: str, gated: bool = True) -> None:
super().__init__(guardrail_name=f"sign-in-{issuer}")
self.issuer: Final = issuer
self.scope: Final = scope
self.gated: Final = gated
self.seen: list[tuple[str, str | None]] = [] # mutable-ok: call recorder
def caller_sign_in(self, server: MCPServer, user_api_key_auth: UserAPIKeyAuth | None) -> CallerSignIn | None:
self.seen.append((server.name, user_api_key_auth.user_id if user_api_key_auth else None))
if not self.gated:
return None
return CallerSignIn(issuers=(self.issuer,), scopes=(self.scope,))
async def preflight_caller_sign_in(
self, server: MCPServer, user_api_key_auth: UserAPIKeyAuth | None, subject_token: str
) -> CallerSignInPreflight:
return SignedIn()
class _PlainGuardrail(CustomGuardrail):
def __init__(self) -> None:
super().__init__(guardrail_name="plain")
def _server(auth_type: MCPAuth | None = None, scopes: list[str] | None = None) -> MCPServer:
return MCPServer(
server_id="tools-id",
name="tools",
server_name="tools",
transport=MCPTransport.http,
url="https://tools.test/mcp",
auth_type=auth_type,
scopes=scopes,
)
@pytest.fixture
def registered() -> Iterator[tuple[_SignInGuardrail, _SignInGuardrail]]:
first: Final = _SignInGuardrail("https://idp-a.test", "scope-a")
second: Final = _SignInGuardrail("https://idp-b.test", "scope-b")
plain: Final = _PlainGuardrail()
for callback in (first, plain, second):
litellm.logging_callback_manager.add_litellm_callback(callback)
try:
yield first, second
finally:
for callback in (first, plain, second):
litellm.logging_callback_manager.remove_callback_from_list_by_object(
litellm.callbacks, callback, require_self=False
)
def test_protocol_matches_only_guardrails_implementing_the_hook():
assert isinstance(_SignInGuardrail("i", "s"), CallerSignInProvider)
assert not isinstance(_PlainGuardrail(), CallerSignInProvider)
def test_no_registered_provider_and_non_obo_advertises_nothing():
assert caller_sign_in_for(_server(), None) is None
def test_registered_providers_merge_in_order_and_dedupe(registered):
first, second = registered
sign_in: Final = caller_sign_in_for(_server(), None)
assert sign_in is not None
assert sign_in.issuers == ("https://idp-a.test", "https://idp-b.test")
assert sign_in.scopes == ("scope-a", "scope-b")
assert first.seen == [("tools", None)]
assert second.seen == [("tools", None)]
def test_provider_returning_none_contributes_nothing(registered):
ungated = _SignInGuardrail("https://idp-c.test", "scope-c", gated=False)
litellm.logging_callback_manager.add_litellm_callback(ungated)
try:
sign_in: Final = caller_sign_in_for(_server(), None)
assert sign_in is not None
assert "https://idp-c.test" not in sign_in.issuers
finally:
litellm.logging_callback_manager.remove_callback_from_list_by_object(
litellm.callbacks, ungated, require_self=False
)
def test_obo_server_contributes_jwt_issuers_and_own_scopes(monkeypatch):
monkeypatch.setenv("JWT_ISSUER", "https://jwt-idp.test")
sign_in: Final = caller_sign_in_for(_server(auth_type=MCPAuth.oauth2_token_exchange, scopes=["read"]), None)
assert sign_in == CallerSignIn(issuers=("https://jwt-idp.test",), scopes=("read",))
def test_obo_server_and_provider_merge_and_dedupe(monkeypatch, registered):
monkeypatch.setenv("JWT_ISSUER", "https://idp-a.test")
sign_in: Final = caller_sign_in_for(
_server(auth_type=MCPAuth.oauth2_token_exchange, scopes=["read", "scope-a"]), None
)
assert sign_in is not None
assert sign_in.issuers == ("https://idp-a.test", "https://idp-b.test")
assert sign_in.scopes == ("read", "scope-a", "scope-b")
def test_obo_server_without_jwt_issuer_still_signs_in_when_a_provider_gates(registered):
sign_in: Final = caller_sign_in_for(_server(auth_type=MCPAuth.oauth2_token_exchange), None)
assert sign_in is not None
assert sign_in.issuers == ("https://idp-a.test", "https://idp-b.test")
def test_oauth_utils_strips_the_route_relative_root_path():
"""Regression: Starlette sets ``app_root_path`` to ``""`` on an unmounted app, so the strip must
fall back to ``root_path`` (which is where ``/mcp`` lands when the MCP app is mounted)."""
from litellm.proxy._experimental.mcp_server.oauth_utils import get_route_relative_request_path
scope: Final[Mapping[str, object]] = {
"type": "http",
"path": "/mcp/catalog",
"root_path": "/mcp",
"app_root_path": "",
}
assert get_route_relative_request_path(scope) == "/catalog" # pyright: ignore[reportArgumentType]
@pytest.mark.asyncio
async def test_preflight_with_no_gating_provider_never_reads_the_request_body(monkeypatch):
"""A plain OBO connect has nothing to pre-flight, so the exchange answers before the body arrives, as it
did before the sign-in seam; reading the body first would stall a client that sends its headers early."""
monkeypatch.setenv("JWT_ISSUER", "https://jwt-idp.test")
body_read: Final = asyncio.Event()
async def connecting() -> bool:
body_read.set()
return True
await preflight_caller_sign_in(
_server(auth_type=MCPAuth.oauth2_token_exchange),
None,
"sub.ject.jws",
root_path="",
resource_metadata=None,
connecting=connecting,
)
assert not body_read.is_set()
@pytest.mark.asyncio
async def test_preflight_with_a_gating_provider_reads_the_body_to_tell_a_connect_apart(registered):
body_read: Final = asyncio.Event()
async def connecting() -> bool:
body_read.set()
return True
await preflight_caller_sign_in(
_server(), None, "sub.ject.jws", root_path="", resource_metadata=None, connecting=connecting
)
assert body_read.is_set()

View file

@ -7,6 +7,7 @@ from base64 import urlsafe_b64encode
from datetime import datetime, timedelta, timezone
from typing import TYPE_CHECKING, Final
from unittest.mock import AsyncMock, MagicMock, patch
from urllib.parse import parse_qs, urlparse
import pytest
from fastapi import HTTPException
@ -392,6 +393,84 @@ def trust_xff():
yield
def _registered_gateway_oauth2_server():
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
from litellm.proxy._types import MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
server: Final = MCPServer(
server_id="oid-7f3a",
name="gwx",
server_name="gwx",
alias="gwx",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
oauth2_flow="authorization_code",
client_id="gw-client",
client_secret="gw-secret",
authorization_url="https://provider.com/oauth/authorize",
token_url="https://provider.com/oauth/token",
scopes=["read"],
)
global_mcp_server_manager.registry.clear()
global_mcp_server_manager.registry[server.server_id] = server
return server
@pytest.mark.asyncio
@pytest.mark.parametrize("lookup", ["GWX", "oid-7f3a"], ids=["alias_case", "server_id"])
async def test_authorization_server_doc_for_a_moved_lookup_matches_the_exact_name_doc(lookup):
"""A case variant or server id now resolves like the exact name, so its AS metadata is the
exact-name doc with the requested spelling in the issuer and endpoint paths."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
oauth_authorization_server_mcp_standard,
)
_registered_gateway_oauth2_server()
request: Final = _mock_callback_request("http://litellm.example.com/")
exact: Final = await oauth_authorization_server_mcp_standard(request=request, mcp_server_name="gwx")
moved: Final = await oauth_authorization_server_mcp_standard(request=request, mcp_server_name=lookup)
assert exact["issuer"] == "http://litellm.example.com/mcp/gwx"
assert moved == {
key: value.replace("/gwx", f"/{lookup}") if isinstance(value, str) else value for key, value in exact.items()
}
@pytest.mark.asyncio
@pytest.mark.parametrize("lookup", ["GWX", "oid-7f3a"], ids=["alias_case", "server_id"])
async def test_authorize_relay_for_a_moved_lookup_redirects_like_the_exact_name(lookup):
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import authorize
_registered_gateway_oauth2_server()
request: Final = _mock_callback_request("http://litellm.example.com/")
with patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.encrypt_value_helper", return_value="sealed"
):
exact: Final = await authorize(
request=request, mcp_server_name="gwx", redirect_uri="http://127.0.0.1:60108/callback", state="s1"
)
moved: Final = await authorize(
request=request, mcp_server_name=lookup, redirect_uri="http://127.0.0.1:60108/callback", state="s1"
)
exact_target: Final = urlparse(exact.headers["location"])
moved_target: Final = urlparse(moved.headers["location"])
exact_query: Final = parse_qs(exact_target.query)
moved_query: Final = parse_qs(moved_target.query)
assert exact.status_code == 307
assert exact_target._replace(query="") == urlparse("https://provider.com/oauth/authorize")
assert exact_query["client_id"] == ["gw-client"]
assert len(exact_query.pop("state")) == 1 and len(moved_query.pop("state")) == 1
assert (moved.status_code, moved_target._replace(query=""), moved_query) == (
exact.status_code,
exact_target._replace(query=""),
exact_query,
)
@pytest.mark.asyncio
async def test_authorize_endpoint_includes_response_type():
"""Test that authorize endpoint includes response_type=code parameter (fixes #15684)"""
@ -3642,8 +3721,8 @@ async def test_authorize_resolves_server_by_id_when_name_lookup_fails():
assert response.status_code == 307
assert "https://provider.com/oauth/authorize" in response.headers["location"]
by_name.assert_called_once_with(server.server_id, client_ip=None)
by_id.assert_called_once_with(server.server_id, client_ip=None)
by_name.assert_called_once_with(server.server_id)
by_id.assert_called_once_with(server.server_id)
@pytest.mark.asyncio
@ -3679,8 +3758,8 @@ async def test_token_endpoint_resolves_server_by_id_when_name_lookup_fails():
)
assert json.loads(result.body)["access_token"] == "token"
by_name.assert_called_once_with(server.server_id, client_ip=None)
by_id.assert_called_once_with(server.server_id, client_ip=None)
by_name.assert_called_once_with(server.server_id)
by_id.assert_called_once_with(server.server_id)
@pytest.mark.asyncio
@ -3711,8 +3790,8 @@ async def test_register_client_resolves_server_by_id_when_name_lookup_fails():
result = await discoverable_endpoints.register_client(request=request, mcp_server_name=server.server_id)
assert json.loads(result.body)["client_id"] == "registered-client"
by_name.assert_called_once_with(server.server_id, client_ip=None)
by_id.assert_called_once_with(server.server_id, client_ip=None)
by_name.assert_called_once_with(server.server_id)
by_id.assert_called_once_with(server.server_id)
@pytest.mark.asyncio
@ -3739,8 +3818,58 @@ async def test_protected_resource_metadata_resolves_server_by_id_when_name_looku
assert result["authorization_servers"] == ["https://llm.example.com/mcp"]
assert result["resource"] == f"https://llm.example.com/mcp/{server.server_id}"
by_name.assert_called_once_with(server.server_id, client_ip=None)
by_id.assert_called_once_with(server.server_id, client_ip=None)
by_name.assert_called_once_with(server.server_id)
by_id.assert_called_once_with(server.server_id)
@pytest.mark.asyncio
async def test_protected_resource_metadata_resolves_the_connected_case_variant():
"""The challenge points clients at the segment they connected with (``/mcp/CATALOG``), so the PRM
route must resolve that same segment through the answering-to fallback."""
from fastapi import Request
from litellm.proxy._experimental.mcp_server import discoverable_endpoints
from litellm.proxy._experimental.mcp_server.caller_sign_in import CallerSignIn
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
from litellm.types.mcp import MCPAuth, MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
server = MCPServer(
server_id="catalog-server-id-001",
name="catalog",
alias="catalog",
server_name="catalog",
url="https://catalog.test/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.none,
mcp_info={"server_name": "catalog"},
)
sign_in: Final = CallerSignIn(
issuers=("https://login.microsoftonline.com/00000000-0000-0000-0000-000000000000/v2.0",),
scopes=("api://22222222-2222-2222-2222-222222222222/access_as_user",),
)
request = MagicMock(spec=Request)
request.base_url = "https://llm.example.com/"
request.headers = {}
with (
patch.object(global_mcp_server_manager, "get_mcp_server_by_name", return_value=None), # test-quality-ok: resolver seam
patch.object(global_mcp_server_manager, "get_mcp_server_by_id", return_value=None), # test-quality-ok: resolver seam
patch.object(
global_mcp_server_manager, "get_filtered_registry", return_value={server.server_id: server}
), # test-quality-ok: resolver seam
patch.object(discoverable_endpoints, "caller_sign_in_for", return_value=sign_in), # test-quality-ok: provider seam
):
result = await discoverable_endpoints._build_oauth_protected_resource_response(
request=request,
mcp_server_name="CATALOG",
use_standard_pattern=True,
)
assert result["authorization_servers"] == (
"https://login.microsoftonline.com/00000000-0000-0000-0000-000000000000/v2.0",
)
assert result["resource"] == "https://llm.example.com/mcp/CATALOG"
def test_authorization_server_metadata_resolves_server_by_id_when_name_lookup_fails():
@ -3765,8 +3894,8 @@ def test_authorization_server_metadata_resolves_server_by_id_when_name_lookup_fa
assert result["scopes_supported"] == server.scopes
assert result["issuer"] == f"https://llm.example.com/{server.server_id}"
by_name.assert_called_once_with(server.server_id, client_ip=None)
by_id.assert_called_once_with(server.server_id, client_ip=None)
by_name.assert_called_once_with(server.server_id)
by_id.assert_called_once_with(server.server_id)
@pytest.mark.asyncio
@ -7288,7 +7417,7 @@ async def test_token_exchange_persists_for_oauth2():
# -------------------------------------------------------------------
_OBO_RESOURCE = "https://litellm.example.com/mcp/obo_mcp"
_PATCH_ISSUERS = "litellm.proxy._experimental.mcp_server.discoverable_endpoints._jwt_auth_issuers"
_PATCH_ISSUERS = "litellm.proxy._experimental.mcp_server.caller_sign_in.jwt_auth_issuers"
def _obo_server(scopes=None):
@ -7307,48 +7436,48 @@ def _obo_server(scopes=None):
)
def test_obo_protected_resource_response_names_jwt_issuers():
def test_caller_sign_in_protected_resource_response_names_jwt_issuers():
"""An OBO server's PRM points authorization_servers at the configured JWT issuers (the IdP that
mints and validates the subject token), with the gateway resource echoed back."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
_obo_protected_resource_response,
_caller_sign_in_protected_resource_response,
)
with patch(_PATCH_ISSUERS, return_value=["https://idp.example.com"]):
response = _obo_protected_resource_response(_obo_server(scopes=["read"]), _OBO_RESOURCE)
response = _caller_sign_in_protected_resource_response(_obo_server(scopes=["read"]), _OBO_RESOURCE)
assert response == {
"authorization_servers": ["https://idp.example.com"],
"authorization_servers": ("https://idp.example.com",),
"resource": _OBO_RESOURCE,
"scopes_supported": ["read"],
"scopes_supported": ("read",),
}
def test_obo_protected_resource_response_scopes_default_empty():
"""A scopeless OBO server reports scopes_supported as [] rather than None."""
def test_caller_sign_in_protected_resource_response_scopes_default_empty():
"""A scopeless OBO server reports scopes_supported as an empty array rather than None."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
_obo_protected_resource_response,
_caller_sign_in_protected_resource_response,
)
with patch(_PATCH_ISSUERS, return_value=["https://idp.example.com"]):
response = _obo_protected_resource_response(_obo_server(scopes=None), _OBO_RESOURCE)
assert response["scopes_supported"] == []
response = _caller_sign_in_protected_resource_response(_obo_server(scopes=None), _OBO_RESOURCE)
assert response["scopes_supported"] == ()
def test_obo_protected_resource_response_falls_back_when_no_issuer():
def test_caller_sign_in_protected_resource_response_falls_back_when_no_issuer():
"""With no JWT issuer configured, the OBO branch returns None so the caller falls back to the
gateway-default PRM (discovery still works, it just can't name the IdP)."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
_obo_protected_resource_response,
_caller_sign_in_protected_resource_response,
)
with patch(_PATCH_ISSUERS, return_value=[]):
assert _obo_protected_resource_response(_obo_server(), _OBO_RESOURCE) is None
assert _caller_sign_in_protected_resource_response(_obo_server(), _OBO_RESOURCE) is None
def test_obo_protected_resource_response_ignores_non_obo_server():
def test_caller_sign_in_protected_resource_response_ignores_non_obo_server():
"""Non-OBO servers are not handled by this branch (returns None -> gateway default)."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
_obo_protected_resource_response,
_caller_sign_in_protected_resource_response,
)
from litellm.proxy._types import MCPTransport
from litellm.types.mcp import MCPAuth
@ -7360,7 +7489,7 @@ def test_obo_protected_resource_response_ignores_non_obo_server():
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
)
assert _obo_protected_resource_response(oauth2_server, _OBO_RESOURCE) is None
assert _caller_sign_in_protected_resource_response(oauth2_server, _OBO_RESOURCE) is None
@pytest.mark.asyncio
@ -7390,7 +7519,7 @@ async def test_build_oauth_protected_resource_response_obo_end_to_end():
mcp_server_name="obo_mcp",
use_standard_pattern=True,
)
assert response["authorization_servers"] == ["https://idp.example.com"]
assert response["authorization_servers"] == ("https://idp.example.com",)
assert response["resource"] == "https://litellm.example.com/mcp/obo_mcp"
finally:
global_mcp_server_manager.registry.clear()

View file

@ -924,6 +924,7 @@ async def test_get_tools_from_mcp_servers():
mock_manager.get_mcp_server_by_id = lambda server_id: (
mock_server_1 if server_id == "server1_id" else mock_server_2
)
mock_manager.get_mcp_server_answering_to = MCPServerManager().get_mcp_server_answering_to
mock_manager._get_tools_from_server = AsyncMock(return_value=[mock_tool_1])
# Mock filter_server_ids_by_ip_with_info to return input unchanged (no IP filtering in test)
mock_manager.filter_server_ids_by_ip_with_info = MagicMock(
@ -1004,6 +1005,7 @@ async def test_get_tools_from_mcp_servers():
if server_id == "server1_id"
else (mock_server_2 if server_id == "server2_id" else mock_server_3)
)
mock_manager.get_mcp_server_answering_to = MagicMock(return_value=None)
mock_manager._get_tools_from_server = AsyncMock(return_value=[mock_tool_1])
# Mock filter_server_ids_by_ip_with_info to return input unchanged (no IP filtering in test)
mock_manager.filter_server_ids_by_ip_with_info = MagicMock(

View file

@ -1637,6 +1637,19 @@ class TestMCPServerManager:
assert server.uses_per_server_oauth_relay is True
assert server.advertises_gateway_authorization_server is False
@pytest.mark.asyncio
async def test_load_servers_from_config_does_not_advertise_gateway_as_for_token_exchange(self):
# keeps_caller_authorization includes oauth2_token_exchange so a sign-in provider may gate it,
# but named discovery must still fall to the server's own PRM rather than the aggregate AS.
manager = MCPServerManager()
with patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=None)):
await manager.load_servers_from_config(self._client_forwarded_config(MCPAuth.oauth2_token_exchange))
server = next(iter(manager.config_mcp_servers.values()))
assert server.keeps_caller_authorization is True
assert server.advertises_gateway_authorization_server is False
@pytest.mark.asyncio
@pytest.mark.parametrize(
"config",
@ -3090,6 +3103,34 @@ class TestMCPServerManager:
www_authenticate = headers.get("WWW-Authenticate") or headers.get("www-authenticate") or ""
assert "resource_metadata" in www_authenticate
@pytest.mark.asyncio
async def test_preflight_rejected_subject_challenge_names_the_connected_segment(self):
"""A subject rejected on ``/mcp/<server_id>`` must point resource_metadata at that same
segment, the way the sign-in preflight does, so the client's discovery fetch resolves."""
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Error
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import CredError
class _FakeProvider:
async def resolve_credentials(self, subject, server):
return Error(CredError.of_unauthorized("subject token rejected by the IdP"))
manager = MCPServerManager(cred_provider=_FakeProvider())
server = self._token_exchange_server("te-preflight-segment")
with pytest.raises(HTTPException) as exc_info:
await manager.preflight_token_exchange(
server=server,
oauth2_headers={"Authorization": "Bearer rejected-subject"},
user_api_key_auth=None,
resource_metadata=f"http://gw.test/.well-known/oauth-protected-resource/mcp/{server.server_id}",
)
headers = exc_info.value.headers or {}
www_authenticate = headers.get("WWW-Authenticate") or headers.get("www-authenticate") or ""
assert (
f'resource_metadata="http://gw.test/.well-known/oauth-protected-resource/mcp/{server.server_id}"'
in www_authenticate
), www_authenticate
@pytest.mark.asyncio
async def test_preflight_token_exchange_maps_gateway_fault_to_public_status(self):
"""A gateway-fault CredError (e.g. invalid_client) must surface its public status (500)
@ -6346,6 +6387,94 @@ class TestMCPServerManager:
assert manager._resolve_mcp_server_for_tool_call("zapier", "create_zap") is zapier
assert manager._resolve_mcp_server_for_tool_call("other", "create_zap") is other
@pytest.mark.parametrize("gh_first", [True, False], ids=["gh-listed-first", "gh-public-listed-first"])
def test_answering_to_prefers_alias_over_earlier_prefix_match(self, gh_first):
manager = MCPServerManager()
gh = MCPServer(
server_id="gh-id", name="gh", server_name="gh", transport=MCPTransport.http, auth_type=MCPAuth.oauth2
)
gh_public = MCPServer(
server_id="gh-public-id", name="gh_public", server_name="gh_public", alias="gh", transport=MCPTransport.http
)
manager.registry = (
{"gh-id": gh, "gh-public-id": gh_public} if gh_first else {"gh-public-id": gh_public, "gh-id": gh}
)
assert manager.get_mcp_server_answering_to("gh") is gh_public
assert manager.get_mcp_server_answering_to("GH") is gh_public
assert manager.get_mcp_server_answering_to("Gh_Public") is gh_public
assert manager.get_mcp_server_answering_to("gh-public-id") is gh_public
assert manager.get_mcp_server_answering_to("GH_PUBLIC") is gh_public
assert manager.get_mcp_server_answering_to("gh-id") is gh
@pytest.mark.parametrize("gh_first", [True, False], ids=["alias-listed-first", "server-name-listed-first"])
def test_answering_to_agrees_with_exact_name_before_case_folding(self, gh_first):
manager = MCPServerManager()
by_alias = MCPServer(
server_id="a-id",
name="a",
server_name="a",
alias="gh",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
)
by_server_name = MCPServer(server_id="b-id", name="b", server_name="Gh", transport=MCPTransport.http)
manager.registry = (
{"a-id": by_alias, "b-id": by_server_name} if gh_first else {"b-id": by_server_name, "a-id": by_alias}
)
for name in ("Gh", "gh"):
assert manager.get_mcp_server_answering_to(name) is manager.get_mcp_server_by_name(name), name
assert manager.get_mcp_server_answering_to("Gh") is by_server_name
assert manager.get_mcp_server_answering_to("gh") is by_alias
assert manager.get_mcp_server_answering_to("GH") is by_alias
@pytest.mark.parametrize("hidden_first", [True, False], ids=["hidden-listed-first", "public-listed-first"])
def test_answering_to_never_reroutes_a_name_hidden_from_an_ip_to_a_case_variant(self, hidden_first):
manager = MCPServerManager()
hidden = MCPServer(
server_id="p-id",
name="gh",
server_name="gh",
transport=MCPTransport.http,
available_on_public_internet=False,
)
public = MCPServer(server_id="u-id", name="u", server_name="u", alias="Gh", transport=MCPTransport.http)
manager.registry = {"p-id": hidden, "u-id": public} if hidden_first else {"u-id": public, "p-id": hidden}
assert manager.get_mcp_server_answering_to("gh", client_ip="203.0.113.7") is None
assert manager.get_mcp_server_answering_to("p-id", client_ip="203.0.113.7") is None
assert manager.get_mcp_server_answering_to("gh", client_ip="10.0.0.7") is hidden
assert manager.get_mcp_server_answering_to("Gh", client_ip="203.0.113.7") is public
@pytest.mark.parametrize("pinned_first", [True, False], ids=["pinned-id-listed-first", "alias-listed-first"])
def test_answering_to_and_discovery_agree_on_a_pinned_id_that_another_alias_case_folds_to(self, pinned_first):
from litellm.proxy._experimental.mcp_server import discoverable_endpoints
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
pinned = MCPServer(server_id="foo", name="pinned", server_name="pinned", transport=MCPTransport.http)
by_alias = MCPServer(
server_id="b-id",
name="b",
server_name="b",
alias="Foo",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
)
registry = {"foo": pinned, "b-id": by_alias} if pinned_first else {"b-id": by_alias, "foo": pinned}
global_mcp_server_manager.registry.clear()
global_mcp_server_manager.registry.update(registry)
try:
for name in ("foo", "Foo", "FOO"):
connected = global_mcp_server_manager.get_mcp_server_answering_to(name)
discovered = discoverable_endpoints._resolve_mcp_server_by_name_or_id(name, client_ip=None)
assert discovered is connected, name
assert global_mcp_server_manager.get_mcp_server_answering_to("foo") is pinned
assert global_mcp_server_manager.get_mcp_server_answering_to("Foo") is by_alias
assert global_mcp_server_manager.get_mcp_server_answering_to("FOO") is by_alias
finally:
global_mcp_server_manager.registry.clear()
def test_remove_server_drops_only_its_own_tool_mapping_rows(self):
manager = self._manager_with_deepwiki_and_huggingface()
@ -6973,6 +7102,185 @@ class TestMCPServerManager:
assert exc_info.value.status_code == 403
@pytest.mark.asyncio
@pytest.mark.parametrize(
("raw_headers", "api_key", "expected_bearer", "expected_subject"),
[
pytest.param(
{"authorization": "Bearer sk-1234"},
"sk-1234",
"sk-1234",
None,
id="litellm-key-as-bearer-is-not-a-subject",
),
pytest.param(
{"x-litellm-api-key": "sk-1234", "authorization": "Bearer eyJ.x.y"},
"sk-1234",
"eyJ.x.y",
"eyJ.x.y",
id="key-admission-plus-idp-bearer-subject",
),
pytest.param(
{"x-litellm-api-key": "sk-1234", "authorization": "bearer eyJ.x.y"},
"sk-1234",
"eyJ.x.y",
"eyJ.x.y",
id="lowercase-bearer-scheme-is-the-subject",
),
pytest.param(
{"x-litellm-api-key": "sk-1234", "authorization": "Basic a.b.c"},
"sk-1234",
None,
None,
id="non-bearer-scheme-is-not-a-subject",
),
pytest.param(
{"x-litellm-api-key": "sk-1234", "authorization": "Digest x.y.z"},
"sk-1234",
None,
None,
id="digest-scheme-is-not-a-subject",
),
pytest.param(
{"x-litellm-api-key": "sk-1234", "authorization": "eyJ.x.y"},
"sk-1234",
None,
None,
id="scheme-less-value-is-not-a-subject",
),
],
)
async def test_pre_call_tool_check_separates_raw_bearer_from_subject(
self, raw_headers, api_key, expected_bearer, expected_subject
):
manager = MCPServerManager()
server = MCPServer(
server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv", allowed_tools=None
)
proxy_logging = MagicMock()
proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value={})
proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={})
proxy_logging.pre_call_hook = AsyncMock(return_value=None)
await manager.pre_call_tool_check(
server_name="srv",
name="turn",
arguments={},
user_api_key_auth=UserAPIKeyAuth(api_key=api_key, user_id="u"),
proxy_logging_obj=proxy_logging,
server=server,
raw_headers=raw_headers,
)
kwargs: Final = proxy_logging._create_mcp_request_object_from_kwargs.call_args.args[0]
assert kwargs["incoming_bearer_token"] == expected_bearer
assert kwargs["incoming_subject_token"] == expected_subject
@pytest.mark.asyncio
@pytest.mark.parametrize(
("raw_headers", "expected_subject"),
[
pytest.param({"authorization": "Bearer gw.master.key"}, None, id="master-key-alone-is-not-a-subject"),
pytest.param(
{"x-litellm-api-key": "sk-1234", "authorization": "Bearer gw.master.key"},
None,
id="master-key-next-to-a-virtual-key-is-not-a-subject",
),
pytest.param(
{"x-litellm-api-key": "sk-1234", "authorization": "Bearer gw.other.jws"},
"gw.other.jws",
id="a-dotted-bearer-that-is-not-the-master-key-is-the-subject",
),
],
)
async def test_pre_call_tool_check_withholds_a_dotted_master_key_from_the_subject(
self, raw_headers, expected_subject
):
"""A master key is a LiteLLM credential whatever its shape, so even one with the two dots of a
compact JWS never becomes the sign-in subject a provider would send to its IdP."""
manager = MCPServerManager()
server = MCPServer(
server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv", allowed_tools=None
)
proxy_logging = MagicMock()
proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value={})
proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={})
proxy_logging.pre_call_hook = AsyncMock(return_value=None)
with patch("litellm.proxy.proxy_server.master_key", "gw.master.key"):
await manager.pre_call_tool_check(
server_name="srv",
name="turn",
arguments={},
user_api_key_auth=UserAPIKeyAuth(api_key="sk-1234", user_id="u"),
proxy_logging_obj=proxy_logging,
server=server,
raw_headers=raw_headers,
)
kwargs: Final = proxy_logging._create_mcp_request_object_from_kwargs.call_args.args[0]
assert kwargs["incoming_bearer_token"] == raw_headers["authorization"].removeprefix("Bearer ")
assert kwargs["incoming_subject_token"] == expected_subject
@pytest.mark.asyncio
@pytest.mark.parametrize(
("raw_headers", "api_key", "custom_auth", "expected_subject"),
[
pytest.param(
{"authorization": "Bearer eyJ.x.y"},
"eyJ.x.y",
True,
"eyJ.x.y",
id="custom-auth-idp-bearer-is-the-subject",
),
pytest.param(
{"authorization": "Bearer eyJ.x.y"},
"eyJ.x.y",
False,
"eyJ.x.y",
id="built-in-oauth2-admission-bearer-is-the-subject",
),
pytest.param({"authorization": "Bearer sk-1234"}, "sk-1234", True, None, id="virtual-key-is-not-a-subject"),
pytest.param(
{"x-litellm-api-key": "ca-key", "authorization": "Bearer ca-key"},
"ca-key",
True,
None,
id="explicit-key-admission-repeated-in-authorization-is-not-a-subject",
),
],
)
async def test_pre_call_tool_check_hands_sign_in_the_bearer_that_admitted_the_caller(
self, raw_headers, api_key, custom_auth, expected_subject
):
"""Custom auth and the built-in OAuth2 admission both admit the caller on its own IdP token in
``Authorization`` with no ``x-litellm-api-key`` and record it as ``api_key``; that token is the
sign-in subject as the raw bearer was before the subject split."""
manager = MCPServerManager()
server = MCPServer(
server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv", allowed_tools=None
)
admitted = UserAPIKeyAuth(api_key=api_key, user_id="u")
admitted.authenticated_by_custom_auth = custom_auth
proxy_logging = MagicMock()
proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value={})
proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={})
proxy_logging.pre_call_hook = AsyncMock(return_value=None)
await manager.pre_call_tool_check(
server_name="srv",
name="turn",
arguments={},
user_api_key_auth=admitted,
proxy_logging_obj=proxy_logging,
server=server,
raw_headers=raw_headers,
)
kwargs: Final = proxy_logging._create_mcp_request_object_from_kwargs.call_args.args[0]
assert kwargs["incoming_bearer_token"] == api_key
assert kwargs["incoming_subject_token"] == expected_subject
@pytest.mark.asyncio
async def test_check_tool_permission_for_key_team_allows_permitted_tool(self):
"""

View file

@ -650,6 +650,10 @@ async def test_per_user_oauth_missing_stored_token_returns_preemptive_401():
"litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_by_name",
return_value=oauth_server,
),
patch(
"litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_answering_to",
return_value=oauth_server,
),
patch.object(
session_manager_stateless,
"handle_request",
@ -738,6 +742,10 @@ async def test_admitted_subject_missing_stored_token_challenged_with_resource_me
"litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_by_name",
return_value=oauth_server,
),
patch(
"litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_answering_to",
return_value=oauth_server,
),
patch.object(
session_manager_stateless,
"handle_request",
@ -1033,6 +1041,10 @@ async def test_per_user_oauth_with_stored_token_skips_preemptive_401():
"litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_by_name",
return_value=oauth_server,
),
patch(
"litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_answering_to",
return_value=oauth_server,
),
patch.object(
session_manager_stateless,
"handle_request",
@ -1136,6 +1148,10 @@ async def test_handle_streamable_http_mcp_delegated_server_without_token_returns
"litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_by_name",
return_value=delegated_server,
),
patch(
"litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_answering_to",
return_value=delegated_server,
),
patch.object(
session_manager_stateful,
"handle_request",
@ -1224,6 +1240,10 @@ async def test_handle_streamable_http_mcp_token_exchange_without_subject_returns
"litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_by_name",
return_value=obo_server,
),
patch(
"litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_answering_to",
return_value=obo_server,
),
patch.object(
session_manager_stateful,
"handle_request",
@ -1323,6 +1343,10 @@ async def test_handle_streamable_http_mcp_oauth_delegate_without_token_returns_g
"litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_by_name",
return_value=od_server,
),
patch(
"litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_answering_to",
return_value=od_server,
),
patch.object(
session_manager_stateful,
"handle_request",
@ -1580,6 +1604,10 @@ async def test_handle_streamable_http_mcp_true_passthrough_without_token_surface
"litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_by_name",
return_value=tp_server,
),
patch(
"litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_answering_to",
return_value=tp_server,
),
patch.object(
session_manager_stateful,
"handle_request",
@ -1648,6 +1676,10 @@ async def test_handle_streamable_http_mcp_true_passthrough_dcr_bridge_challenges
"litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_by_name",
return_value=bridge_server,
),
patch(
"litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_answering_to",
return_value=bridge_server,
),
patch.object(
session_manager_stateful,
"handle_request",

File diff suppressed because it is too large Load diff

View file

@ -39,6 +39,7 @@ def test_convert_mcp_to_llm_format_returns_synthetic_data(proxy_logging, make_mc
"user_api_key_hash": "hash",
"user_api_key_request_route": "/mcp",
"incoming_bearer_token": "tok",
"incoming_subject_token": "a.b.c",
},
)
snapshot = {
@ -47,6 +48,7 @@ def test_convert_mcp_to_llm_format_returns_synthetic_data(proxy_logging, make_mc
"mcp_tool_name": out["mcp_tool_name"],
"mcp_arguments": out["mcp_arguments"],
"incoming_bearer_token": out["incoming_bearer_token"],
"incoming_subject_token": out["incoming_subject_token"],
"message_role": out["messages"][0]["role"],
}
assert snapshot == {
@ -55,6 +57,7 @@ def test_convert_mcp_to_llm_format_returns_synthetic_data(proxy_logging, make_mc
"mcp_tool_name": "search",
"mcp_arguments": {"q": "hello"},
"incoming_bearer_token": "tok",
"incoming_subject_token": "a.b.c",
"message_role": "user",
}
@ -66,12 +69,14 @@ def test_convert_mcp_to_llm_format_defaults_model(proxy_logging, make_mcp_reques
"model": out["model"],
"mcp_tool_name": out["mcp_tool_name"],
"incoming_bearer_token": out["incoming_bearer_token"],
"incoming_subject_token": out["incoming_subject_token"],
"user_id": out["user_api_key_user_id"],
}
assert snapshot == {
"model": "mcp-tool-call",
"mcp_tool_name": "calculator",
"incoming_bearer_token": None,
"incoming_subject_token": None,
"user_id": None,
}