feat(mcp): fold gateway sign-in into a caller sign-in contract on the challenge path

Replace the parallel gateway sign-in provider registry with
CallerSignInProvider, merged into the existing oauth2_token_exchange
challenge: gates key off caller_sign_in_for() returning non-None, the
resolver answers by server_id, case-insensitive name, and short prefix
like the router, keeps_caller_authorization covers OBO servers, and the
Agent 365 guardrail exchanges through an injected TokenExchanger.

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-09-26 09:32:46 +00:00
parent 9c68ca4412
commit 4cadb402e5
15 changed files with 1005 additions and 477 deletions

View file

@ -0,0 +1,124 @@
"""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
from collections.abc import Mapping
from dataclasses import dataclass
from typing import TYPE_CHECKING, Final, Protocol, cast, runtime_checkable
from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError
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, ...]
@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)."""
...
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]] = cast(Mapping[str, object], general_settings)
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 = [
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(issuer for contribution in contributions for issuer in contribution.issuers))
scopes: Final = tuple(dict.fromkeys(scope for contribution in contributions for scope in contribution.scopes))
return CallerSignIn(issuers=issuers, scopes=scopes)

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,
@ -2580,9 +2581,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 +2604,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": list(sign_in.issuers),
"resource": resource_url,
"scopes_supported": (mcp_server.scopes if mcp_server.scopes else []),
"scopes_supported": list(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

@ -168,6 +168,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,
)
@ -4241,6 +4242,12 @@ 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,
)
if subject_token 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())
return
resolved_server: Final = await self.ensure_oauth_metadata_discovered(server)
spec: Final = _to_server_spec_fail_closed(resolved_server)
@ -5964,13 +5971,8 @@ 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 ") :]
# Admission credentials are never handed to guardrails as the caller's assertion.
incoming_bearer_token: Final = self._extract_subject_token(None, raw_headers, user_api_key_auth)
pre_hook_kwargs: Final = {
"guardrail_context": guardrail_context,
@ -7189,6 +7191,18 @@ class MCPServerManager:
return server
return None
def get_mcp_server_answering_to(self, name: str, client_ip: str | None = None) -> MCPServer | None:
"""The server a scoped ``/mcp/{name}`` connect resolves to, matched the way the router matches
it: case-insensitive over server_id, name and every published prefix form."""
return next(
(
server
for server in self.get_filtered_registry(client_ip).values()
if 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,
)
@ -496,8 +497,7 @@ 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)
async def raise_denied_scoped_mcp_access(
@ -1768,7 +1768,11 @@ def _challenge_missing_token_exchange_subject(
warm path already raises. Gated to servers the key may reach so an unauthorized caller
learns nothing about the catalog.
"""
if server is None or server.auth_type != MCPAuth.oauth2_token_exchange:
from litellm.proxy._experimental.mcp_server.caller_sign_in import (
caller_sign_in_for, # noqa: PLC0415 # lazy: caller_sign_in pulls the proxy graph
)
if server is None or caller_sign_in_for(server, user_api_key_auth) is None:
return
if requested_server is not None and requested_server.server_id != server.server_id:
return

View file

@ -79,6 +79,11 @@ async def _post_exchange_endpoint(
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:

View file

@ -36,6 +36,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,
@ -1580,6 +1581,21 @@ if MCP_AVAILABLE:
)
return user_api_key_auth.model_copy(update={"object_permission": updated_op, "mcp_toolset_id": toolset_id})
async def _key_granted_single_server(
server: MCPServer,
mcp_servers: Sequence[str] | None,
user_api_key_auth: UserAPIKeyAuth | None,
client_ip: str | None,
) -> bool:
"""Sign-in challenges are issued 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."""
if len(mcp_servers or []) != 1:
return False
allowed: Final = await operations._get_allowed_mcp_servers(
user_api_key_auth=user_api_key_auth, mcp_servers=mcp_servers, client_ip=client_ip
)
return any(granted.server_id == server.server_id for granted in allowed)
async def _raise_preemptive_401_for_unauthenticated_servers(
scope: Scope,
mcp_servers: list[str] | None,
@ -1602,7 +1618,7 @@ 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)
server = operations.global_mcp_server_manager.get_mcp_server_answering_to(server_name, client_ip=client_ip)
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
@ -1698,12 +1714,22 @@ 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. Non-OBO gates fire only on a
# single-server connect the key's grant admits.
sign_in = caller_sign_in_for(server, user_api_key_auth) if server is not None else None
if (
server
and sign_in is not None
and operations.global_mcp_server_manager._extract_subject_token( # pyright: ignore[reportPrivateUsage] # the manager owns the subject/admission filter shared with the preflight
oauth2_headers, raw_headers, user_api_key_auth
)
is None
and (
server.auth_type == MCPAuth.oauth2_token_exchange
or await _key_granted_single_server(server, mcp_servers, user_api_key_auth, client_ip)
)
):
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import ( # noqa: PLC0415 # lazy: adapter pulls MCP subgraph
raise_token_exchange_challenge,
)

View file

@ -370,6 +370,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

@ -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,18 @@ from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
)
from litellm.proxy._experimental.mcp_server.caller_sign_in import CallerSignIn
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Error, Ok
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchange_provider import (
build_token_exchanger,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchanger import TokenExchanger
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
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,38 +57,14 @@ 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 ()
def _parse_tool_input_schema(raw: object) -> Mapping[str, object] | None:
@ -129,27 +113,6 @@ class _BlockedDetail(TypedDict):
correlation_id: ReadOnly[str | None]
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
@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)
class Agent365MalformedResponseError(Exception):
pass
class Agent365ThrottledError(Exception):
def __init__(self, status_code: int) -> None:
super().__init__(f"HTTP {status_code}")
@ -173,6 +136,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 +156,19 @@ 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()
verbose_proxy_logger.info("Initialized Microsoft Agent 365 guardrail: %s", guardrail_name)
@staticmethod
@ -233,42 +208,40 @@ class Agent365Guardrail(CustomGuardrail):
)
try:
obo_token: Final = await self._get_obo_token(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})",
)
except Agent365ThrottledError as exc:
self._handle_throttled(
data=data,
tool_name=tool_name,
reason=f"the Entra token endpoint returned HTTP {exc.status_code}",
latency_ms=None,
)
exchange_result: Final = await self._exchange_caller_assertion(assertion)
except (httpx.HTTPError, LitellmTimeout, TimeoutError) as exc:
return self._handle_unavailable(
data=data,
tool_name=tool_name,
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),
)
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 "misconfigured":
return self._handle_unavailable(
data=data,
tool_name=tool_name,
reason=(
f"Entra rejected the gateway's own Agent 365 credentials ({error.misconfigured}); "
"check the guardrail's client_id and client_secret"
),
)
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 +257,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 +281,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,62 +430,29 @@ 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]
response: Final = await self._post_allowing_error_status(
url=TOKEN_ENDPOINT_TEMPLATE.format(tenant_id=self.tenant_id),
data={ # mutable-ok: OAuth form body; AsyncHTTPHandler.post requires dict
"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"}, # mutable-ok: httpx header dict
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 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
if user_api_key_auth is not None:
probe: Final[dict[str, Mapping[str, object]]] = { # pyright: ignore[reportUnknownVariableType] # UserAPIKeyAuth metadata dicts are untyped
"metadata": {
"user_api_key_metadata": user_api_key_auth.metadata, # pyright: ignore[reportUnknownMemberType] # UserAPIKeyAuth.metadata is a raw dict
"user_api_key_team_metadata": user_api_key_auth.team_metadata, # 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=(GATEWAY_SCOPE_TEMPLATE.format(client_id=self.client_id),),
)
if response.status_code in (408, 429):
raise Agent365ThrottledError(status_code=response.status_code)
if response.status_code >= 500:
raise httpx.HTTPStatusError(
f"Entra token endpoint returned {response.status_code}",
request=response.request,
response=response,
)
try:
parsed_body: Final = 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
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")),
)
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:
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
async def _post_allowing_error_status(
self,
@ -577,11 +519,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

@ -313,10 +313,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,
@ -326,11 +326,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

@ -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, ServerSpec
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchange_provider import (
_post_exchange_endpoint,
build_token_exchanger,
@ -17,17 +19,19 @@ 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:
@ -80,7 +84,7 @@ 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"}, {})
@ -93,7 +97,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"}, {})
@ -128,7 +132,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 +141,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 +152,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,134 @@
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,
CallerSignInProvider,
caller_sign_in_for,
)
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,))
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]

View file

@ -7288,7 +7288,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,15 +7307,15 @@ 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"],
"resource": _OBO_RESOURCE,
@ -7323,32 +7323,32 @@ def test_obo_protected_resource_response_names_jwt_issuers():
}
def test_obo_protected_resource_response_scopes_default_empty():
def test_caller_sign_in_protected_resource_response_scopes_default_empty():
"""A scopeless OBO server reports scopes_supported as [] 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)
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 +7360,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

View file

@ -1,4 +1,3 @@
from litellm.proxy._experimental.mcp_server import operations as mcp_operations
import asyncio
import contextlib
import contextvars
@ -29,6 +28,9 @@ from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS, LATEST_HANDSHAKE_VERS
from pydantic import TypeAdapter
from starlette.types import Message, Receive, Scope, Send
import litellm
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.proxy._experimental.mcp_server import operations as mcp_operations
from litellm.proxy._experimental.mcp_server.mcp_context import active_mcp_request_ctx_var
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ListedToolsCaller
from litellm.proxy._types import (
@ -1159,6 +1161,7 @@ async def test_mcp_read_resource_success():
)
async def test_read_resource_preserves_content_metadata(_mcp_request_ctx, kind, metadata):
from mcp.types import ReadResourceRequestParams
from litellm.proxy._experimental.mcp_server import operations, server
uri: Final = "https://example.com/resource"
@ -9881,7 +9884,7 @@ class TestPreemptive401ModeAware:
with (
patch.object(
mcp_operations.global_mcp_server_manager,
"get_mcp_server_by_name",
"get_mcp_server_answering_to",
return_value=server,
),
patch.object(
@ -9980,7 +9983,7 @@ class TestPreemptive401ModeAware:
patch.dict(os.environ, {"SERVER_ROOT_PATH": "/litellm"}),
patch.object(
mcp_operations.global_mcp_server_manager,
"get_mcp_server_by_name",
"get_mcp_server_answering_to",
return_value=server,
),
patch.object(
@ -10087,7 +10090,7 @@ class TestSingleServerPreflightReachesIdJag:
with (
patch.object( # test-quality-ok: route wiring must use the manager's configured server
mcp_operations.global_mcp_server_manager,
"get_mcp_server_by_name",
"get_mcp_server_answering_to",
return_value=server,
),
patch.object( # test-quality-ok: route wiring must invoke the manager preflight
@ -10147,7 +10150,7 @@ class TestSingleServerPreflightReachesIdJag:
with (
patch.object( # test-quality-ok: route wiring must use the manager's configured server
mcp_operations.global_mcp_server_manager,
"get_mcp_server_by_name",
"get_mcp_server_answering_to",
return_value=token_exchange,
),
patch.object( # test-quality-ok: route wiring must invoke the manager preflight
@ -10212,7 +10215,7 @@ class TestOboPreflightScopedToAllowedServers:
preflight = AsyncMock()
with (
patch.object( # test-quality-ok: route handler reads the module-level manager, no injection seam
mcp_operations.global_mcp_server_manager, "get_mcp_server_by_name", return_value=requested
mcp_operations.global_mcp_server_manager, "get_mcp_server_answering_to", return_value=requested
),
patch.object( # test-quality-ok: the exchanger is the observable; a real one would call an IdP
mcp_operations.global_mcp_server_manager, "preflight_token_exchange", preflight
@ -10228,6 +10231,10 @@ class TestOboPreflightScopedToAllowedServers:
mcp_server_auth_headers=None,
user_api_key_auth=user_api_key_auth,
client_ip="10.0.0.7",
raw_headers={
"x-litellm-api-key": user_api_key_auth.api_key if user_api_key_auth else "",
"authorization": self.SUBJECT_HEADERS["Authorization"],
},
)
return allowed_lookup, preflight
@ -10253,7 +10260,13 @@ class TestOboPreflightScopedToAllowedServers:
_, preflight = await self._run(requested, allowed=[requested], user_api_key_auth=key)
preflight.assert_awaited_once_with(
server=requested, oauth2_headers=self.SUBJECT_HEADERS, user_api_key_auth=key, raw_headers=None
server=requested,
oauth2_headers=self.SUBJECT_HEADERS,
user_api_key_auth=key,
raw_headers={
"x-litellm-api-key": key.api_key,
"authorization": self.SUBJECT_HEADERS["Authorization"],
},
)
@ -10774,8 +10787,8 @@ async def test_native_listing_preserves_empty_result_on_auth_failure(_mcp_reques
@pytest.mark.asyncio
@pytest.mark.parametrize("failure_hook_raises", [False, True])
async def test_tool_listing_preserves_permission_denial_when_failure_logging_fails(failure_hook_raises):
from litellm.proxy._experimental.mcp_server import operations
from litellm.proxy import proxy_server
from litellm.proxy._experimental.mcp_server import operations
auth = UserAPIKeyAuth(user_id="denied-caller")
denial = HTTPException(status_code=403, detail="scope denied")
@ -10811,8 +10824,9 @@ async def test_legacy_sse_mount_emits_message_endpoint(
) -> None:
from starlette.applications import Starlette
from starlette.routing import Mount
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import AggregateToolListing
from litellm.proxy._experimental.mcp_server import server as mcp_server
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import AggregateToolListing
app: Final = Starlette(routes=[Mount("/mcp", app=mcp_server.app)])
incoming: Final[asyncio.Queue[Message]] = asyncio.Queue()
@ -10970,6 +10984,7 @@ def test_protocol_header_respects_configured_advertisement(revision, rejected):
@pytest.mark.asyncio
async def test_discovery_adapter_preserves_authenticated_context(_mcp_request_ctx):
from mcp.types import DiscoverResult, RequestParams, ServerCapabilities
from litellm.proxy._experimental.mcp_server import server
expected = DiscoverResult(supported_versions=["2025-11-25"], capabilities=ServerCapabilities())
@ -10988,3 +11003,151 @@ async def test_discovery_adapter_preserves_authenticated_context(_mcp_request_ct
context = dispatched.await_args.args[1]
assert context.user_api_key_auth.user_id == "discover-caller"
assert context.mcp_servers == ("allowed",)
def _catalog_server() -> MCPServer:
return 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"},
)
class _CallerSignInGuardrail(CustomGuardrail):
def caller_sign_in(self, server, user_api_key_auth):
from litellm.proxy._experimental.mcp_server.caller_sign_in import CallerSignIn
return CallerSignIn(issuers=("https://idp.test",), scopes=("scope-a",))
class TestConnectChallengeResolver:
"""The connect-time sign-in challenge must resolve the server the same way the router resolves
``/mcp/{name}``: server_id, case-insensitive name, and the short prefix all reach it."""
@pytest.mark.asyncio
@pytest.mark.parametrize(
"route_name",
[
"catalog-server-id-001",
"CATALOG",
],
ids=["server_id", "uppercase_name"],
)
async def test_provider_gated_server_challenged_on_every_route_spelling(self, route_name):
from litellm.proxy._experimental.mcp_server import server as server_module
server = _catalog_server()
guardrail = _CallerSignInGuardrail(guardrail_name="sign-in-stub")
litellm.logging_callback_manager.add_litellm_callback(guardrail)
preflight = AsyncMock()
try:
with (
patch.object(
mcp_operations.global_mcp_server_manager,
"get_filtered_registry",
return_value={server.server_id: server},
),
patch.object(
mcp_operations.global_mcp_server_manager,
"preflight_token_exchange",
preflight,
),
patch.object(
mcp_operations,
"_get_allowed_mcp_servers",
AsyncMock(return_value=[server]),
),
pytest.raises(HTTPException) as exc,
):
await server_module._raise_preemptive_401_for_unauthenticated_servers(
scope={"type": "http", "method": "POST", "path": f"/mcp/{route_name}", "headers": []},
mcp_servers=[route_name],
oauth2_headers=None,
mcp_server_auth_headers=None,
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"),
client_ip=None,
)
finally:
litellm.logging_callback_manager.remove_callback_from_list_by_object(
litellm.callbacks, guardrail, require_self=False
)
assert exc.value.status_code == 401
headers = exc.value.headers or {}
assert "resource_metadata" in (headers.get("WWW-Authenticate") or "")
preflight.assert_not_awaited()
@pytest.mark.asyncio
async def test_provider_gated_server_challenged_on_short_prefix_route(self):
from litellm.proxy._experimental.mcp_server import server as server_module
from litellm.proxy._experimental.mcp_server.utils import compute_short_server_prefix
server = _catalog_server()
route_name = compute_short_server_prefix(server.server_id)
guardrail = _CallerSignInGuardrail(guardrail_name="sign-in-stub")
litellm.logging_callback_manager.add_litellm_callback(guardrail)
try:
with (
patch.object(
mcp_operations.global_mcp_server_manager,
"get_filtered_registry",
return_value={server.server_id: server},
),
patch.object(
mcp_operations,
"_get_allowed_mcp_servers",
AsyncMock(return_value=[server]),
),
pytest.raises(HTTPException) as exc,
):
await server_module._raise_preemptive_401_for_unauthenticated_servers(
scope={"type": "http", "method": "POST", "path": f"/mcp/{route_name}", "headers": []},
mcp_servers=[route_name],
oauth2_headers=None,
mcp_server_auth_headers=None,
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"),
client_ip=None,
)
finally:
litellm.logging_callback_manager.remove_callback_from_list_by_object(
litellm.callbacks, guardrail, require_self=False
)
assert exc.value.status_code == 401
@pytest.mark.asyncio
async def test_obo_challenge_www_authenticate_matches_main_byte_for_byte(self, monkeypatch):
"""The provider redesign must not change what an OBO server challenges with: the relative
RFC 9728 resource_metadata path plus the RFC 6750 invalid_token triple, exactly as main."""
from litellm.proxy._experimental.mcp_server import server as server_module
monkeypatch.delenv("SERVER_ROOT_PATH", raising=False)
obo = _make_obo_server("obo")
with (
patch.object(
mcp_operations.global_mcp_server_manager,
"get_mcp_server_answering_to",
return_value=obo,
),
pytest.raises(HTTPException) as exc,
):
await server_module._raise_preemptive_401_for_unauthenticated_servers(
scope={"type": "http", "method": "POST", "path": "/mcp/obo", "headers": []},
mcp_servers=["obo"],
oauth2_headers=None,
mcp_server_auth_headers=None,
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"),
client_ip=None,
)
assert exc.value.status_code == 401
assert (exc.value.headers or {}).get("WWW-Authenticate") == (
'Bearer resource_metadata="/.well-known/oauth-protected-resource/mcp/obo", '
'error="invalid_token", '
'error_description="Missing or invalid subject token; authenticate with the IdP and retry"'
)

View file

@ -1464,6 +1464,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",

View file

@ -14,6 +14,14 @@ from litellm.exceptions import Timeout as LitellmTimeout
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.secret_redaction import redact_string
from litellm.proxy._experimental.mcp_server.caller_sign_in import CallerSignIn, caller_sign_in_for
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.types import (
CredError,
ServerSpec,
TokenExchangeConfig,
)
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails.guardrail_hooks.agent_365 import (
Agent365Guardrail,
@ -27,9 +35,12 @@ from litellm.types.guardrails import (
LitellmParams,
SupportedGuardrailIntegrations,
)
from litellm.types.mcp import MCPAuth, MCPTransport
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,
AGENT_365_SCOPE_NAME,
Agent365GuardrailConfigModel,
)
@ -45,8 +56,46 @@ def _response(status_code: int, payload: Any = None, text: str | None = None) ->
return httpx.Response(status_code=status_code, text=text or "", request=request)
def _token_response(access_token: str = "obo-access-token", expires_in: int = 3599) -> httpx.Response:
return _response(200, {"access_token": access_token, "expires_in": expires_in})
class StubTokenExchanger:
"""The TokenExchanger the guardrail is injected with in tests: programmed Result queue plus a
per-subject cache honoring ``expires_at``, so cache and evaluate-401-invalidate behavior is
exercised the way the real OboTokenExchanger drives it."""
def __init__(self, results: list[Result[OAuthToken, CredError] | BaseException] | None = None):
self._results = list(results or [])
self._cache: dict[str, OAuthToken] = {}
self.calls: list[tuple[str, ServerSpec, TokenExchangeConfig]] = []
self.invalidations: list[str] = []
async def exchange(
self, subject_token: str, server: ServerSpec, config: TokenExchangeConfig, *, tenant_id: str = ""
) -> Result[OAuthToken, CredError]:
cached: Final = self._cache.get(subject_token)
if cached is not None and (cached.expires_at is None or cached.expires_at > time.time()):
return Ok(cached)
self.calls.append((subject_token, server, config))
if not self._results:
raise AssertionError("StubTokenExchanger ran out of programmed results")
result = self._results.pop(0)
if isinstance(result, BaseException):
raise result
if isinstance(result, Ok):
self._cache[subject_token] = result.ok
return result
async def invalidate(
self, subject_token: str, server: ServerSpec, config: TokenExchangeConfig, *, tenant_id: str = ""
) -> None:
self.invalidations.append(subject_token)
self._cache.pop(subject_token, None)
def _ok_exchange(access_token: str = "obo-access-token", expires_in: int = 3599) -> Ok[OAuthToken, CredError]:
return Ok(OAuthToken(access_token=access_token, expires_at=time.time() + expires_in))
def _obo_ok(access_token: str = "obo-access-token") -> list[Result[OAuthToken, CredError]]:
return [_ok_exchange(access_token)]
def _allow_response(correlation_id: str = "corr-1") -> httpx.Response:
@ -121,7 +170,9 @@ class FakeHandler:
def _make_guardrail(
handler: FakeHandler,
*,
exchanger: StubTokenExchanger | None = None,
unreachable_fallback: str = "fail_closed",
default_on: bool = True,
) -> Agent365Guardrail:
return Agent365Guardrail(
guardrail_name="agent-365-guard",
@ -130,21 +181,28 @@ def _make_guardrail(
client_secret="secret-123",
unreachable_fallback=unreachable_fallback,
async_handler=handler,
token_exchanger=exchanger if exchanger is not None else StubTokenExchanger(_obo_ok()),
event_hook="pre_mcp_call",
default_on=True,
default_on=default_on,
)
def _default_fallback_guardrail(handler: FakeHandler) -> Agent365Guardrail:
return Agent365Guardrail(
guardrail_name="agent-365-guard",
tenant_id="tenant-abc",
client_id="client-xyz",
client_secret="secret-123",
async_handler=handler,
event_hook="pre_mcp_call",
default_on=True,
)
def _default_fallback_guardrail(
handler: FakeHandler, exchanger: StubTokenExchanger | None = None
) -> Agent365Guardrail:
return _make_guardrail(handler, exchanger=exchanger)
def _server(**overrides: Any) -> MCPServer:
kwargs: Final[dict] = {
"server_id": "outlook-id",
"name": "outlook_mcp",
"server_name": "outlook_mcp",
"transport": MCPTransport.http,
"url": "https://outlook.test/mcp",
}
kwargs.update(overrides)
return MCPServer(**kwargs)
def _mcp_data(**overrides: Any) -> dict:
@ -257,14 +315,13 @@ class TestInitializeGuardrail:
resource_app_id="00000000-0000-0000-0000-000000000000",
agent_id="yaml-agent",
)
handler: Final = FakeHandler([_token_response(), _allow_response()])
handler: Final = FakeHandler([_allow_response()])
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
guardrail: Final = initialize_guardrail(params, {"guardrail_name": "a365-stale"}, async_handler=handler)
initialize_guardrail(params, {"guardrail_name": "a365-stale"}, async_handler=handler)
assert "ignoring api_base, resource_app_id, agent_id" in caplog.text
guardrail: Final = _make_guardrail(handler)
await _run(guardrail, _mcp_data())
token_call, evaluate_call = handler.calls
assert token_call.url == TOKEN_URL
assert token_call.data["scope"] == f"{AGENT_365_PROD_RESOURCE_APP_ID}/ThreatProtection.Evaluate.All"
evaluate_call: Final = handler.calls[0]
assert evaluate_call.url == EVALUATE_URL
assert evaluate_call.json["agentId"] == "my-agent-key"
@ -305,7 +362,7 @@ def _guardrail_info(data: dict) -> dict:
class TestAllowFlow:
@pytest.mark.asyncio
async def test_allowed_call_passes_through(self):
handler: Final = FakeHandler([_token_response(), _allow_response()])
handler: Final = FakeHandler([_allow_response()])
guardrail: Final = _make_guardrail(handler)
data: Final = _mcp_data()
result: Final = await _run(guardrail, data)
@ -319,25 +376,27 @@ class TestAllowFlow:
assert info["guardrail_response"]["latency_ms"] >= 0
@pytest.mark.asyncio
async def test_obo_exchange_form(self):
handler: Final = FakeHandler([_token_response(), _allow_response()])
guardrail: Final = _make_guardrail(handler)
async def test_obo_exchange_uses_the_built_entra_obo_config(self):
exchanger: Final = StubTokenExchanger(_obo_ok())
handler: Final = FakeHandler([_allow_response()])
guardrail: Final = _make_guardrail(handler, exchanger=exchanger)
await _run(guardrail, _mcp_data())
token_call: Final = handler.calls[0]
assert token_call.url == TOKEN_URL
assert token_call.data["grant_type"] == "urn:ietf:params:oauth:grant-type:jwt-bearer"
assert token_call.data["requested_token_use"] == "on_behalf_of"
assert token_call.data["assertion"] == FAKE_ASSERTION
assert token_call.data["client_id"] == "client-xyz"
assert token_call.data["client_secret"] == "secret-123"
assert token_call.data["scope"] == f"{AGENT_365_PROD_RESOURCE_APP_ID}/ThreatProtection.Evaluate.All"
subject_token, server, config = exchanger.calls[0]
assert subject_token == FAKE_ASSERTION
assert server.server_id == "agent-365:tenant-abc"
assert server.resource == AGENT_365_PROD_API_BASE
assert config.profile == "entra_obo"
assert config.token_exchange_endpoint == TOKEN_URL
assert config.client_id == "client-xyz"
assert config.client_secret is not None and config.client_secret.get_secret_value() == "secret-123"
assert config.scopes == (f"{AGENT_365_PROD_RESOURCE_APP_ID}/{AGENT_365_SCOPE_NAME}",)
@pytest.mark.asyncio
async def test_evaluate_payload(self):
handler: Final = FakeHandler([_token_response(), _allow_response()])
handler: Final = FakeHandler([_allow_response()])
guardrail: Final = _make_guardrail(handler)
await _run(guardrail, _mcp_data())
evaluate_call: Final = handler.calls[1]
evaluate_call: Final = handler.calls[0]
assert evaluate_call.url == EVALUATE_URL
assert evaluate_call.headers["Authorization"] == "Bearer obo-access-token"
assert evaluate_call.json["tool"] == {"name": "send_email"}
@ -348,11 +407,11 @@ class TestAllowFlow:
@pytest.mark.asyncio
async def test_evaluate_payload_includes_listed_tool_metadata(self):
handler: Final = FakeHandler([_token_response(), _allow_response()])
handler: Final = FakeHandler([_allow_response()])
guardrail: Final = _make_guardrail(handler)
schema: Final = {"type": "object", "properties": {"to": {"type": "string"}}, "required": ["to"]}
await _run(guardrail, _mcp_data(mcp_tool_description="Send an email", mcp_input_schema=schema))
assert handler.calls[1].json["tool"] == {
await _run(guardrail, _mcp_data(mcp_tool_description="Send an email", mcp_tool_input_schema=schema))
assert handler.calls[0].json["tool"] == {
"name": "send_email",
"description": "Send an email",
"inputSchema": schema,
@ -364,10 +423,17 @@ class TestAllowFlow:
[(None, None), ("", None), (None, ["not", "a", "schema"]), (42, "type: object")],
)
async def test_evaluate_payload_omits_missing_or_malformed_tool_metadata(self, description, schema):
handler: Final = FakeHandler([_token_response(), _allow_response()])
handler: Final = FakeHandler([_allow_response()])
guardrail: Final = _make_guardrail(handler)
await _run(guardrail, _mcp_data(mcp_tool_description=description, mcp_input_schema=schema))
assert handler.calls[1].json["tool"] == {"name": "send_email"}
await _run(guardrail, _mcp_data(mcp_tool_description=description, mcp_tool_input_schema=schema))
assert handler.calls[0].json["tool"] == {"name": "send_email"}
@pytest.mark.asyncio
async def test_agent_id_falls_back_to_key_alias(self):
handler: Final = FakeHandler([_allow_response()])
guardrail: Final = _make_guardrail(handler)
await _run(guardrail, _mcp_data())
assert handler.calls[0].json["agentId"] == "my-agent-key"
@pytest.mark.asyncio
async def test_non_mcp_call_type_skipped(self):
@ -385,26 +451,26 @@ class TestConversationId:
@pytest.mark.asyncio
async def test_two_calls_in_one_session_share_the_conversation_id(self):
handler: Final = FakeHandler([_token_response(), _allow_response(), _allow_response()])
handler: Final = FakeHandler([_allow_response(), _allow_response()])
guardrail: Final = _make_guardrail(handler)
for call_id in ("call-1", "call-2"):
await _run(
guardrail,
_mcp_data(litellm_call_id=call_id, litellm_logging_obj=_logging_obj(call_id, mcp_session_id="sess-A")),
)
assert [call.json["conversationId"] for call in handler.calls[1:]] == ["sess-A", "sess-A"]
assert [call.json["conversationId"] for call in handler.calls[:]] == ["sess-A", "sess-A"]
@pytest.mark.asyncio
async def test_server_recorded_session_beats_the_client_header(self):
handler: Final = FakeHandler([_token_response(), _allow_response()])
handler: Final = FakeHandler([_allow_response()])
guardrail: Final = _make_guardrail(handler)
data: Final = _mcp_data(litellm_logging_obj=_logging_obj("call-id-1", mcp_session_id="sess-from-logging"))
await _run(guardrail, data)
assert handler.calls[1].json["conversationId"] == "sess-from-logging"
assert handler.calls[0].json["conversationId"] == "sess-from-logging"
@pytest.mark.asyncio
async def test_sessionless_call_falls_back_to_the_request_call_id(self):
handler: Final = FakeHandler([_token_response(), _allow_response()])
handler: Final = FakeHandler([_allow_response()])
guardrail: Final = _make_guardrail(handler)
data: Final = _mcp_data(
metadata={"headers": {}},
@ -412,37 +478,37 @@ class TestConversationId:
litellm_logging_obj=_logging_obj("call-id-from-logging"),
)
await _run(guardrail, data)
assert handler.calls[1].json["conversationId"] == "call-id-from-data"
assert handler.calls[0].json["conversationId"] == "call-id-from-data"
@pytest.mark.asyncio
async def test_sessionless_call_without_request_call_id_uses_the_logging_call_id(self):
handler: Final = FakeHandler([_token_response(), _allow_response()])
handler: Final = FakeHandler([_allow_response()])
guardrail: Final = _make_guardrail(handler)
data: Final = _mcp_data(metadata={"headers": {}}, litellm_logging_obj=_logging_obj("call-id-from-logging"))
await _run(guardrail, data)
assert handler.calls[1].json["conversationId"] == "call-id-from-logging"
assert handler.calls[0].json["conversationId"] == "call-id-from-logging"
@pytest.mark.asyncio
async def test_session_id_header_case_insensitive(self):
handler: Final = FakeHandler([_token_response(), _allow_response()])
handler: Final = FakeHandler([_allow_response()])
guardrail: Final = _make_guardrail(handler)
data: Final = _mcp_data(metadata={"headers": {"Mcp-Session-Id": "sess-CASED"}})
await _run(guardrail, data)
assert handler.calls[1].json["conversationId"] == "sess-CASED"
assert handler.calls[0].json["conversationId"] == "sess-CASED"
@pytest.mark.asyncio
async def test_generates_uuid_when_no_identifier_available(self):
handler: Final = FakeHandler([_token_response(), _allow_response()])
handler: Final = FakeHandler([_allow_response()])
guardrail: Final = _make_guardrail(handler)
await _run(guardrail, _mcp_data(metadata={"headers": {}}, litellm_logging_obj=_logging_obj("")))
conversation_id: Final = handler.calls[1].json["conversationId"]
conversation_id: Final = handler.calls[0].json["conversationId"]
assert uuid.UUID(conversation_id).version == 4
class TestBlockFlow:
@pytest.mark.asyncio
async def test_blocked_call_raises_400(self):
handler: Final = FakeHandler([_token_response(), _block_response(message="Injection detected")])
handler: Final = FakeHandler([_block_response(message="Injection detected")])
guardrail: Final = _make_guardrail(handler)
data: Final = _mcp_data()
with pytest.raises(HTTPException) as exc_info:
@ -458,7 +524,7 @@ class TestBlockFlow:
@pytest.mark.asyncio
async def test_blocked_even_with_fail_open(self):
handler: Final = FakeHandler([_token_response(), _block_response()])
handler: Final = FakeHandler([_block_response()])
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open")
with pytest.raises(HTTPException) as exc_info:
await _run(guardrail, _mcp_data())
@ -467,7 +533,7 @@ class TestBlockFlow:
@pytest.mark.asyncio
@pytest.mark.parametrize("status", ["Skipped", "FailedOpen"])
async def test_explicit_block_wins_over_non_evaluated_status(self, status):
handler: Final = FakeHandler([_token_response(), _block_response(status=status)])
handler: Final = FakeHandler([_block_response(status=status)])
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open")
data: Final = _mcp_data()
with pytest.raises(HTTPException) as exc_info:
@ -483,7 +549,7 @@ class TestDefenderNotEvaluated:
@pytest.mark.asyncio
@pytest.mark.parametrize("status", ["Skipped", "FailedOpen"])
async def test_fail_closed_blocks_allowed_but_unevaluated_call(self, status):
handler: Final = FakeHandler([_token_response(), _not_evaluated_response(status)])
handler: Final = FakeHandler([_not_evaluated_response(status)])
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_closed")
data: Final = _mcp_data()
with pytest.raises(HTTPException) as exc_info:
@ -500,7 +566,7 @@ class TestDefenderNotEvaluated:
@pytest.mark.asyncio
@pytest.mark.parametrize("status", ["Skipped", "FailedOpen"])
async def test_fail_open_allows_unevaluated_call_as_unscanned(self, status):
handler: Final = FakeHandler([_token_response(), _not_evaluated_response(status)])
handler: Final = FakeHandler([_not_evaluated_response(status)])
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open")
data: Final = _mcp_data()
result: Final = await _run(guardrail, data)
@ -514,7 +580,7 @@ class TestDefenderNotEvaluated:
@pytest.mark.asyncio
@pytest.mark.parametrize("payload", [{"allowed": True}, {"allowed": True, "defender": {"verdict": "Allow"}}])
async def test_allowed_without_defender_status_is_not_an_evaluated_allow(self, payload):
handler: Final = FakeHandler([_token_response(), _response(200, payload)])
handler: Final = FakeHandler([_response(200, payload)])
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_closed")
data: Final = _mcp_data()
with pytest.raises(HTTPException) as exc_info:
@ -525,7 +591,7 @@ class TestDefenderNotEvaluated:
@pytest.mark.asyncio
async def test_http_400_always_blocks_even_fail_open(self):
handler: Final = FakeHandler([_token_response(), _response(400, text="Bad request: serverName missing")])
handler: Final = FakeHandler([_response(400, text="Bad request: serverName missing")])
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open")
with pytest.raises(HTTPException) as exc_info:
await _run(guardrail, _mcp_data())
@ -534,18 +600,20 @@ class TestDefenderNotEvaluated:
AVAILABILITY_FAILURES: Final = (
pytest.param([_token_response(), httpx.ReadTimeout("timed out")], id="evaluate-timeout"),
pytest.param([_token_response(), _response(502, text="bad gateway")], id="evaluate-5xx"),
pytest.param([_token_response(), _not_evaluated_response("Skipped")], id="evaluate-skipped"),
pytest.param([_response(503, text="entra down")], id="entra-5xx"),
pytest.param([httpx.ReadTimeout("timed out")], _obo_ok(), id="evaluate-timeout"),
pytest.param([_response(502, text="bad gateway")], _obo_ok(), id="evaluate-5xx"),
pytest.param([_not_evaluated_response("Skipped")], _obo_ok(), id="evaluate-skipped"),
pytest.param([], [Error(CredError.of_upstream_unavailable("entra down"))], id="entra-5xx"),
)
class TestFailOpenOptIn:
@pytest.mark.asyncio
@pytest.mark.parametrize("responses", AVAILABILITY_FAILURES)
async def test_constructor_default_blocks_each_availability_failure_with_503(self, responses):
guardrail: Final = _default_fallback_guardrail(FakeHandler(responses))
@pytest.mark.parametrize(("responses", "exchange_results"), AVAILABILITY_FAILURES)
async def test_constructor_default_blocks_each_availability_failure_with_503(self, responses, exchange_results):
guardrail: Final = _default_fallback_guardrail(
FakeHandler(responses), StubTokenExchanger(exchange_results)
)
assert guardrail.unreachable_fallback == "fail_closed"
with pytest.raises(HTTPException) as exc_info:
await _run(guardrail, _mcp_data())
@ -553,9 +621,15 @@ class TestFailOpenOptIn:
assert "fail_closed" in exc_info.value.detail["message"]
@pytest.mark.asyncio
@pytest.mark.parametrize("responses", AVAILABILITY_FAILURES)
async def test_opted_in_fail_open_lets_each_availability_failure_through_as_failed_to_respond(self, responses):
guardrail: Final = _make_guardrail(FakeHandler(responses), unreachable_fallback="fail_open")
@pytest.mark.parametrize(("responses", "exchange_results"), AVAILABILITY_FAILURES)
async def test_opted_in_fail_open_lets_each_availability_failure_through_as_failed_to_respond(
self, responses, exchange_results
):
guardrail: Final = _make_guardrail(
FakeHandler(responses),
exchanger=StubTokenExchanger(exchange_results),
unreachable_fallback="fail_open",
)
data: Final = _mcp_data()
assert await _run(guardrail, data) is data
info: Final = _guardrail_info(data)
@ -564,7 +638,7 @@ class TestFailOpenOptIn:
@pytest.mark.asyncio
async def test_opted_in_fail_open_logs_the_unscanned_call_at_error_level(self, caplog):
handler: Final = FakeHandler([_token_response(), httpx.ReadTimeout("timed out")])
handler: Final = FakeHandler([httpx.ReadTimeout("timed out")])
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open")
with caplog.at_level(logging.ERROR, logger="LiteLLM Proxy"):
await _run(guardrail, _mcp_data())
@ -573,7 +647,7 @@ class TestFailOpenOptIn:
@pytest.mark.asyncio
async def test_opted_in_fail_open_still_blocks_a_policy_block(self):
handler: Final = FakeHandler([_token_response(), _block_response()])
handler: Final = FakeHandler([_block_response()])
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open")
with pytest.raises(HTTPException) as exc_info:
await _run(guardrail, _mcp_data())
@ -584,10 +658,7 @@ class TestUnreachableFallback:
@pytest.mark.asyncio
async def test_evaluate_litellm_timeout_fail_closed(self):
handler: Final = FakeHandler(
[
_token_response(),
LitellmTimeout(message="Connection timed out", model="default-model-name", llm_provider="httpx"),
]
[LitellmTimeout(message="Connection timed out", model="default-model-name", llm_provider="httpx")]
)
guardrail: Final = _make_guardrail(handler)
with pytest.raises(HTTPException) as exc_info:
@ -596,7 +667,7 @@ class TestUnreachableFallback:
@pytest.mark.asyncio
async def test_evaluate_timeout_fail_closed(self):
handler: Final = FakeHandler([_token_response(), httpx.ReadTimeout("timed out")])
handler: Final = FakeHandler([httpx.ReadTimeout("timed out")])
guardrail: Final = _make_guardrail(handler)
with pytest.raises(HTTPException) as exc_info:
await _run(guardrail, _mcp_data())
@ -605,7 +676,7 @@ class TestUnreachableFallback:
@pytest.mark.asyncio
async def test_evaluate_timeout_fail_open(self):
handler: Final = FakeHandler([_token_response(), httpx.ReadTimeout("timed out")])
handler: Final = FakeHandler([httpx.ReadTimeout("timed out")])
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open")
data: Final = _mcp_data()
result: Final = await _run(guardrail, data)
@ -616,7 +687,7 @@ class TestUnreachableFallback:
@pytest.mark.asyncio
async def test_evaluate_5xx_fail_closed(self):
handler: Final = FakeHandler([_token_response(), _response(502, text="bad gateway")])
handler: Final = FakeHandler([_response(502, text="bad gateway")])
guardrail: Final = _make_guardrail(handler)
with pytest.raises(HTTPException) as exc_info:
await _run(guardrail, _mcp_data())
@ -634,11 +705,13 @@ class TestUnreachableFallback:
@pytest.mark.asyncio
async def test_non_jwt_bearer_token_fail_closed(self):
exchanger: Final = StubTokenExchanger()
handler: Final = FakeHandler([])
guardrail: Final = _make_guardrail(handler)
guardrail: Final = _make_guardrail(handler, exchanger=exchanger)
with pytest.raises(HTTPException) as exc_info:
await _run(guardrail, _mcp_data(incoming_bearer_token="sk-litellm-virtual-key"))
assert exc_info.value.status_code == 401
assert exchanger.calls == []
@pytest.mark.asyncio
async def test_missing_bearer_token_blocks_even_fail_open(self):
@ -655,18 +728,19 @@ class TestUnreachableFallback:
@pytest.mark.asyncio
async def test_obo_rejected_blocks_even_fail_open(self):
handler: Final = FakeHandler(
[_response(400, {"error": "invalid_grant", "error_description": "AADSTS50013: bad assertion"})]
exchanger: Final = StubTokenExchanger(
[Error(CredError.of_unauthorized("IdP rejected the subject token (HTTP 400)", claims="invalid_grant"))]
)
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open")
handler: Final = FakeHandler([])
guardrail: Final = _make_guardrail(handler, exchanger=exchanger, unreachable_fallback="fail_open")
with pytest.raises(HTTPException) as exc_info:
await _run(guardrail, _mcp_data())
assert exc_info.value.status_code == 401
assert "invalid_grant" in exc_info.value.detail["message"]
assert "On-Behalf-Of token exchange was rejected" in exc_info.value.detail["message"]
@pytest.mark.asyncio
async def test_evaluate_4xx_blocks_even_fail_open(self):
handler: Final = FakeHandler([_token_response(), _response(403, text="obo token lacks the scope")])
handler: Final = FakeHandler([_response(403, text="obo token lacks the scope")])
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open")
data: Final = _mcp_data()
with pytest.raises(HTTPException) as exc_info:
@ -680,24 +754,24 @@ class TestUnreachableFallback:
@pytest.mark.asyncio
async def test_obo_rejected_fail_closed(self):
handler: Final = FakeHandler(
[_response(400, {"error": "invalid_grant", "error_description": "AADSTS50013: bad assertion"})]
exchanger: Final = StubTokenExchanger(
[Error(CredError.of_unauthorized("IdP rejected the subject token (HTTP 400)"))]
)
guardrail: Final = _make_guardrail(handler)
handler: Final = FakeHandler([])
guardrail: Final = _make_guardrail(handler, exchanger=exchanger)
with pytest.raises(HTTPException) as exc_info:
await _run(guardrail, _mcp_data())
assert exc_info.value.status_code == 401
assert "invalid_grant" in exc_info.value.detail["message"]
assert "On-Behalf-Of token exchange was rejected" in exc_info.value.detail["message"]
@pytest.mark.asyncio
@pytest.mark.parametrize(
"error_code", ["invalid_client", "unauthorized_client", "invalid_scope", "invalid_resource"]
)
async def test_gateway_credential_rejection_is_unavailable_not_a_caller_401(self, error_code: str):
handler: Final = FakeHandler(
[_response(401, {"error": error_code, "error_description": "AADSTS7000215: invalid client secret"})]
)
guardrail: Final = _make_guardrail(handler)
exchanger: Final = StubTokenExchanger([Error(CredError.of_misconfigured(error_code))])
handler: Final = FakeHandler([])
guardrail: Final = _make_guardrail(handler, exchanger=exchanger)
data: Final = _mcp_data()
with pytest.raises(HTTPException) as exc_info:
await _run(guardrail, data)
@ -710,23 +784,12 @@ class TestUnreachableFallback:
assert "client_secret" in info["guardrail_response"]["reason"]
@pytest.mark.asyncio
@pytest.mark.parametrize("aadsts_code", [5002710, 5002723], ids=["malformed-header", "no-kid"])
async def test_malformed_assertion_reported_as_invalid_client_is_a_caller_401(self, aadsts_code: int):
"""Entra answers ``invalid_client`` for a forged or garbled assertion (AADSTS50027xx) exactly as for a
bad gateway secret; the sub-code is what says the caller, not the gateway, has to fix it."""
handler: Final = FakeHandler(
[
_response(
401,
{
"error": "invalid_client",
"error_description": f"AADSTS{aadsts_code}: Invalid JWT token.",
"error_codes": [aadsts_code],
},
)
]
async def test_caller_rejection_reason_does_not_blame_the_gateway_credentials(self):
exchanger: Final = StubTokenExchanger(
[Error(CredError.of_unauthorized("IdP rejected the subject token (HTTP 401)"))]
)
guardrail: Final = _make_guardrail(handler)
handler: Final = FakeHandler([])
guardrail: Final = _make_guardrail(handler, exchanger=exchanger)
data: Final = _mcp_data()
with pytest.raises(HTTPException) as exc_info:
await _run(guardrail, data)
@ -735,10 +798,9 @@ class TestUnreachableFallback:
@pytest.mark.asyncio
async def test_gateway_credential_rejection_follows_fail_open(self):
handler: Final = FakeHandler(
[_response(401, {"error": "invalid_client", "error_description": "AADSTS7000215: invalid client secret"})]
)
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open")
exchanger: Final = StubTokenExchanger([Error(CredError.of_misconfigured("invalid_client"))])
handler: Final = FakeHandler([])
guardrail: Final = _make_guardrail(handler, exchanger=exchanger, unreachable_fallback="fail_open")
data: Final = _mcp_data()
result: Final = await _run(guardrail, data)
assert result is data
@ -747,10 +809,43 @@ class TestUnreachableFallback:
assert info["guardrail_response"]["verdict"] == "Unscanned"
assert "invalid_client" in info["guardrail_response"]["reason"]
@pytest.mark.asyncio
async def test_exchange_upstream_unavailable_is_unavailable_with_the_summary_not_a_caller_401(self):
exchanger: Final = StubTokenExchanger(
[Error(CredError.of_upstream_unavailable("token exchange did not return a usable access token"))]
)
handler: Final = FakeHandler([])
guardrail: Final = _make_guardrail(handler, exchanger=exchanger)
data: Final = _mcp_data()
with pytest.raises(HTTPException) as exc_info:
await _run(guardrail, data)
assert exc_info.value.status_code == 503
info: Final = _guardrail_info(data)
assert info["guardrail_status"] == "guardrail_failed_to_respond"
assert info["guardrail_response"]["verdict"] == "Unavailable"
assert (
info["guardrail_response"]["reason"]
== "the Entra token exchange failed (upstream unavailable: token exchange did not return a usable access token)"
)
@pytest.mark.asyncio
async def test_exchange_upstream_unavailable_follows_fail_open(self):
exchanger: Final = StubTokenExchanger([Error(CredError.of_upstream_unavailable("Entra throttled the exchange"))])
handler: Final = FakeHandler([])
guardrail: Final = _make_guardrail(handler, exchanger=exchanger, unreachable_fallback="fail_open")
data: Final = _mcp_data()
result: Final = await _run(guardrail, data)
assert result is data
info: Final = _guardrail_info(data)
assert info["guardrail_status"] == "guardrail_failed_to_respond"
assert info["guardrail_response"]["verdict"] == "Unscanned"
assert "Entra throttled the exchange" in info["guardrail_response"]["reason"]
@pytest.mark.asyncio
async def test_obo_endpoint_5xx_fail_open(self):
handler: Final = FakeHandler([_response(503, text="entra down")])
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open")
exchanger: Final = StubTokenExchanger([Error(CredError.of_upstream_unavailable("entra down"))])
handler: Final = FakeHandler([])
guardrail: Final = _make_guardrail(handler, exchanger=exchanger, unreachable_fallback="fail_open")
data: Final = _mcp_data()
result: Final = await _run(guardrail, data)
assert result is data
@ -762,47 +857,33 @@ class TestUnreachableFallback:
class TestOboTokenCache:
@pytest.mark.asyncio
async def test_same_assertion_reuses_token(self):
handler: Final = FakeHandler([_token_response(), _allow_response(), _allow_response()])
guardrail: Final = _make_guardrail(handler)
exchanger: Final = StubTokenExchanger(_obo_ok())
handler: Final = FakeHandler([_allow_response(), _allow_response()])
guardrail: Final = _make_guardrail(handler, exchanger=exchanger)
await _run(guardrail, _mcp_data())
await _run(guardrail, _mcp_data())
token_calls: Final = [c for c in handler.calls if c.url == TOKEN_URL]
assert len(token_calls) == 1
assert len(exchanger.calls) == 1
@pytest.mark.asyncio
async def test_different_assertions_get_distinct_tokens(self):
other_assertion: Final = "eyJhbGciOi.eyJvdGhlciI.b3RoZXJzaWc"
handler: Final = FakeHandler(
[
_token_response(access_token="token-a"),
_allow_response(),
_token_response(access_token="token-b"),
_allow_response(),
]
)
guardrail: Final = _make_guardrail(handler)
exchanger: Final = StubTokenExchanger([_ok_exchange("token-a"), _ok_exchange("token-b")])
handler: Final = FakeHandler([_allow_response(), _allow_response()])
guardrail: Final = _make_guardrail(handler, exchanger=exchanger)
await _run(guardrail, _mcp_data())
await _run(guardrail, _mcp_data(incoming_bearer_token=other_assertion))
token_calls: Final = [c for c in handler.calls if c.url == TOKEN_URL]
assert len(token_calls) == 2
assert handler.calls[3].headers["Authorization"] == "Bearer token-b"
assert len(exchanger.calls) == 2
assert handler.calls[1].headers["Authorization"] == "Bearer token-b"
@pytest.mark.asyncio
async def test_expired_token_refreshed(self):
handler: Final = FakeHandler(
[
_token_response(access_token="short-lived", expires_in=1),
_allow_response(),
_token_response(access_token="fresh"),
_allow_response(),
]
)
guardrail: Final = _make_guardrail(handler)
exchanger: Final = StubTokenExchanger([_ok_exchange("short-lived", -1), _ok_exchange("fresh")])
handler: Final = FakeHandler([_allow_response(), _allow_response()])
guardrail: Final = _make_guardrail(handler, exchanger=exchanger)
await _run(guardrail, _mcp_data())
await _run(guardrail, _mcp_data())
token_calls: Final = [c for c in handler.calls if c.url == TOKEN_URL]
assert len(token_calls) == 2
assert handler.calls[3].headers["Authorization"] == "Bearer fresh"
assert len(exchanger.calls) == 2
assert handler.calls[1].headers["Authorization"] == "Bearer fresh"
class TestEarlyPhasePassthrough:
@ -835,34 +916,27 @@ class TestRegistryDiscovery:
class TestMalformedResponses:
@pytest.mark.asyncio
async def test_obo_html_body_fail_open(self):
handler: Final = FakeHandler([_response(200, text="<html>blocked by egress proxy</html>")])
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open")
async def test_obo_upstream_unavailable_fail_open(self):
exchanger: Final = StubTokenExchanger([Error(CredError.of_upstream_unavailable("entra returned no token"))])
handler: Final = FakeHandler([])
guardrail: Final = _make_guardrail(handler, exchanger=exchanger, unreachable_fallback="fail_open")
data: Final = _mcp_data()
result: Final = await _run(guardrail, data)
assert result is data
assert _guardrail_info(data)["guardrail_response"]["verdict"] == "Unscanned"
@pytest.mark.asyncio
async def test_obo_html_body_fail_closed(self):
handler: Final = FakeHandler([_response(200, text="<html>outage</html>")])
guardrail: Final = _make_guardrail(handler)
with pytest.raises(HTTPException) as exc_info:
await _run(guardrail, _mcp_data())
assert exc_info.value.status_code == 503
assert "non-JSON" in exc_info.value.detail["message"]
@pytest.mark.asyncio
async def test_obo_non_object_json_fail_closed(self):
handler: Final = FakeHandler([_response(200, ["not", "a", "dict"])])
guardrail: Final = _make_guardrail(handler)
async def test_obo_upstream_unavailable_fail_closed(self):
exchanger: Final = StubTokenExchanger([Error(CredError.of_upstream_unavailable("entra returned no token"))])
handler: Final = FakeHandler([])
guardrail: Final = _make_guardrail(handler, exchanger=exchanger)
with pytest.raises(HTTPException) as exc_info:
await _run(guardrail, _mcp_data())
assert exc_info.value.status_code == 503
@pytest.mark.asyncio
async def test_evaluate_html_body_fail_open(self):
handler: Final = FakeHandler([_token_response(), _response(200, text="<html>waf page</html>")])
handler: Final = FakeHandler([_response(200, text="<html>waf page</html>")])
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open")
data: Final = _mcp_data()
result: Final = await _run(guardrail, data)
@ -871,7 +945,7 @@ class TestMalformedResponses:
@pytest.mark.asyncio
async def test_evaluate_html_body_fail_closed(self):
handler: Final = FakeHandler([_token_response(), _response(200, text="<html>waf page</html>")])
handler: Final = FakeHandler([_response(200, text="<html>waf page</html>")])
guardrail: Final = _make_guardrail(handler)
with pytest.raises(HTTPException) as exc_info:
await _run(guardrail, _mcp_data())
@ -879,7 +953,7 @@ class TestMalformedResponses:
@pytest.mark.asyncio
async def test_evaluate_non_object_json_fail_closed(self):
handler: Final = FakeHandler([_token_response(), _response(200, "allowed")])
handler: Final = FakeHandler([_response(200, "allowed")])
guardrail: Final = _make_guardrail(handler)
with pytest.raises(HTTPException) as exc_info:
await _run(guardrail, _mcp_data())
@ -892,7 +966,7 @@ class TestMalformedResponses:
ids=["missing", "null", "string-true", "int-one", "string-false"],
)
async def test_evaluate_non_boolean_allowed_fail_closed(self, verdict: dict):
handler: Final = FakeHandler([_token_response(), _response(200, verdict)])
handler: Final = FakeHandler([_response(200, verdict)])
guardrail: Final = _make_guardrail(handler)
data: Final = _mcp_data()
with pytest.raises(HTTPException) as exc_info:
@ -908,7 +982,7 @@ class TestMalformedResponses:
ids=["missing", "null", "string-true", "int-one", "string-false"],
)
async def test_evaluate_non_boolean_allowed_fail_open(self, verdict: dict):
handler: Final = FakeHandler([_token_response(), _response(200, verdict)])
handler: Final = FakeHandler([_response(200, verdict)])
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open")
data: Final = _mcp_data()
result: Final = await _run(guardrail, data)
@ -917,21 +991,21 @@ class TestMalformedResponses:
assert _guardrail_info(data)["guardrail_status"] == "guardrail_failed_to_respond"
@pytest.mark.asyncio
async def test_bad_expires_in_still_allows(self):
handler: Final = FakeHandler(
[_response(200, {"access_token": "tok-1", "expires_in": "soon"}), _allow_response()]
)
guardrail: Final = _make_guardrail(handler)
async def test_obo_token_without_expiry_still_allows(self):
exchanger: Final = StubTokenExchanger([Ok(OAuthToken(access_token="tok-1", expires_at=None))])
handler: Final = FakeHandler([_allow_response()])
guardrail: Final = _make_guardrail(handler, exchanger=exchanger)
data: Final = _mcp_data()
result: Final = await _run(guardrail, data)
assert result is data
@pytest.mark.asyncio
async def test_obo_litellm_timeout_fail_open(self):
handler: Final = FakeHandler(
exchanger: Final = StubTokenExchanger(
[LitellmTimeout(message="Connection timed out", model="default-model-name", llm_provider="httpx")]
)
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open")
handler: Final = FakeHandler([])
guardrail: Final = _make_guardrail(handler, exchanger=exchanger, unreachable_fallback="fail_open")
data: Final = _mcp_data()
result: Final = await _run(guardrail, data)
assert result is data
@ -940,28 +1014,18 @@ class TestMalformedResponses:
class TestDeltaHardening:
@pytest.mark.asyncio
async def test_non_string_access_token_fail_closed(self):
handler: Final = FakeHandler([_response(200, {"access_token": None, "expires_in": 3599})])
guardrail: Final = _make_guardrail(handler)
with pytest.raises(HTTPException) as exc_info:
await _run(guardrail, _mcp_data())
assert exc_info.value.status_code == 503
assert "access_token" in exc_info.value.detail["message"]
@pytest.mark.asyncio
async def test_numeric_string_expires_in_honored(self):
handler: Final = FakeHandler(
[_response(200, {"access_token": "tok-9", "expires_in": "120"}), _allow_response()]
)
guardrail: Final = _make_guardrail(handler)
async def test_unexpired_exchange_result_is_reused(self):
exchanger: Final = StubTokenExchanger([_ok_exchange("tok-9", 120)])
handler: Final = FakeHandler([_allow_response(), _allow_response()])
guardrail: Final = _make_guardrail(handler, exchanger=exchanger)
await _run(guardrail, _mcp_data())
entries: Final = list(guardrail._obo_token_cache.values())
assert len(entries) == 1
assert entries[0][1] - time.time() < 200
await _run(guardrail, _mcp_data())
assert len(exchanger.calls) == 1
assert handler.calls[1].headers["Authorization"] == "Bearer tok-9"
@pytest.mark.asyncio
async def test_evaluate_400_records_intervention(self):
handler: Final = FakeHandler([_token_response(), _response(400, text="bad request shape")])
handler: Final = FakeHandler([_response(400, text="bad request shape")])
guardrail: Final = _make_guardrail(handler)
data: Final = _mcp_data()
with pytest.raises(HTTPException) as exc_info:
@ -975,25 +1039,19 @@ class TestDeltaHardening:
class TestVeriaHardening:
@pytest.mark.asyncio
async def test_evaluate_401_evicts_cached_obo_token(self):
handler: Final = FakeHandler(
[
_token_response(),
_response(401, text="token expired"),
_token_response(access_token="tok-2"),
_allow_response(),
]
)
guardrail: Final = _make_guardrail(handler)
exchanger: Final = StubTokenExchanger(_obo_ok() + _obo_ok("tok-2"))
handler: Final = FakeHandler([_response(401, text="token expired"), _allow_response()])
guardrail: Final = _make_guardrail(handler, exchanger=exchanger)
with pytest.raises(HTTPException):
await _run(guardrail, _mcp_data())
result: Final = await _run(guardrail, _mcp_data())
assert result is not None
token_calls: Final = [c for c in handler.calls if c.url == TOKEN_URL]
assert len(token_calls) == 2
assert exchanger.invalidations == [FAKE_ASSERTION]
assert len(exchanger.calls) == 2
@pytest.mark.asyncio
async def test_evaluate_429_blocks_even_fail_open_as_throttled(self):
handler: Final = FakeHandler([_token_response(), _response(429, text="slow down")])
handler: Final = FakeHandler([_response(429, text="slow down")])
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open")
data: Final = _mcp_data()
with pytest.raises(HTTPException) as exc_info:
@ -1006,7 +1064,7 @@ class TestVeriaHardening:
@pytest.mark.asyncio
async def test_evaluate_500_is_unavailable(self):
handler: Final = FakeHandler([_token_response(), _response(500, text="oops")])
handler: Final = FakeHandler([_response(500, text="oops")])
guardrail: Final = _make_guardrail(handler)
with pytest.raises(HTTPException) as exc_info:
await _run(guardrail, _mcp_data())
@ -1014,51 +1072,23 @@ class TestVeriaHardening:
assert "500" in exc_info.value.detail["message"]
@pytest.mark.asyncio
async def test_token_endpoint_429_blocks_even_fail_open_as_throttled(self):
handler: Final = FakeHandler(
[_response(429, {"error": "temporarily_throttled", "error_description": "AADSTS90056"})]
async def test_token_endpoint_unauthorized_is_a_caller_401_even_fail_open(self):
exchanger: Final = StubTokenExchanger(
[Error(CredError.of_unauthorized("IdP rejected the subject token (HTTP 429)"))]
)
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open")
handler: Final = FakeHandler([])
guardrail: Final = _make_guardrail(handler, exchanger=exchanger, unreachable_fallback="fail_open")
data: Final = _mcp_data()
with pytest.raises(HTTPException) as exc_info:
await _run(guardrail, data)
assert exc_info.value.status_code == 503
assert "429" in exc_info.value.detail["message"]
assert exc_info.value.status_code == 401
info: Final = _guardrail_info(data)
assert info["guardrail_status"] == "guardrail_failed_to_respond"
assert info["guardrail_response"]["verdict"] == "Throttled"
@pytest.mark.asyncio
async def test_token_endpoint_408_non_json_blocks_as_throttled(self):
handler: Final = FakeHandler([_response(408, text="Request Timeout")])
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open")
data: Final = _mcp_data()
with pytest.raises(HTTPException) as exc_info:
await _run(guardrail, data)
assert exc_info.value.status_code == 503
assert _guardrail_info(data)["guardrail_response"]["verdict"] == "Throttled"
@pytest.mark.asyncio
async def test_token_endpoint_4xx_html_stays_infra_fail_open(self):
handler: Final = FakeHandler([_response(403, text="<html>waf block page</html>")])
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open")
data: Final = _mcp_data()
result: Final = await _run(guardrail, data)
assert result is data
assert _guardrail_info(data)["guardrail_response"]["verdict"] == "Unscanned"
@pytest.mark.asyncio
async def test_entra_200_missing_access_token_is_malformed(self):
handler: Final = FakeHandler([_response(200, {"token_type": "Bearer"})])
guardrail: Final = _make_guardrail(handler)
with pytest.raises(HTTPException) as exc_info:
await _run(guardrail, _mcp_data())
assert exc_info.value.status_code == 503
assert "access_token" in exc_info.value.detail["message"]
assert info["guardrail_status"] == "guardrail_intervened"
assert info["guardrail_response"]["verdict"] == "Rejected"
@pytest.mark.asyncio
async def test_evaluate_5xx_fail_open_allows_unscanned_once(self):
handler: Final = FakeHandler([_token_response(), _response(502, text='{"error": "bad gateway"}')])
handler: Final = FakeHandler([_response(502, text='{"error": "bad gateway"}')])
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open")
data: Final = _mcp_data()
result: Final = await _run(guardrail, data)
@ -1099,7 +1129,7 @@ class TestFinalArgumentsEvaluated:
@pytest.mark.asyncio
@pytest.mark.parametrize("agent_365_first", [True, False], ids=["agent_365_then_masker", "masker_then_agent_365"])
async def test_agent_365_receives_the_arguments_sent_upstream(self, agent_365_first: bool):
handler: Final = FakeHandler([_token_response(), _allow_response()])
handler: Final = FakeHandler([_allow_response()])
guardrail: Final = _make_guardrail(handler)
masker: Final = _ArgumentMasker("arg-rewrite")
registered: Final = (guardrail, masker) if agent_365_first else (masker, guardrail)
@ -1116,4 +1146,61 @@ class TestFinalArgumentsEvaluated:
litellm.callbacks, callback, require_self=False
)
assert result["modified_arguments"] == {"turn": "please [REWRITE_ME_REDACTED] now"}
assert handler.calls[1].json["arguments"] == {"turn": "please [REWRITE_ME_REDACTED] now"}
assert handler.calls[0].json["arguments"] == {"turn": "please [REWRITE_ME_REDACTED] now"}
class TestCallerSignIn:
"""The guardrail is a CallerSignInProvider: it tells the MCP connect path which issuer and scope
the caller must sign in for before the first tool call, and only for servers that keep the
caller's Authorization with the gateway."""
def test_gated_server_advertises_entra_issuer_and_gateway_scope(self):
guardrail: Final = _make_guardrail(FakeHandler([]))
sign_in: Final = guardrail.caller_sign_in(_server(), None)
assert sign_in == CallerSignIn(
issuers=("https://login.microsoftonline.com/tenant-abc/v2.0",),
scopes=("api://client-xyz/access_as_user",),
)
def test_default_off_guardrail_does_not_gate(self):
guardrail: Final = _make_guardrail(FakeHandler([]), default_on=False)
assert guardrail.caller_sign_in(_server(), None) is None
def test_server_that_fills_authorization_itself_is_not_gated(self):
guardrail: Final = _make_guardrail(FakeHandler([]))
assert guardrail.caller_sign_in(_server(auth_type=MCPAuth.oauth2), None) is None
assert guardrail.caller_sign_in(_server(extra_headers=["authorization"]), None) is None
def test_opted_out_key_does_not_gate(self):
class _OptedOut(Agent365Guardrail):
def should_run_guardrail(self, data, event_type) -> bool:
return False
guardrail: Final = _OptedOut(
guardrail_name="a365-off",
tenant_id="tenant-abc",
client_id="client-xyz",
client_secret="secret-123",
token_exchanger=StubTokenExchanger(),
default_on=True,
)
assert guardrail.caller_sign_in(_server(), UserAPIKeyAuth(api_key="k", user_id="u-1")) is None
assert guardrail.caller_sign_in(_server(), None) is not None
def test_obo_server_with_provider_advertises_both_issuers_and_scopes(self, monkeypatch):
monkeypatch.setenv("JWT_ISSUER", "https://jwt-idp.test")
guardrail: Final = _make_guardrail(FakeHandler([]))
litellm.logging_callback_manager.add_litellm_callback(guardrail)
try:
server: Final = _server(auth_type=MCPAuth.oauth2_token_exchange, scopes=["read"])
sign_in: Final = caller_sign_in_for(server, None)
finally:
litellm.logging_callback_manager.remove_callback_from_list_by_object(
litellm.callbacks, guardrail, require_self=False
)
assert sign_in is not None
assert sign_in.issuers == (
"https://jwt-idp.test",
"https://login.microsoftonline.com/tenant-abc/v2.0",
)
assert sign_in.scopes == ("read", "api://client-xyz/access_as_user")