From 4cadb402e5ab71650d9ca64c13984b2aeded49b8 Mon Sep 17 00:00:00 2001 From: yucheng Date: Sat, 26 Sep 2026 09:32:46 +0000 Subject: [PATCH] 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> --- .../mcp_server/caller_sign_in.py | 124 ++++ .../mcp_server/discoverable_endpoints.py | 58 +- .../mcp_server/mcp_server_manager.py | 28 +- .../_experimental/mcp_server/operations.py | 10 +- .../token_exchange_provider.py | 5 + .../proxy/_experimental/mcp_server/server.py | 40 +- .../proxy/_experimental/mcp_server/utils.py | 8 + .../guardrail_hooks/agent_365/agent_365.py | 229 +++---- .../types/mcp_server/mcp_server_manager.py | 23 +- .../test_token_exchange_provider.py | 42 +- .../mcp_server/test_caller_sign_in.py | 134 +++++ .../mcp_server/test_discoverable_endpoints.py | 26 +- .../mcp_server/test_mcp_server.py | 181 +++++- .../mcp_server/test_mcp_server_manager.py | 13 + .../guardrail_hooks/test_agent_365.py | 561 ++++++++++-------- 15 files changed, 1005 insertions(+), 477 deletions(-) create mode 100644 litellm/proxy/_experimental/mcp_server/caller_sign_in.py create mode 100644 tests/test_litellm/proxy/_experimental/mcp_server/test_caller_sign_in.py diff --git a/litellm/proxy/_experimental/mcp_server/caller_sign_in.py b/litellm/proxy/_experimental/mcp_server/caller_sign_in.py new file mode 100644 index 00000000000..9d9b8b7297e --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/caller_sign_in.py @@ -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) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 7a0f59c3c2b..6940e8d4f02 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -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) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 09714b65f86..c38d156dd46 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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. diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index 891c37f4fe7..abc7c52f5a7 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/token_exchange_provider.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/token_exchange_provider.py index 2501f751d8e..7f5c0c99145 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/token_exchange_provider.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/token_exchange_provider.py @@ -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: diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 2412e83b9d9..b73d183db86 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -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, ) diff --git a/litellm/proxy/_experimental/mcp_server/utils.py b/litellm/proxy/_experimental/mcp_server/utils.py index 7411dc5c4f0..8234a62c4ef 100644 --- a/litellm/proxy/_experimental/mcp_server/utils.py +++ b/litellm/proxy/_experimental/mcp_server/utils.py @@ -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. diff --git a/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py b/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py index 94de0b2df7b..e5397a7269d 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py +++ b/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py @@ -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 diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index 91ae95eff48..b7dd62559f0 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_token_exchange_provider.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_token_exchange_provider.py index cf15fdb3e26..002e70da9d2 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_token_exchange_provider.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_token_exchange_provider.py @@ -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" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_caller_sign_in.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_caller_sign_in.py new file mode 100644 index 00000000000..aa853696462 --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_caller_sign_in.py @@ -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] diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 4e27ec134d4..7349db67d5d 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 070d87eb012..8d2642a66f1 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -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"' + ) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index f361b8bfd09..97b455bfd39 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -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", diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py index e72b716665c..cfd6745308d 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py @@ -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="blocked by egress proxy")]) - 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="outage")]) - 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="waf page")]) + handler: Final = FakeHandler([_response(200, text="waf page")]) 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="waf page")]) + handler: Final = FakeHandler([_response(200, text="waf page")]) 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="waf block page")]) - 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")