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..d948606dcf1 --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/caller_sign_in.py @@ -0,0 +1,196 @@ +"""Caller-side sign-in requirements for MCP connects. + +A guardrail that evaluates tool calls in the caller's own identity (an On-Behalf-Of exchange of the caller's +bearer) needs the caller signed in with its issuer before the first tool call, and a tool call's JSON-RPC +error cannot carry ``WWW-Authenticate``. ``token_exchange`` (OBO) servers have the same need: the caller must +present a subject token the gateway can exchange. Both cases share one contract: a connect that carries no +usable subject answers 401 with the RFC 9728 challenge, and the protected-resource metadata advertises the +issuers and scopes the caller signs in for. + +Guardrails implement :class:`CallerSignInProvider`; :func:`caller_sign_in_for` merges the OBO server's own +requirement with every registered provider's so the challenge and the metadata always agree. The MCP package +never imports a concrete provider. +""" + +from __future__ import annotations + +import itertools +from collections.abc import Awaitable, Callable, Mapping +from dataclasses import dataclass +from typing import TYPE_CHECKING, Final, Protocol, runtime_checkable + +from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError +from typing_extensions import assert_never + +import litellm +from litellm.integrations.custom_guardrail import CustomGuardrail +from litellm.types.mcp import MCPAuth + +if TYPE_CHECKING: + from litellm.proxy._types import UserAPIKeyAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + +@dataclass(frozen=True, slots=True) +class CallerSignIn: + """The issuers a caller signs in with and the scopes it requests before calling a server.""" + + issuers: tuple[str, ...] + scopes: tuple[str, ...] + + +@dataclass(frozen=True, slots=True) +class SignedIn: + """The subject token the caller presented satisfies this provider's sign-in.""" + + +@dataclass(frozen=True, slots=True) +class Rejected: + """The caller's identity provider rejected the presented subject token.""" + + detail: str + claims: str | None = None + + +@dataclass(frozen=True, slots=True) +class Unavailable: + """The provider could not reach a verdict; ``fail_open`` is the provider's own fallback policy.""" + + detail: str + fail_open: bool + + +CallerSignInPreflight = SignedIn | Rejected | Unavailable + + +@runtime_checkable +class CallerSignInProvider(Protocol): + def caller_sign_in(self, server: MCPServer, user_api_key_auth: UserAPIKeyAuth | None) -> CallerSignIn | None: + """The sign-in this provider requires of callers hitting ``server``; ``None`` when it does not gate + the server for this caller (``user_api_key_auth=None`` is the anonymous metadata fetch that follows + a challenge).""" + ... + + async def preflight_caller_sign_in( + self, server: MCPServer, user_api_key_auth: UserAPIKeyAuth | None, subject_token: str + ) -> CallerSignInPreflight: + """Validate ``subject_token`` against this provider at connect time, where a challenge's + ``WWW-Authenticate`` still reaches the client.""" + ... + + +class _JwtIssuerEntry(BaseModel): + model_config = ConfigDict(extra="ignore") + + issuer: str | None = None + + +class _JwtAuthConfig(BaseModel): + model_config = ConfigDict(extra="ignore") + + issuers: list[_JwtIssuerEntry] = [] # mutable-ok: pydantic copies the default per instance + + +_JWT_AUTH_ADAPTER: Final = TypeAdapter(_JwtAuthConfig) + + +def _jwt_auth_issuer_entries(jwtauth: object) -> tuple[_JwtIssuerEntry, ...]: + try: + return tuple(_JWT_AUTH_ADAPTER.validate_python(jwtauth, from_attributes=True).issuers) + except ValidationError: + return () + + +def _providers() -> tuple[CallerSignInProvider, ...]: + return tuple( + callback + for callback in litellm.logging_callback_manager.get_custom_loggers_for_type(CustomGuardrail) + if isinstance(callback, CallerSignInProvider) + ) + + +def jwt_auth_issuers() -> tuple[str, ...]: + """The OAuth issuer identifier(s) LiteLLM's JWT auth trusts, for the OBO PRM authorization_servers. + + In token_exchange the IdP that issues the subject JWT is the same one LiteLLM validates it + against, so OBO discovery points clients at the JWT-auth issuer to obtain a subject token. + Sourced from ``JWT_ISSUER`` and any configured ``litellm_jwtauth.issuers``. + """ + import os # noqa: PLC0415 + + from litellm.proxy.proxy_server import ( # noqa: PLC0415 # lazy: proxy_server pulls the whole proxy graph + general_settings, # pyright: ignore[reportUnknownVariableType] # proxy_server.general_settings is a raw untyped dict + ) + + env_issuer: Final = os.getenv("JWT_ISSUER") + env: Final[tuple[str, ...]] = (env_issuer,) if env_issuer else () + + settings: Final[Mapping[str, object]] = general_settings if isinstance(general_settings, Mapping) else {} + configured: Final = tuple( + entry.issuer for entry in _jwt_auth_issuer_entries(settings.get("litellm_jwtauth")) if entry.issuer + ) + return tuple(dict.fromkeys((*env, *configured))) + + +def caller_sign_in_for(server: MCPServer, user_api_key_auth: UserAPIKeyAuth | None) -> CallerSignIn | None: + """The merged sign-in requirement for ``server``: the OBO server's own issuer/scopes plus every + registered provider's contribution. ``None`` when nothing requires sign-in, which is also the gate the + connect-time challenge branches on.""" + contributions: Final = tuple( + contribution + for contribution in ( + *( + (CallerSignIn(issuers=jwt_auth_issuers(), scopes=tuple(server.scopes or ())),) + if server.auth_type == MCPAuth.oauth2_token_exchange + else () + ), + *(provider.caller_sign_in(server, user_api_key_auth) for provider in _providers()), + ) + if contribution is not None + ) + if not contributions: + return None + issuers: Final = tuple(dict.fromkeys(itertools.chain.from_iterable(c.issuers for c in contributions))) + scopes: Final = tuple(dict.fromkeys(itertools.chain.from_iterable(c.scopes for c in contributions))) + return CallerSignIn(issuers=issuers, scopes=scopes) + + +async def preflight_caller_sign_in( + server: MCPServer, + user_api_key_auth: UserAPIKeyAuth | None, + subject_token: str, + *, + root_path: str, + resource_metadata: str | None, + connecting: Callable[[], Awaitable[bool]], +) -> None: + """Run every provider's connect-time check against the subject token while ``connecting``, so a bearer + the IdP will reject surfaces as a challenge here rather than a JSON-RPC error at the first tool call, and + a fail-closed provider outage is the connect's 503. On an open session nothing is exchanged here: the + tool-call hook runs the one exchange and answers inside the JSON-RPC envelope, with its guardrail Logs + row.""" + from fastapi import HTTPException # noqa: PLC0415 # lazy: fastapi import stays off the cold path + + from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import ( # noqa: PLC0415 # lazy: adapter pulls MCP subgraph + raise_token_exchange_challenge, + ) + + gating: Final = tuple( + provider for provider in _providers() if provider.caller_sign_in(server, user_api_key_auth) is not None + ) + if not gating or not await connecting(): + return + for provider in gating: + match await provider.preflight_caller_sign_in(server, user_api_key_auth, subject_token): + case SignedIn(): + continue + case Rejected(detail=_, claims=claims): + raise_token_exchange_challenge( + server, root_path=root_path, claims=claims, resource_metadata=resource_metadata + ) + case Unavailable(fail_open=True): + continue + case Unavailable(detail=detail, fail_open=False): + raise HTTPException(status_code=503, detail=detail) + case _ as verdict: + assert_never(verdict) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 7a0f59c3c2b..9e262432ade 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, @@ -527,10 +528,7 @@ def _resolve_mcp_server_by_name_or_id(lookup: str, client_ip: str | None) -> MCP global_mcp_server_manager, ) - by_name: Final = global_mcp_server_manager.get_mcp_server_by_name(lookup, client_ip=client_ip) - if by_name is not None: - return by_name - return global_mcp_server_manager.get_mcp_server_by_id(lookup, client_ip=client_ip) + return global_mcp_server_manager.get_mcp_server_answering_to(lookup, client_ip=client_ip) def _resolve_oauth2_server_for_root_endpoints( @@ -2580,9 +2578,9 @@ async def _build_oauth_protected_resource_response( detail=(f"Upstream oauth-protected-resource metadata unavailable for MCP server {mcp_server.name!r}"), ) - obo_response: Final = _obo_protected_resource_response(mcp_server, resource_url) - if obo_response is not None: - return obo_response + sign_in_response: Final = _caller_sign_in_protected_resource_response(mcp_server, resource_url) + if sign_in_response is not None: + return sign_in_response if mcp_server is not None and mcp_server.advertises_gateway_authorization_server: return { @@ -2603,51 +2601,30 @@ async def _build_oauth_protected_resource_response( } -def _obo_protected_resource_response(mcp_server: MCPServer | None, resource_url: str) -> dict | None: - """The OBO (token_exchange) PRM, or None when this server is not OBO / no issuer is configured. +def _caller_sign_in_protected_resource_response( + mcp_server: MCPServer | None, resource_url: str +) -> dict[str, object] | None: + """The caller sign-in PRM: the OBO issuer(s) LiteLLM trusts merged with every registered + ``CallerSignInProvider``'s contribution, or None when no sign-in gates this server. - The client SSOs with the IdP to obtain a subject token, which LiteLLM then exchanges, so discovery - points at the JWT-auth issuer(s) LiteLLM trusts (the same IdP that issues and validates the - subject), not the gateway. None falls the caller back to the gateway default so discovery still - returns metadata; it just can't name the IdP. + The client SSOs with the IdP to obtain a subject token, which LiteLLM then exchanges (or a + guardrail consumes directly), so discovery points at the issuer(s), not the gateway. None falls + the caller back to the gateway default so discovery still returns metadata; it just can't name + the IdP. The anonymous metadata fetch passes ``user_api_key_auth=None`` because it cannot see + which key selected a provider. """ - if mcp_server is None or mcp_server.auth_type != MCPAuth.oauth2_token_exchange: + if mcp_server is None: return None - issuers: Final = _jwt_auth_issuers() - if not issuers: + sign_in: Final = caller_sign_in_for(mcp_server, None) + if sign_in is None or not sign_in.issuers: return None return { - "authorization_servers": issuers, + "authorization_servers": sign_in.issuers, "resource": resource_url, - "scopes_supported": (mcp_server.scopes if mcp_server.scopes else []), + "scopes_supported": sign_in.scopes, } -def _jwt_auth_issuers() -> list: - """The OAuth issuer identifier(s) LiteLLM's JWT auth trusts, for the OBO PRM authorization_servers. - - In token_exchange the IdP that issues the subject JWT is the same one LiteLLM validates it - against, so OBO discovery points clients at the JWT-auth issuer to obtain a subject token. - Sourced from ``JWT_ISSUER`` and any configured ``litellm_jwtauth.issuers``. - """ - import os # noqa: PLC0415 - - from litellm.proxy.proxy_server import general_settings # noqa: PLC0415 - - issuers: Final[list] = [] - env_issuer: Final = os.getenv("JWT_ISSUER") - if env_issuer: - issuers.append(env_issuer) - - jwtauth: Final = general_settings.get("litellm_jwtauth") if isinstance(general_settings, Mapping) else None - raw_issuers: Final = jwtauth.get("issuers") if isinstance(jwtauth, dict) else getattr(jwtauth, "issuers", None) - for cfg in raw_issuers or []: - issuer = cfg.get("issuer") if isinstance(cfg, dict) else getattr(cfg, "issuer", None) - if issuer and issuer not in issuers: - issuers.append(issuer) - return issuers - - @router.get("/.well-known/oauth-protected-resource") def oauth_protected_resource_root(request: Request) -> dict[str, str | tuple[str, ...]]: request_base_url: Final = get_request_base_url(request) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index c5102c303af..f1fe97ff72b 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -13,6 +13,7 @@ import json import math import os import re +import secrets import time from collections.abc import ( AsyncIterator, @@ -169,6 +170,7 @@ from litellm.proxy._experimental.mcp_server.utils import ( normalize_server_name, openapi_tool_name, parse_admin_env_vars, + server_answers_to_name, strip_known_server_prefix, validate_mcp_server_name, ) @@ -1163,6 +1165,10 @@ def _raw_header_value(raw_headers: Mapping[str, str] | None, name: str) -> str | return next((v for k, v in (raw_headers or {}).items() if isinstance(k, str) and k.lower() == name), None) +def _is_master_key(bearer: str, master_key: str | None) -> bool: + return bool(master_key) and secrets.compare_digest(bearer.encode(), (master_key or "").encode()) + + def _has_explicit_litellm_admission_header(raw_headers: Mapping[str, str] | None) -> bool: """Admission only consumes a non-empty ``x-litellm-api-key``; an empty one falls back to ``Authorization``.""" return bool(_raw_header_value(raw_headers, "x-litellm-api-key")) @@ -3897,6 +3903,31 @@ class MCPServerManager: return None return bearer + @staticmethod + def _caller_sign_in_subject_token( + oauth2_headers: Mapping[str, str] | None, + raw_headers: Mapping[str, str] | None, + ) -> str | None: + """The ``Bearer`` credential a caller sign-in provider validates. An admission that consumed + ``Authorization`` (custom auth, built-in OAuth2, JWT) did so on the caller's own IdP token, so that token + is the subject; any other scheme, a LiteLLM key (virtual or master) and a bearer repeating + ``x-litellm-api-key`` are withheld.""" + from litellm.proxy.proxy_server import master_key # noqa: PLC0415 # circular import + + authorization: Final = (oauth2_headers or {}).get("Authorization") or _raw_header_value( + raw_headers, "authorization" + ) + scheme_and_credential: Final = (authorization or "").split(None, 1) + if len(scheme_and_credential) != 2 or scheme_and_credential[0].lower() != "bearer": + return None + bearer: Final = scheme_and_credential[1] + if bearer.startswith(LITELLM_VIRTUAL_KEY_PREFIX) or _is_master_key(bearer, master_key): + return None + admission_header: Final = _raw_header_value(raw_headers, "x-litellm-api-key") + if admission_header and strip_auth_scheme(admission_header, "Bearer") == bearer: + return None + return bearer + def _obo_subject_token( self, server: MCPServer, @@ -4109,6 +4140,7 @@ class MCPServerManager: oauth2_headers: dict[str, str] | None, user_api_key_auth: UserAPIKeyAuth | None, raw_headers: Mapping[str, str] | None = None, + resource_metadata: str | None = None, ) -> None: """Mint an exchange-backed server's upstream credential at the transport edge. @@ -4140,13 +4172,24 @@ class MCPServerManager: if subject_token is not None: return case _: + from litellm.proxy._experimental.mcp_server.caller_sign_in import ( # noqa: PLC0415 # lazy: provider discovery pulls the guardrail registry + caller_sign_in_for, + ) + + sign_in_subject: Final = self._caller_sign_in_subject_token(oauth2_headers, raw_headers) + if sign_in_subject is None and caller_sign_in_for(server, user_api_key_auth) is not None: + raise_token_exchange_challenge( + server, root_path=get_request_root_path(), resource_metadata=resource_metadata + ) return resolved_server: Final = await self.ensure_oauth_metadata_discovered(server) spec: Final = to_server_spec_fail_closed(resolved_server) if spec is None or not isinstance(spec.config, (TokenExchangeConfig, IdJagConfig)): return if subject_token is None and isinstance(spec.config, TokenExchangeConfig): - raise_token_exchange_challenge(resolved_server, root_path=get_request_root_path()) + raise_token_exchange_challenge( + resolved_server, root_path=get_request_root_path(), resource_metadata=resource_metadata + ) match await self._cred_provider.resolve_credentials(to_subject(user_api_key_auth, subject_token), spec): case Ok(_): return @@ -4156,6 +4199,7 @@ class MCPServerManager: resolved_server, root_path=get_request_root_path(), claims=err.unauthorized.claims, + resource_metadata=resource_metadata, ) raise_public(err) @@ -5747,13 +5791,11 @@ class MCPServerManager: if proxy_logging_obj is None: return hook_result - # Extract incoming Bearer token from raw request headers so - # guardrails like MCPJWTSigner can verify + re-sign it (FR-5). - normalized_raw: Final = {k.lower(): v for k, v in (raw_headers or {}).items()} - incoming_bearer_token: str | None = None - auth_hdr: Final = normalized_raw.get("authorization", "") - if auth_hdr.lower().startswith("bearer "): - incoming_bearer_token = auth_hdr[len("bearer ") :] + inbound_authorization: Final = _raw_header_value(raw_headers, "authorization") or "" + incoming_bearer_token: Final = ( + inbound_authorization[len("bearer ") :] if inbound_authorization.lower().startswith("bearer ") else None + ) + incoming_subject_token: Final = self._caller_sign_in_subject_token(None, raw_headers) pre_hook_kwargs: Final = { "guardrail_context": guardrail_context, @@ -5769,6 +5811,7 @@ class MCPServerManager: ), "user_api_key_hash": (getattr(user_api_key_auth, "api_key_hash", None) if user_api_key_auth else None), "incoming_bearer_token": incoming_bearer_token, + "incoming_subject_token": incoming_subject_token, "headers": logging_safe_mcp_headers(raw_headers), "tool_description": tool.description if tool is not None else None, "tool_input_schema": tool.input_schema if tool is not None else None, @@ -6981,6 +7024,70 @@ class MCPServerManager: return server return None + def get_mcp_server_answering_to( + self, name: str, client_ip: str | None = None, *, among: Sequence[MCPServer] | None = None + ) -> MCPServer | None: + """The one server a ``/mcp/{name}`` segment denotes, shared by the connect preflight, the scoped + router, and RFC 9728 discovery so all three name the same server: the exact ``get_mcp_server_by_name`` + priority first, then the exact ``server_id``, then the name priority case-insensitively, then any prefix + form routing accepts. A name that denotes a server hidden from ``client_ip`` resolves to ``None`` at the + pass that found it: it never falls through to a looser pass that could name another server. ``among`` + runs the same passes over those servers alone instead of the registry, which is how the scoped router + picks the caller's granted server answering to ``name``.""" + if among is not None: + return self._server_among_answering_to(name, tuple(among), client_ip) + exact: Final = self.get_mcp_server_by_name(name) + if exact is not None: + return exact if self._is_server_accessible_from_ip(exact, client_ip) else None + by_id: Final = self.get_mcp_server_by_id(name) + if by_id is not None: + return by_id if self._is_server_accessible_from_ip(by_id, client_ip) else None + requested: Final = name.lower() + servers: Final = tuple(self.get_registry().values()) + identifiers: Final[tuple[Callable[[MCPServer], str | None], ...]] = ( + lambda server: server.alias, + lambda server: server.server_name, + lambda server: server.name, + ) + for identifier in identifiers: + if (found := next((s for s in servers if (identifier(s) or "").lower() == requested), None)) is not None: + return found if self._is_server_accessible_from_ip(found, client_ip) else None + return next( + ( + server + for server in self.get_filtered_registry(client_ip).values() + if server_answers_to_name(server, name) + ), + None, + ) + + def _server_among_answering_to( + self, name: str, servers: Sequence[MCPServer], client_ip: str | None + ) -> MCPServer | None: + """``get_mcp_server_answering_to`` over ``servers`` instead of the registry: the same passes in the + same order, with a server hidden from ``client_ip`` resolving to ``None`` at the pass that found it.""" + requested: Final = name.lower() + passes: Final[tuple[Callable[[MCPServer], bool], ...]] = ( + lambda server: server.alias == name, + lambda server: server.server_name == name, + lambda server: server.name == name, + lambda server: server.server_id == name, + lambda server: (server.alias or "").lower() == requested, + lambda server: (server.server_name or "").lower() == requested, + lambda server: (server.name or "").lower() == requested, + ) + for matches in passes: + if (found := next((server for server in servers if matches(server)), None)) is not None: + return found if self._is_server_accessible_from_ip(found, client_ip) else None + return next( + ( + server + for server in servers + if self._is_server_accessible_from_ip(server, client_ip) and server_answers_to_name(server, name) + ), + None, + ) + def get_filtered_registry(self, client_ip: str | None = None) -> dict[str, MCPServer]: """ Get registry filtered by client IP access control. diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index b0410e87103..7ae019840de 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, ) @@ -435,6 +436,7 @@ async def _dispatch_virtual_mcp_tool( async def _get_allowed_mcp_servers_from_mcp_server_names( mcp_servers: Sequence[str] | None, allowed_mcp_servers: list[MCPServer], + client_ip: str | None = None, ) -> list[MCPServer]: """ Get the filtered MCP servers from the MCP server names. @@ -451,15 +453,10 @@ async def _get_allowed_mcp_servers_from_mcp_server_names( # Filter servers based on mcp_servers parameter if provided if mcp_servers is not None: for server_or_group in mcp_servers: - server_name_matched = False - - for server in allowed_mcp_servers: - if server and _server_answers_to(server, server_or_group): - filtered_server[server.server_id] = server - server_name_matched = True - break - - if not server_name_matched: + scoped = _scoped_server(server_or_group, allowed_mcp_servers, client_ip) + if scoped is not None: + filtered_server[scoped.server_id] = scoped + else: try: access_group_server_ids = await MCPRequestHandler._get_mcp_servers_from_access_groups( [server_or_group] @@ -493,8 +490,20 @@ def _http_detail_message(detail: object) -> str: def _server_answers_to(server: MCPServer, name: str) -> bool: - requested: Final = name.lower() - return any(requested == known.lower() for known in iter_known_server_prefixes(server) if known) + return server_answers_to_name(server, name) + + +def _scoped_server(name: str, allowed_mcp_servers: Sequence[MCPServer], client_ip: str | None) -> MCPServer | None: + """The granted server a scoped ``name`` selects for the caller: the registry's own pass order run over + ``allowed_mcp_servers`` alone, so a granted server wins over an ungranted alias or case variant the registry + would pick. ``None`` when no granted server answers, or when the registry's own pick for ``name`` is a server + hidden from ``client_ip``; the caller then retries the name as an access group it holds.""" + if ( + global_mcp_server_manager.get_mcp_server_answering_to(name, client_ip=client_ip) is None + and global_mcp_server_manager.get_mcp_server_answering_to(name) is not None + ): + return None + return global_mcp_server_manager.get_mcp_server_answering_to(name, client_ip=client_ip, among=allowed_mcp_servers) async def raise_denied_scoped_mcp_access( @@ -675,6 +684,7 @@ async def _get_allowed_mcp_servers( allowed_mcp_servers = await _get_allowed_mcp_servers_from_mcp_server_names( mcp_servers=mcp_servers, allowed_mcp_servers=allowed_mcp_servers, + client_ip=client_ip, ) return allowed_mcp_servers @@ -1776,7 +1786,9 @@ def _challenge_missing_token_exchange_subject( The listing that fills a cold catalog absorbs the upstream 401 by design, so without this check a missing subject surfaces as an unknown-tool error instead of the challenge the warm path already raises. Gated to servers the key may reach so an unauthorized caller - learns nothing about the catalog. + learns nothing about the catalog. Guardrail-only sign-in is challenged at connect instead: a + tool call's JSON-RPC error drops ``WWW-Authenticate``, so the guardrail's own rejection is + the more useful answer there. """ if server is None or server.auth_type != MCPAuth.oauth2_token_exchange: return @@ -2439,6 +2451,7 @@ async def call_mcp_tool( allowed_mcp_servers = await _get_allowed_mcp_servers_from_mcp_server_names( mcp_servers=mcp_servers, allowed_mcp_servers=allowed_mcp_servers, + client_ip=client_ip, ) if mcp_servers and not allowed_mcp_servers: await raise_denied_scoped_mcp_access( diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/adapter.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/adapter.py index 42947e39530..abbde8f16f6 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/adapter.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/adapter.py @@ -348,6 +348,7 @@ def raise_token_exchange_challenge( *, root_path: str, claims: str | None = None, + resource_metadata: str | None = None, ) -> NoReturn: """Raise the RFC 9728 / RFC 6750 challenge an OBO (``token_exchange``) server returns when the caller's subject token is missing or the IdP rejected it. @@ -365,8 +366,12 @@ def raise_token_exchange_challenge( ``error="invalid_token"`` and is byte-identical to the static one. Both the error value (one of two literals) and the base64 claims draw from a fixed alphabet, so nothing from the IdP body reaches the header unescaped. + + ``resource_metadata`` is the absolute metadata URL of the route the client connected on (RFC 9728 + 5.1 names the parameter a URL, and the MCP SDK fetches it verbatim), supplied by the connect gate + that still holds the request; without it the challenge falls back to the alias's relative path. """ - resource_metadata: Final = oauth_protected_resource_path(root_path, server) + metadata_url: Final = resource_metadata or oauth_protected_resource_path(root_path, server) encoded_claims: Final = base64.b64encode(claims.encode()).decode() if claims else None error: Final = "insufficient_claims" if encoded_claims else "invalid_token" error_description: Final = ( @@ -376,7 +381,7 @@ def raise_token_exchange_challenge( ) www_authenticate: Final = ", ".join( ( - f'Bearer resource_metadata="{resource_metadata}"', + f'Bearer resource_metadata="{metadata_url}"', f'error="{error}"', f'error_description="{error_description}"', *((f'claims="{encoded_claims}"',) if encoded_claims else ()), 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..7ce84624847 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 @@ -9,6 +9,7 @@ call), so it needs no lazy wrapper. from __future__ import annotations +from dataclasses import dataclass from typing import Final import httpx @@ -24,6 +25,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_sto InMemoryTokenCacheBackend, ) from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchanger import ( + ExchangeHttpPost, OboTokenExchanger, SubjectTokenRejected, TokenExchangeClientError, @@ -34,11 +36,27 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchanger _GATEWAY_FAULT_OAUTH_ERRORS: Final = frozenset( {"invalid_client", "unauthorized_client", "unsupported_grant_type", "invalid_target", "invalid_scope"} ) +_INVALID_ASSERTION_AADSTS_PREFIX: Final = "50027" -def _oauth_error_fields(response: httpx.Response) -> tuple[str | None, str | None]: - """Read the RFC 6749 5.2 ``error`` code and the IdP's step-up ``claims`` blob from a - token-endpoint error body, as ``(error, claims)`` with None for whatever is absent. +@dataclass(frozen=True, slots=True) +class OAuthErrorBody: + error: str | None + claims: str | None + error_codes: tuple[str, ...] + + @property + def gateway_fault(self) -> str | None: + if self.error is None or self.error not in _GATEWAY_FAULT_OAUTH_ERRORS: + return None + if any(code.startswith(_INVALID_ASSERTION_AADSTS_PREFIX) for code in self.error_codes): + return None + return self.error + + +def oauth_error_fields(response: httpx.Response) -> OAuthErrorBody: + """Read the RFC 6749 5.2 ``error`` code, the IdP's step-up ``claims`` blob and Entra's + ``error_codes`` sub-codes from a token-endpoint error body, None or empty for whatever is absent. ``claims`` is the Entra Conditional Access / CAE challenge (a JSON string the client must replay to the IdP to satisfy the step-up); it is the caller's own requirement, not an IdP @@ -48,14 +66,18 @@ def _oauth_error_fields(response: httpx.Response) -> tuple[str | None, str | Non try: body: Final[object] = response.json() except Exception: # noqa: BLE001 - return None, None + return OAuthErrorBody(error=None, claims=None, error_codes=()) if not isinstance(body, dict): - return None, None + return OAuthErrorBody(error=None, claims=None, error_codes=()) code: Final = body.get("error") claims: Final = body.get("claims") - return ( - code if isinstance(code, str) else None, - claims if isinstance(claims, str) and claims else None, + raw_codes: Final = body.get("error_codes") + return OAuthErrorBody( + error=code if isinstance(code, str) else None, + claims=claims if isinstance(claims, str) and claims else None, + error_codes=tuple(str(c) for c in raw_codes if isinstance(c, (int, str))) + if isinstance(raw_codes, list) + else (), ) @@ -74,24 +96,32 @@ async def _post_exchange_endpoint( headers: Final = {"Accept": "application/json", **client_auth_headers} try: client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.MCP) # pyright: ignore - response: Final = await client.post(url, headers=headers, data=form) # pyright: ignore + response: Final = await client.post( # pyright: ignore[reportUnknownMemberType] # untyped handler + url, headers=headers, data=form + ) response.raise_for_status() # pyright: ignore parsed: Final[object] = response.json() # pyright: ignore except httpx.HTTPStatusError as status_err: status_code: Final = status_err.response.status_code + if status_code in (408, 429): + # Retry hints, not subject rejections: the IdP is shedding load, so a 401 would tell the + # caller to sign in again for nothing; surface it like a transport failure. + verbose_logger.warning("MCP token exchange throttled or timed out (HTTP %d)", status_code) + return None if 400 <= status_code < 500: - oauth_error, claims = _oauth_error_fields(status_err.response) - if oauth_error in _GATEWAY_FAULT_OAUTH_ERRORS: + oauth_error: Final = oauth_error_fields(status_err.response) + gateway_fault: Final = oauth_error.gateway_fault + if gateway_fault is not None: verbose_logger.warning( "MCP token exchange rejected as %s (HTTP %d); check the gateway client credentials, " "audience, and scope for this server", - oauth_error, + gateway_fault, status_code, ) - raise TokenExchangeClientError(oauth_error) from status_err + raise TokenExchangeClientError(gateway_fault) from status_err raise SubjectTokenRejected( f"IdP rejected the subject token (HTTP {status_code})", - claims=claims, + claims=oauth_error.claims, ) from status_err verbose_logger.warning("MCP token exchange request failed: %s", status_err) return None @@ -106,9 +136,9 @@ async def _post_exchange_endpoint( return parsed # pyright: ignore -def build_token_exchanger() -> OboTokenExchanger: +def build_token_exchanger(*, post: ExchangeHttpPost = _post_exchange_endpoint) -> OboTokenExchanger: return OboTokenExchanger( - _post_exchange_endpoint, + post, cache=InMemoryTokenCacheBackend(max_size=MCP_TOKEN_EXCHANGE_CACHE_MAX_SIZE), default_ttl_seconds=MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL, min_ttl_seconds=MCP_OAUTH2_TOKEN_CACHE_MIN_TTL, diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 2090a0c7421..ddbbe5c638f 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -8,12 +8,13 @@ import asyncio import contextlib import contextvars import hashlib +import itertools import json import os import time import types from collections import Counter -from collections.abc import AsyncGenerator, AsyncIterator, Callable, Iterable, Mapping, Sequence +from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Callable, Iterable, Iterator, Mapping, Sequence from typing import TYPE_CHECKING, Final, NoReturn, Protocol import httpx @@ -36,6 +37,7 @@ from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( MCPRequestHandler, _is_mcp_admitted_user_subject, ) +from litellm.proxy._experimental.mcp_server.caller_sign_in import caller_sign_in_for from litellm.proxy._experimental.mcp_server.client_allowlist import ( MCPClientAllowlist, check_mcp_client_allowed, @@ -62,6 +64,7 @@ from litellm.proxy._experimental.mcp_server.mcp_debug import ( ) from litellm.proxy._experimental.mcp_server.oauth_utils import ( _redact_mcp_resource_url, + get_passthrough_resource_metadata_url, get_passthrough_www_authenticate, get_route_relative_request_path, well_known_root_suffix, @@ -1424,6 +1427,41 @@ if MCP_AVAILABLE: return consumed_messages, b"".join(body_chunks) + class _ConnectBodyPeek: + """Reads a session-less ``POST`` body only once a gate asks whether it is ``initialize``, so a challenge + that needs no body still answers before the body arrives; consumed messages replay through ``receive``.""" + + def __init__(self, receive: Receive, peekable: bool) -> None: + self._receive: Final = receive + self._peekable: Final = peekable + self._body: bytes | None = None + self._replay: Iterator[Message] = iter(()) + + async def read(self) -> bytes: + messages, body = await _read_request_body_for_routing(self._receive) + self._replay = itertools.chain(self._replay, messages) + return body + + async def body(self) -> bytes: + if not self._peekable: + return b"" + if self._body is None: + self._body = await self.read() + return self._body + + async def connecting(self) -> bool: + return _is_initialize_request(await self.body()) + + async def receive(self) -> Message: + replayed: Final = next(self._replay, None) + return replayed if replayed is not None else await self._receive() + + def _known_connecting(value: bool) -> Callable[[], Awaitable[bool]]: + async def answer() -> bool: + return value + + return answer + async def _handle_stale_mcp_session( scope: Scope, receive: Receive, @@ -1606,6 +1644,7 @@ if MCP_AVAILABLE: mcp_server_auth_headers: dict[str, dict[str, str]] | None, user_api_key_auth: UserAPIKeyAuth | None, client_ip: str | None, + connecting: Callable[[], Awaitable[bool]], allowed_server_ids: set[str] | None = None, raw_headers: Mapping[str, str] | None = None, ) -> None: @@ -1621,7 +1660,28 @@ if MCP_AVAILABLE: a server it will be 403'd on immediately after authentication. """ for server_name in mcp_servers or []: - server = operations.global_mcp_server_manager.get_mcp_server_by_name(server_name, client_ip=client_ip) + registry_pick = operations.global_mcp_server_manager.get_mcp_server_answering_to( + server_name, client_ip=client_ip + ) + allowed_single = ( + await operations._get_allowed_mcp_servers( + user_api_key_auth=user_api_key_auth, mcp_servers=mcp_servers, client_ip=client_ip + ) + if registry_pick and mcp_servers is not None and len(mcp_servers) == 1 + else () + ) + granted = ( + operations.global_mcp_server_manager.get_mcp_server_answering_to( + server_name, client_ip=client_ip, among=allowed_single + ) + if allowed_single + else None + ) + server = granted if granted is not None else registry_pick + granted_single = granted is not None + obo_without_subject = ( + server is not None and server.auth_type == MCPAuth.oauth2_token_exchange and not oauth2_headers + ) if server is not None and allowed_server_ids is not None and server.server_id not in allowed_server_ids: # Caller's narrowed scope excludes this server — skip the # preemptive challenge and let downstream authorization @@ -1717,12 +1777,21 @@ if MCP_AVAILABLE: # reaches the token_exchange / pass-through blocks below. continue - # token_exchange (OBO): the caller supplied no subject token. Challenge at connect - # (transport level, where WWW-Authenticate survives) with the RFC 9728 resource_metadata - # so the client discovers the IdP, SSOs, and retries with a subject token, which LiteLLM - # then exchanges. A tool-call-time 401 would be wrapped into a JSON-RPC error and the - # header lost, so the discovery flow needs this pre-emptive challenge. - if server and server.auth_type == MCPAuth.oauth2_token_exchange and not oauth2_headers: + # Caller sign-in: challenge at connect because a tool-call-time 401 is wrapped into a + # JSON-RPC error and the WWW-Authenticate header is lost. OBO keeps its connect gate; + # guardrail-only gates fire only on a single-server connect the key's grant admits, so a + # key without access gets the grant's 403 instead of a sign-in it could not use. The one + # admission lookup above serves the challenge, the sign-in preflight and the exchange. + sign_in = caller_sign_in_for(server, user_api_key_auth) if server is not None else None + resource_metadata = get_passthrough_resource_metadata_url(scope, server_name) + subject_token = ( + operations.global_mcp_server_manager._caller_sign_in_subject_token( # pyright: ignore[reportPrivateUsage] # the manager owns the subject/admission filter shared with the preflight + oauth2_headers, raw_headers + ) + if server is not None + else None + ) + if server and sign_in is not None and subject_token is None and (obo_without_subject or granted_single): from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import ( # noqa: PLC0415 # lazy: adapter pulls MCP subgraph raise_token_exchange_challenge, ) @@ -1730,7 +1799,25 @@ if MCP_AVAILABLE: get_request_root_path, ) - raise_token_exchange_challenge(server, root_path=get_request_root_path()) + raise_token_exchange_challenge( + server, root_path=get_request_root_path(), resource_metadata=resource_metadata + ) + if server and sign_in is not None and subject_token is not None and granted_single: + from litellm.proxy._experimental.mcp_server.caller_sign_in import ( # noqa: PLC0415 # lazy: provider discovery pulls the guardrail registry + preflight_caller_sign_in, + ) + from litellm.proxy.middleware.per_request_root_path_middleware import ( # noqa: PLC0415 # lazy: middleware imports proxy utils + get_request_root_path, + ) + + await preflight_caller_sign_in( + server, + user_api_key_auth, + subject_token, + root_path=get_request_root_path(), + resource_metadata=resource_metadata, + connecting=connecting, + ) # Exchange-backed modes (token_exchange's OBO mint, id_jag's stored-assertion mint): run # the exchange here at the transport edge, so a rejected subject raises the RFC 9728 @@ -1739,22 +1826,13 @@ if MCP_AVAILABLE: # and what each mints from. Gated to single-server routes the key may reach; the # multi-server aggregate keeps absorbing per-server auth failures so one bad server # cannot 401 the whole connect. - if ( - server - and len(mcp_servers or []) == 1 - and server.server_id - in frozenset( - allowed.server_id - for allowed in await operations._get_allowed_mcp_servers( - user_api_key_auth=user_api_key_auth, mcp_servers=mcp_servers, client_ip=client_ip - ) - ) - ): + if server and granted_single: await operations.global_mcp_server_manager.preflight_token_exchange( server=server, oauth2_headers=oauth2_headers, user_api_key_auth=user_api_key_auth, raw_headers=raw_headers, + resource_metadata=resource_metadata, ) # Pass-through OAuth: when the admin has opted a server into @@ -2030,6 +2108,20 @@ if MCP_AVAILABLE: user_api_key_auth = await _apply_toolset_scope(user_api_key_auth, active_toolset_id) toolset_allowed_server_ids = await _toolset_server_ids(active_toolset_id) + named_session_id: Final = _get_session_id_from_scope(scope) + names_live_session: Final = ( + named_session_id is not None and named_session_id in _stateful_server_instances() + ) + request_owner: Final = _owner_fingerprint_for(user_api_key_auth, oauth2_headers, _client_ip) + expected_owner: Final = ( + _stateful_session_owners.get(named_session_id) if named_session_id is not None else None + ) + owner_mismatch: Final = expected_owner is not None and expected_owner != request_owner + connect_peek: Final = _ConnectBodyPeek( + receive, peekable=scope.get("method") == "POST" and not names_live_session and not owner_mismatch + ) + receive = connect_peek.receive + # https://datatracker.ietf.org/doc/html/rfc9728#name-www-authenticate-response # Must run after toolset scoping so the challenge set is derived # from the fully-authorized server set: a passthrough server that @@ -2042,6 +2134,7 @@ if MCP_AVAILABLE: mcp_server_auth_headers=mcp_server_auth_headers, user_api_key_auth=user_api_key_auth, client_ip=_client_ip, + connecting=connect_peek.connecting, allowed_server_ids=toolset_allowed_server_ids, raw_headers=raw_headers, ) @@ -2077,8 +2170,6 @@ if MCP_AVAILABLE: # - No session ID + initialize → stateful (so client gets mcp-session-id) # - No session ID + other → stateless (curl, Inspector, Notion) session_id = _get_session_id_from_scope(scope) - is_initialize = False - consumed_messages: list[Message] = [] # Owner-binding: a live stateful session may only be driven by the # caller that created it. Reject mismatches with 403 so a leaked @@ -2086,12 +2177,9 @@ if MCP_AVAILABLE: # # Run before ``_handle_stale_mcp_session`` so a non-owner cannot # force-clean another caller's residual tracking entries via a - # stale DELETE, and before peeking the request body so the 403 - # response sees a pristine ``receive`` channel. + # stale DELETE. if session_id: - expected_owner: Final = _stateful_session_owners.get(session_id) - request_owner = _owner_fingerprint_for(user_api_key_auth, oauth2_headers, _client_ip) - if expected_owner is not None and expected_owner != request_owner: + if owner_mismatch: verbose_logger.warning( "Rejecting MCP request: session '%s' owner mismatch.", session_id, @@ -2116,10 +2204,11 @@ if MCP_AVAILABLE: return session_id = _get_session_id_from_scope(scope) - body = b"" - if scope.get("method") == "POST": - consumed_messages, body = await _read_request_body_for_routing(receive) - is_initialize = _is_initialize_request(body) + session_body: Final = ( + await connect_peek.read() if scope.get("method") == "POST" and names_live_session else b"" + ) + body: Final = await connect_peek.body() or session_body + is_initialize: Final = _is_initialize_request(body) use_stateful: Final = bool(session_id or is_initialize) target_manager: Final = session_manager_stateful if use_stateful else session_manager_stateless @@ -2134,7 +2223,6 @@ if MCP_AVAILABLE: # session. Cap how many a single caller can hold so an authenticated # client cannot spam `initialize` and exhaust memory. if is_initialize and not session_id: - request_owner = _owner_fingerprint_for(user_api_key_auth, oauth2_headers, _client_ip) if not await _enforce_stateful_session_cap_for_owner(request_owner): verbose_logger.warning( "Rejecting MCP initialize: caller already holds the maximum number of active stateful sessions." @@ -2149,17 +2237,6 @@ if MCP_AVAILABLE: await too_many_response(scope, receive, send) return - # Replay body messages if we consumed them for peeking - original_receive: Final = receive - if consumed_messages: - - async def wrapped_receive(): - if consumed_messages: - return consumed_messages.pop(0) - return await original_receive() - - receive = wrapped_receive - # Serialize requests on the same stateful session so concurrent # callers don't clobber each other's auth context mid-flight. # @@ -2389,6 +2466,7 @@ if MCP_AVAILABLE: mcp_server_auth_headers=mcp_server_auth_headers, user_api_key_auth=user_api_key_auth, client_ip=_sse_client_ip, + connecting=_known_connecting(scope["method"] == "GET"), allowed_server_ids=toolset_allowed_server_ids, raw_headers=raw_headers, ) diff --git a/litellm/proxy/_experimental/mcp_server/utils.py b/litellm/proxy/_experimental/mcp_server/utils.py index 7c9d75457b5..b317f711414 100644 --- a/litellm/proxy/_experimental/mcp_server/utils.py +++ b/litellm/proxy/_experimental/mcp_server/utils.py @@ -362,6 +362,14 @@ def iter_known_server_prefixes(server: _McpServerLike) -> Iterator[str]: yield from _emit(server_id) +def server_answers_to_name(server: _McpServerLike, name: str) -> bool: + """Whether a scoped ``/mcp/{name}`` connect resolves to ``server``: case-insensitive over every prefix + form routing accepts (alias, server_name, server_id, short prefix), the same match + ``_server_answers_to`` applies when the router scopes a request.""" + requested: Final = name.lower() + return any(requested == known.lower() for known in iter_known_server_prefixes(server) if known) + + def iter_known_tool_name_spellings(tool_name: str, server: MCPServer) -> Iterator[str]: """Yield every name that denotes the bare ``tool_name`` on ``server``: the bare name, then its wire spelling under each prefix ``iter_known_server_prefixes`` accepts. diff --git a/litellm/proxy/guardrails/guardrail_hooks/agent_365/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/agent_365/__init__.py index 836d82eb851..29ed52d932b 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/agent_365/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/agent_365/__init__.py @@ -7,6 +7,7 @@ from .agent_365 import Agent365Guardrail if TYPE_CHECKING: from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler + from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchanger import TokenExchanger from litellm.types.guardrails import Guardrail, LitellmParams @@ -15,6 +16,7 @@ def initialize_guardrail( guardrail: "Guardrail", *, async_handler: "AsyncHTTPHandler | None" = None, + token_exchanger: "TokenExchanger | None" = None, ) -> Agent365Guardrail: import litellm from litellm.secret_managers.main import get_secret_str @@ -64,6 +66,7 @@ def initialize_guardrail( request_timeout=litellm_params.timeout if litellm_params.timeout is not None else 10.0, unreachable_fallback=litellm_params.unreachable_fallback, async_handler=async_handler, + token_exchanger=token_exchanger, event_hook=litellm_params.mode, default_on=litellm_params.default_on, ) 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 c683e1f2c5d..ed501b007cd 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,30 @@ from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, ) +from litellm.proxy._experimental.mcp_server.caller_sign_in import ( + CallerSignIn, + CallerSignInPreflight, + Rejected, + SignedIn, + Unavailable, +) +from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import OAuthToken +from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Error, Ok, Result +from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchange_provider import ( + build_token_exchanger, + oauth_error_fields, +) +from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchanger import ( + SubjectTokenRejected, + TokenExchanger, +) +from litellm.proxy._experimental.mcp_server.outbound_credentials.types import ( + CredError, + ServerSpec, + TokenExchangeConfig, +) from litellm.types.guardrails import GuardrailEventHooks +from litellm.types.mcp_server.mcp_server_manager import MCPServer from litellm.types.proxy.guardrails.guardrail_hooks.agent_365 import ( AGENT_365_PROD_API_BASE, AGENT_365_PROD_RESOURCE_APP_ID, @@ -49,43 +69,19 @@ if TYPE_CHECKING: from litellm.types.utils import GuardrailStatus TOKEN_ENDPOINT_TEMPLATE: Final = "https://login.microsoftonline.com/{tenant_id}/oauth2/v2.0/token" +ENTRA_ISSUER_TEMPLATE: Final = "https://login.microsoftonline.com/{tenant_id}/v2.0" EVALUATE_URL: Final = f"{AGENT_365_PROD_API_BASE}/agents/tool-evaluation/evaluate" OBO_SCOPE: Final = f"{AGENT_365_PROD_RESOURCE_APP_ID}/{AGENT_365_SCOPE_NAME}" MCP_SESSION_ID_HEADER: Final = "mcp-session-id" DEFENDER_STATUS_EVALUATED: Final = "Evaluated" -_GATEWAY_OWNED_TOKEN_ERRORS: Final = frozenset( - {"invalid_client", "unauthorized_client", "invalid_scope", "invalid_resource"} -) -# Entra reports a malformed or unverifiable assertion as ``invalid_client`` too; only its AADSTS50027xx -# (InvalidJwtToken) sub-codes tell that apart from a bad gateway secret. -_INVALID_ASSERTION_AADSTS_PREFIX: Final = "50027" -_AADSTS_CODES_ADAPTER: Final = TypeAdapter(tuple[int, ...]) +GATEWAY_SCOPE_TEMPLATE: Final = "api://{client_id}/access_as_user" _MCP_CALL_TYPES: Final[tuple[str, ...]] = ("mcp_call", "call_mcp_tool") -_TOOL_INPUT_SCHEMA_ADAPTER: Final = TypeAdapter(dict[str, object]) -_OBO_CACHE_MAX_ENTRIES: Final = 1000 -_DEFAULT_TOKEN_TTL_SECONDS: Final = 3599.0 -_TOKEN_EXPIRY_SLACK_SECONDS: Final = 60.0 - - -def _parse_expires_in(raw: object) -> float: - if not isinstance(raw, (int, float, str)): - return _DEFAULT_TOKEN_TTL_SECONDS - try: - return float(raw) - except ValueError: - return _DEFAULT_TOKEN_TTL_SECONDS - - -def _parse_aadsts_codes(raw: object) -> tuple[int, ...]: - try: - return _AADSTS_CODES_ADAPTER.validate_python(raw) - except ValidationError: - return () +_JSON_OBJECT_ADAPTER: Final = TypeAdapter(dict[str, object]) def _parse_tool_input_schema(raw: object) -> Mapping[str, object] | None: try: - return _TOOL_INPUT_SCHEMA_ADAPTER.validate_python(raw) + return _JSON_OBJECT_ADAPTER.validate_python(raw) except ValidationError: return None @@ -108,6 +104,15 @@ class _EvaluateResponse(TypedDict, total=False): correlationId: ReadOnly[str] +class _AdmissionMetadata(TypedDict): + user_api_key_metadata: ReadOnly[dict | None] + user_api_key_team_metadata: ReadOnly[dict | None] + + +class _AdmissionProbe(TypedDict): + metadata: ReadOnly[_AdmissionMetadata] + + class _ToolReference(BaseModel): model_config = ConfigDict(frozen=True) @@ -130,20 +135,12 @@ class _BlockedDetail(TypedDict): class Agent365TokenExchangeError(Exception): - def __init__(self, status_code: int, error_code: str, description: str, aadsts_codes: tuple[int, ...] = ()) -> None: - super().__init__(f"{error_code}: {description}") - self.status_code = status_code - self.error_code = error_code - self.description = description - self.aadsts_codes = aadsts_codes + """Entra refused the gateway's own client credentials, scope or resource; the caller cannot fix that by + signing in again, so it is the gateway's outage, never a 401.""" - @property - def gateway_owned(self) -> bool: - """Whether the gateway's own client credentials, scope or resource were refused, as opposed to the - caller's assertion. The caller cannot fix a gateway-owned rejection by signing in again.""" - if self.error_code not in _GATEWAY_OWNED_TOKEN_ERRORS: - return False - return not any(str(code).startswith(_INVALID_ASSERTION_AADSTS_PREFIX) for code in self.aadsts_codes) + def __init__(self, error_code: str) -> None: + super().__init__(error_code) + self.error_code = error_code class Agent365MalformedResponseError(Exception): @@ -156,6 +153,13 @@ class Agent365ThrottledError(Exception): self.status_code = status_code +def _gateway_fault_reason(error_code: str) -> str: + return ( + f"Entra rejected the gateway's own Agent 365 credentials ({error_code}); " + "check the guardrail's client_id and client_secret" + ) + + class Agent365Guardrail(CustomGuardrail): """Pre-MCP-call guardrail enforcing Microsoft Agent 365 tool-evaluation verdicts. @@ -173,6 +177,7 @@ class Agent365Guardrail(CustomGuardrail): request_timeout: float = 10.0, unreachable_fallback: Literal["fail_closed", "fail_open"] = "fail_closed", async_handler: AsyncHTTPHandler | None = None, + token_exchanger: TokenExchanger | None = None, **kwargs, # noqa: ANN003 # kwargs-ok: forwarded verbatim to CustomGuardrail (event_hook, default_on) ) -> None: super().__init__( @@ -192,8 +197,23 @@ class Agent365Guardrail(CustomGuardrail): self.async_handler = async_handler or get_async_httpx_client( llm_provider=httpxSpecialProvider.GuardrailCallback ) - self._obo_token_cache: OrderedDict[str, tuple[str, float]] = OrderedDict() # mutable-ok: lock-guarded LRU - self._obo_cache_lock = threading.Lock() + self._exchange_config: Final = TokenExchangeConfig( + profile="entra_obo", + token_exchange_endpoint=TOKEN_ENDPOINT_TEMPLATE.format(tenant_id=tenant_id), + client_id=client_id, + client_secret=SecretStr(client_secret), + scopes=(OBO_SCOPE,), + ) + self._exchange_server: Final = ServerSpec( + server_id=f"agent-365:{tenant_id}", + resource=AGENT_365_PROD_API_BASE, + config=self._exchange_config, + ) + self._token_exchanger: Final = ( + token_exchanger + if token_exchanger is not None + else build_token_exchanger(post=self._post_entra_token_endpoint) + ) verbose_proxy_logger.info("Initialized Microsoft Agent 365 guardrail: %s", guardrail_name) @staticmethod @@ -220,7 +240,7 @@ class Agent365Guardrail(CustomGuardrail): return data tool_name: Final = str(data.get("mcp_tool_name") or "") - assertion: Final = entra_assertion(data.get("incoming_bearer_token")) + assertion: Final = entra_assertion(data.get("incoming_subject_token")) if assertion is None: self._handle_caller_fault( data=data, @@ -233,22 +253,10 @@ class Agent365Guardrail(CustomGuardrail): ) try: - obo_token: Final = await self._get_obo_token(assertion) + exchange_result: Final = await self._exchange_caller_assertion(assertion) except Agent365TokenExchangeError as exc: - if exc.gateway_owned: - return self._handle_unavailable( - data=data, - tool_name=tool_name, - reason=( - f"Entra rejected the gateway's own Agent 365 credentials ({exc.error_code}); " - "check the guardrail's client_id and client_secret" - ), - ) - self._handle_caller_fault( - data=data, - tool_name=tool_name, - status_code=401, - reason=f"the Entra On-Behalf-Of token exchange was rejected ({exc.error_code})", + return self._handle_unavailable( + data=data, tool_name=tool_name, reason=_gateway_fault_reason(exc.error_code) ) except Agent365ThrottledError as exc: self._handle_throttled( @@ -264,11 +272,25 @@ class Agent365Guardrail(CustomGuardrail): reason=f"the Entra token endpoint could not be reached ({type(exc).__name__})", ) except Agent365MalformedResponseError as exc: - return self._handle_unavailable( - data=data, - tool_name=tool_name, - reason=str(exc), - ) + return self._handle_unavailable(data=data, tool_name=tool_name, reason=str(exc)) + match exchange_result: + case Ok(token): + obo_token: Final = token.access_token + case Error(error): + match error.tag: + case "unauthorized": + self._handle_caller_fault( + data=data, + tool_name=tool_name, + status_code=401, + reason=f"the Entra On-Behalf-Of token exchange was rejected ({error.unauthorized.detail})", + ) + case _: + return self._handle_unavailable( + data=data, + tool_name=tool_name, + reason=f"the Entra token exchange failed ({error.summary})", + ) start: Final = time.perf_counter() try: @@ -284,14 +306,14 @@ class Agent365Guardrail(CustomGuardrail): reason=f"the Agent 365 endpoint could not be reached ({type(exc).__name__})", ) latency_ms: Final = (time.perf_counter() - start) * 1000.0 - fallback: Final = self._handle_evaluate_error( + fallback: Final = await self._handle_evaluate_error( data=data, tool_name=tool_name, assertion=assertion, response=response, latency_ms=latency_ms ) if fallback is not None: return fallback return self._enforce_verdict(data=data, tool_name=tool_name, response=response, latency_ms=latency_ms) - def _handle_evaluate_error( + async def _handle_evaluate_error( self, data: dict, # mutable-ok: guardrail logging appends into the request metadata in place tool_name: str, @@ -308,7 +330,9 @@ class Agent365Guardrail(CustomGuardrail): ) if 400 <= response.status_code < 500: if response.status_code == 401: - self._evict_obo_token(assertion) + await self._token_exchanger.invalidate( + assertion, self._exchange_server, self._exchange_config, tenant_id=self.tenant_id + ) self._record_verdict( data=data, verdict="Rejected", @@ -455,26 +479,54 @@ class Agent365Guardrail(CustomGuardrail): return call_id return str(uuid.uuid4()) - async def _get_obo_token(self, assertion: str) -> str: - cache_key: Final = hashlib.sha256(assertion.encode("utf-8")).hexdigest() - now: Final = time.time() - with self._obo_cache_lock: - cached: Final = self._obo_token_cache.get(cache_key) - if cached and cached[1] > now + _TOKEN_EXPIRY_SLACK_SECONDS: - self._obo_token_cache.move_to_end(cache_key) - return cached[0] + def caller_sign_in(self, server: MCPServer, user_api_key_auth: "UserAPIKeyAuth | None") -> CallerSignIn | None: + """The Entra sign-in this guardrail requires of callers: only a ``default_on`` guardrail whose mode gates a + tagless MCP connect and that the caller's key or team has not opted out of, because the anonymous + metadata fetch that follows a challenge cannot see which key selected a guardrail and would advertise + the wrong issuer. Only servers that leave the caller's top-level ``Authorization`` with the gateway + qualify: a forwarded API-key header travels upstream in its own slot and does not displace the Entra + assertion.""" + if not (self.default_on and server.keeps_caller_authorization): + return None + probe: Final[_AdmissionProbe] = { + "metadata": { + "user_api_key_metadata": user_api_key_auth.metadata if user_api_key_auth else None, # pyright: ignore[reportUnknownMemberType] # UserAPIKeyAuth.metadata is a raw dict + "user_api_key_team_metadata": user_api_key_auth.team_metadata if user_api_key_auth else None, # pyright: ignore[reportUnknownMemberType] # UserAPIKeyAuth.team_metadata is a raw dict + } + } + if ( + self.should_run_guardrail( # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # should_run_guardrail takes an untyped data dict + data=probe, event_type=GuardrailEventHooks.pre_mcp_call + ) + is not True + ): + return None + return CallerSignIn( + issuers=(ENTRA_ISSUER_TEMPLATE.format(tenant_id=self.tenant_id),), + scopes=tuple(server.scopes) + if server.scopes + else (GATEWAY_SCOPE_TEMPLATE.format(client_id=self.client_id),), + ) + async def _exchange_caller_assertion(self, assertion: str) -> Result[OAuthToken, CredError]: + """The one Entra OBO exchange call both the tool-call path and the connect preflight run; a + successful result is cached by the exchanger, so the session reuses what the preflight minted.""" + return await self._token_exchanger.exchange( + assertion, self._exchange_server, self._exchange_config, tenant_id=self.tenant_id + ) + + async def _post_entra_token_endpoint( + self, + url: str, + form: dict[str, str], # mutable-ok: ExchangeHttpPost contract + client_auth_headers: dict[str, str], # mutable-ok: ExchangeHttpPost contract + ) -> dict[str, object] | None: + """The exchanger's HTTP edge for this guardrail: every way Entra can fail keeps its own exception, so + the verdict reason the Logs row carries names the OAuth error code or the transport fault.""" response: Final = await self._post_allowing_error_status( - url=TOKEN_ENDPOINT_TEMPLATE.format(tenant_id=self.tenant_id), - data={ - "grant_type": "urn:ietf:params:oauth:grant-type:jwt-bearer", - "client_id": self.client_id, - "client_secret": self.client_secret, - "assertion": assertion, - "scope": OBO_SCOPE, - "requested_token_use": "on_behalf_of", - }, - headers={"Content-Type": "application/x-www-form-urlencoded"}, + url=url, + data=form, + headers={"Content-Type": "application/x-www-form-urlencoded", **client_auth_headers}, ) if response.status_code in (408, 429): raise Agent365ThrottledError(status_code=response.status_code) @@ -485,32 +537,54 @@ class Agent365Guardrail(CustomGuardrail): response=response, ) try: - parsed_body: Final = response.json() + parsed_body: Final[object] = response.json() except ValueError as exc: raise Agent365MalformedResponseError("the Entra token endpoint returned a non-JSON body") from exc - if not isinstance(parsed_body, dict): - raise Agent365MalformedResponseError("the Entra token endpoint returned a non-object JSON body") - body: Final = parsed_body + try: + body: Final = _JSON_OBJECT_ADAPTER.validate_python(parsed_body) + except ValidationError as exc: + raise Agent365MalformedResponseError("the Entra token endpoint returned a non-object JSON body") from exc if response.status_code >= 400: - raise Agent365TokenExchangeError( - status_code=response.status_code, - error_code=str(body.get("error", "invalid_grant")), - description=str(body.get("error_description", ""))[:512], - aadsts_codes=_parse_aadsts_codes(body.get("error_codes")), - ) + oauth_error: Final = oauth_error_fields(response) + gateway_fault: Final = oauth_error.gateway_fault + if gateway_fault is not None: + raise Agent365TokenExchangeError(error_code=gateway_fault) + raise SubjectTokenRejected(oauth_error.error or "invalid_grant", claims=oauth_error.claims) if "access_token" not in body: raise Agent365MalformedResponseError("the Entra token endpoint returned no access_token") - raw_access_token: Final = body.get("access_token") - if not isinstance(raw_access_token, str) or not raw_access_token: + access_token: Final = body["access_token"] + if not isinstance(access_token, str) or not access_token: raise Agent365MalformedResponseError("the Entra token endpoint returned a non-string access_token") - access_token: Final = raw_access_token - expires_at: Final = time.time() + _parse_expires_in(body.get("expires_in", 3599)) - with self._obo_cache_lock: - self._obo_token_cache[cache_key] = (access_token, expires_at) - self._obo_token_cache.move_to_end(cache_key) - while len(self._obo_token_cache) > _OBO_CACHE_MAX_ENTRIES: - self._obo_token_cache.popitem(last=False) - return access_token + return body + + async def preflight_caller_sign_in( + self, server: MCPServer, user_api_key_auth: "UserAPIKeyAuth | None", subject_token: str + ) -> CallerSignInPreflight: + """The connect-time check the preemptive gate runs: a bearer Entra rejects, or one it could never + accept, gets the sign-in challenge here, where ``WWW-Authenticate`` still reaches the client, instead of + surfacing as a JSON-RPC error on every tools/call. ``subject_token=None`` stays the challenge gate's job.""" + assertion: Final = entra_assertion(subject_token) + if assertion is None: + return Rejected(detail="the caller's bearer is not an Entra token; sign in with Entra and retry") + fail_open: Final = self.unreachable_fallback == "fail_open" + try: + exchange_result: Final = await self._exchange_caller_assertion(assertion) + except Agent365TokenExchangeError as exc: + return Unavailable(detail=_gateway_fault_reason(exc.error_code), fail_open=fail_open) + except Agent365ThrottledError as exc: + return Unavailable(detail=f"the Entra token endpoint returned HTTP {exc.status_code}", fail_open=False) + except (httpx.HTTPError, LitellmTimeout, TimeoutError) as exc: + return Unavailable( + detail=f"the Entra token endpoint could not be reached ({type(exc).__name__})", fail_open=fail_open + ) + except Agent365MalformedResponseError as exc: + return Unavailable(detail=str(exc), fail_open=fail_open) + if isinstance(exchange_result, Ok): + return SignedIn() + error: Final = exchange_result.error + if error.tag == "unauthorized": + return Rejected(detail=error.unauthorized.detail, claims=error.unauthorized.claims) + return Unavailable(detail=f"the Entra token exchange failed ({error.summary})", fail_open=fail_open) async def _post_allowing_error_status( self, @@ -577,11 +651,6 @@ class Agent365Guardrail(CustomGuardrail): } raise HTTPException(status_code=503, detail=throttled_detail) - def _evict_obo_token(self, assertion: str) -> None: - cache_key: Final = hashlib.sha256(assertion.encode("utf-8")).hexdigest() - with self._obo_cache_lock: - self._obo_token_cache.pop(cache_key, None) - def _handle_unavailable( self, data: dict, # mutable-ok: guardrail logging appends into the request metadata in place diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index e447bcca703..d5cbc4bb40c 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -1518,6 +1518,7 @@ class ProxyLogging: # (e.g. MCPJWTSigner) to independently verify the caller's identity # before re-signing an outbound token (FR-5 verify+re-sign). "incoming_bearer_token": kwargs.get("incoming_bearer_token"), + "incoming_subject_token": kwargs.get("incoming_subject_token"), "metadata": synthetic_metadata, } user_api_key_auth: Final = kwargs.get("user_api_key_auth") diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index 2ee19b3e59a..e71ce56d3d6 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -315,10 +315,10 @@ class MCPServer(BaseModel): return self.per_server_oauth_discovery and self.auth_type == MCPAuth.oauth2 and not self.has_client_credentials @property - def advertises_gateway_authorization_server(self) -> bool: - """Whether named discovery should advertise the aggregate gateway authorization server.""" - if self.auth_type == MCPAuth.oauth2: - return self.is_gateway_managed_oauth2 and not self.uses_per_server_oauth_relay + def keeps_caller_authorization(self) -> bool: + """Whether the caller's top-level ``Authorization`` stays with the gateway: the server neither relays + it upstream nor runs an OAuth mode that fills that slot itself, so a gateway guardrail may consume it + as the caller's own assertion. Forwarding a separate API-key header leaves the slot untouched.""" if self.auth_type not in ( None, MCPAuth.none, @@ -328,11 +328,20 @@ class MCPServer(BaseModel): MCPAuth.authorization, MCPAuth.token, MCPAuth.aws_sigv4, + MCPAuth.oauth2_token_exchange, ): return False - return not any( - header.lower() in ("authorization", "x-api-key", "api-key", "apikey") - for header in (self.extra_headers or ()) + return not any(header.lower() == "authorization" for header in (self.extra_headers or ())) + + @property + def advertises_gateway_authorization_server(self) -> bool: + """Whether named discovery should advertise the aggregate gateway authorization server.""" + if self.auth_type == MCPAuth.oauth2: + return self.is_gateway_managed_oauth2 and not self.uses_per_server_oauth_relay + if self.auth_type == MCPAuth.oauth2_token_exchange: + return False + return self.keeps_caller_authorization and not any( + header.lower() in ("x-api-key", "api-key", "apikey") for header in (self.extra_headers or ()) ) @property diff --git a/tests/integration/mcp/test_mcp_access_matrix.py b/tests/integration/mcp/test_mcp_access_matrix.py index e458911cc68..e9cecffaed9 100644 --- a/tests/integration/mcp/test_mcp_access_matrix.py +++ b/tests/integration/mcp/test_mcp_access_matrix.py @@ -1,3 +1,4 @@ +import json import uuid from typing import Final @@ -92,6 +93,36 @@ def test_subject_grant_lists_only_reachable_tools_and_denies_the_rest( assert not any(name.startswith(denied_alias) for name in denied_listed.tools), denied_listed.tools +@pytest.mark.parametrize("entry", ("mcp", "server_mcp")) +def test_access_group_named_like_an_ungranted_server_still_routes_the_groups_servers( + gateway: Gateway, entry: EntryPoint +) -> None: + with peer_of("http") as shadow, peer_of("http") as member, gateway.scenario() as scenario: + group: Final = "docs" + uuid.uuid4().hex[:8] + member_alias: Final = "mem" + uuid.uuid4().hex[:8] + register_mcp(scenario, shadow, group) + register_mcp(scenario, member, member_alias, mcp_access_groups=[group]) + key: Final = scenario.key(object_permission={"mcp_access_groups": [group]}) + unmatched: Final = "none" + uuid.uuid4().hex[:8] + denied: Final = McpCaller( + gateway, key, entry, unmatched, {"x-mcp-servers": unmatched} if entry == "mcp" else {} + ).list_tools() + assert (denied.status, json.loads(denied.error or "null"), denied.tools) == ( + ( + 200, + {"code": -32600, "message": f"The key is not allowed to access the requested MCP servers: {unmatched}"}, + (), + ) + if entry == "mcp" + else (404, {"detail": f"MCP server, toolset, or access group '{unmatched}' not found"}, ()) + ), denied.raw + selection: Final = {"x-mcp-servers": group} if entry == "mcp" else {} + listed: Final = McpCaller(gateway, key, entry, group, selection).list_tools() + assert listed.ok, listed.raw + assert set(listed.tools) == {f"{member_alias}-{tool}" for tool in ("add", "multiply", "fail")}, listed.tools + assert tool_calls(shadow.drain()) == () + + @pytest.mark.parametrize("entry", ("mcp", "server_mcp", "rest")) def test_key_without_any_grant_sees_no_scoped_server(gateway: Gateway, entry: EntryPoint) -> None: with peer_of("http") as peer, gateway.scenario() as scenario: diff --git a/tests/integration/mcp/test_mcp_agent_365_guardrail.py b/tests/integration/mcp/test_mcp_agent_365_guardrail.py index ee746bf3dbf..6ed19426fd3 100644 --- a/tests/integration/mcp/test_mcp_agent_365_guardrail.py +++ b/tests/integration/mcp/test_mcp_agent_365_guardrail.py @@ -16,6 +16,7 @@ from integration._support.client import Gateway, eventually from integration._support.database import read_rows from integration._support.mcp import ( ENTRY_POINTS, + INITIALIZE, EntryPoint, McpCaller, McpPeer, @@ -38,6 +39,10 @@ GUARDRAIL_ROWS: Final = ( FALLBACKS: Final = (None, "fail_open", "fail_closed") +def _origin(gateway: Gateway) -> str: + return str(gateway.client.base_url).rstrip("/") + + def _generic_guardrail_outage(request: Request) -> Reply: assert request.target == "/beta/litellm_basic_guardrail_api", request.target return Reply(status=503, body=json.dumps({"error": "synthetic sibling guardrail outage"}).encode()) @@ -136,15 +141,26 @@ def test_a_missing_or_malformed_caller_bearer_blocks_on_every_entry_point_whatev ) -> None: with _rig(gateway, tmp_path, fallback) as rig: assert f"{rig.alias}-add" in rig.caller().list_tools().tools, "the catalog needs only the virtual key" + metadata_url: Final = f"{_origin(rig.candidate)}/.well-known/oauth-protected-resource/{rig.alias}/mcp" for entry in ENTRY_POINTS: - missing: Final = rig.caller(entry).call(f"{rig.alias}-add", {"entry": entry}, server_id=rig.server_id) - assert missing.error is not None and REJECTED in missing.raw, f"{entry} without a bearer: {missing.raw}" - malformed: Final = rig.caller(entry, "not-a-jws").call( - f"{rig.alias}-add", {"entry": entry}, server_id=rig.server_id - ) - assert malformed.error is not None and REJECTED in malformed.raw, f"{entry} opaque bearer: {malformed.raw}" + for label, bearer in (("without a bearer", None), ("opaque bearer", "not-a-jws")): + caller: Final = rig.caller(entry, bearer) + if entry == "server_mcp": + challenged: Final = caller.rpc("initialize", INITIALIZE) + assert challenged.status_code == 401, f"{entry} {label}: {challenged.text}" + authenticate: Final = challenged.headers.get("www-authenticate", "") + assert f'resource_metadata="{metadata_url}"' in authenticate, f"{entry} {label}: {authenticate!r}" + assert 'error="invalid_token"' in authenticate, f"{entry} {label}: {authenticate!r}" + if bearer is None: + unsigned: Final = caller.rpc( + "tools/call", {"name": f"{rig.alias}-add", "arguments": {"entry": entry}} + ) + assert unsigned.status_code == 401, f"{entry} {label}: {unsigned.text}" + continue + outcome: Final = caller.call(f"{rig.alias}-add", {"entry": entry}, server_id=rig.server_id) + assert outcome.error is not None and REJECTED in outcome.raw, f"{entry} {label}: {outcome.raw}" assert rig.upstream_tool_names() == () - expected: Final = 2 * len(ENTRY_POINTS) + expected: Final = 2 * len(ENTRY_POINTS) - 1 assert rig.guardrail_statuses("call_mcp_tool", expected) == ["guardrail_intervened"] * expected diff --git a/tests/integration/mcp/test_mcp_caller_sign_in.py b/tests/integration/mcp/test_mcp_caller_sign_in.py new file mode 100644 index 00000000000..ae01857cc91 --- /dev/null +++ b/tests/integration/mcp/test_mcp_caller_sign_in.py @@ -0,0 +1,463 @@ +import json +import uuid +from collections.abc import Mapping +from pathlib import Path +from types import MappingProxyType +from typing import Final + +import httpx +import yaml +from cryptography.hazmat.primitives.asymmetric import rsa +from integration._support.client import Gateway +from integration._support.mcp import ( + INITIALIZE, + McpCaller, + forget_mcp, + mcp_peer, + register_mcp, + tool_calls, +) +from integration._support.process import owned_proxy +from integration._support.wire import Reply, Request, wire_server +from jwt import algorithms as jwt_algorithms + +ADD: Final = {"a": 2, "b": 3} +ACCEPT: Final = {"Accept": "application/json, text/event-stream"} + + +def _rpc( + gateway: Gateway, + path: str, + key: str, + headers: dict[str, str], + method: str = "initialize", + params: Mapping[str, object] = INITIALIZE, +) -> httpx.Response: + return gateway.client.post( + path, + json={"jsonrpc": "2.0", "id": 1, "method": method, "params": dict(params)}, + headers={"x-litellm-api-key": key, **ACCEPT, **headers}, + ) + + +def _sse_data(response: httpx.Response) -> str: + return next(line[5:].strip() for line in response.text.splitlines() if line.startswith("data:")) + + +def _advertised(gateway: Gateway, segment: str) -> tuple[int, tuple[str, ...], object]: + response: Final = gateway.client.get(f"/.well-known/oauth-protected-resource/mcp/{segment}") + document: Final = response.json() + issuers: Final = tuple( + str(issuer).removesuffix(f"/{segment}") for issuer in document.get("authorization_servers", ()) + ) + return response.status_code, issuers, document.get("scopes_supported") + + +def _origin(gateway: Gateway) -> str: + return str(gateway.client.base_url).rstrip("/") + + +def _sign_in_config( + guardrail_params: dict[str, object], path: Path, general_settings: Mapping[str, object] = MappingProxyType({}) +) -> Path: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["guardrails"] = [{"guardrail_name": "signin" + uuid.uuid4().hex, "litellm_params": guardrail_params}] + config["general_settings"] = {**config.get("general_settings", {}), **general_settings} + path.write_text(yaml.safe_dump(config)) + return path + + +def test_multi_server_connect_with_a_litellm_key_in_authorization_admits_an_obo_server( + gateway: Gateway, +) -> None: + with mcp_peer() as obo_peer, mcp_peer() as math_peer, gateway.scenario() as scenario: + obo_alias: Final = "obo" + uuid.uuid4().hex[:8] + math_alias: Final = "math" + uuid.uuid4().hex[:8] + obo_id: Final = register_mcp( + scenario, + obo_peer, + obo_alias, + auth_type="oauth2_token_exchange", + token_exchange_endpoint="http://127.0.0.1:9/token", + credentials={"client_id": "obo-client", "client_secret": "obo-secret"}, + ) + math_id: Final = register_mcp(scenario, math_peer, math_alias) + key: Final = scenario.key(object_permission={"mcp_servers": [obo_id, math_id]}) + caller: Final = McpCaller( + gateway, + key, + "root", + headers={ + "x-mcp-servers": f"{obo_alias},{math_alias}", + "Authorization": f"Bearer {key}", + }, + ) + init: Final = caller.initialize() + assert init.ok, init.raw + listed: Final = caller.list_tools() + assert listed.ok, listed.raw + assert any(name.endswith("add") for name in listed.tools), listed.tools + + +def test_alias_first_lookup_wins_over_a_server_whose_name_matches_the_alias(gateway: Gateway, tmp_path: Path) -> None: + with mcp_peer() as math_peer, mcp_peer() as obo_peer: + name: Final = "gh" + uuid.uuid4().hex[:6] + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["mcp_servers"] = { + name: { + "transport": "http", + "url": obo_peer.url, + "auth_type": "oauth2_token_exchange", + "token_exchange_endpoint": "http://127.0.0.1:9/token", + "credentials": {"client_id": "obo-client", "client_secret": "obo-secret"}, + } + } + path: Final = tmp_path / "collision.yaml" + path.write_text(yaml.safe_dump(config)) + with ( + owned_proxy(gateway, tmp_path, {}, config=path) as candidate, + candidate.scenario() as scenario, + ): + second: Final = candidate.request( + "POST", + "/v1/mcp/server", + {"server_name": name + "_public", "alias": name, **math_peer.registration()}, + ) + assert second.status_code == 201, second.text + second_id: Final = second.json()["server_id"] + scenario.cleanups.callback(forget_mcp, candidate, second_id) + key: Final = scenario.key(object_permission={"mcp_servers": [second_id]}) + + init: Final = _rpc(candidate, f"/mcp/{name}", key, {}) + assert init.status_code == 200, init.text + listing: Final = candidate.client.post( + f"/mcp/{name}", + json={"jsonrpc": "2.0", "id": 2, "method": "tools/list", "params": {}}, + headers={"x-litellm-api-key": key, **ACCEPT}, + ) + assert listing.status_code == 200, listing.text + body: Final = json.loads( + next(line[5:].strip() for line in listing.text.splitlines() if line.startswith("data:")) + ) + names: Final = {tool["name"] for tool in body["result"]["tools"]} + assert any(name.endswith("add") for name in names), names + + +def test_exact_name_wins_over_a_case_folded_config_alias_for_connect_discovery_and_calls( + gateway: Gateway, tmp_path: Path +) -> None: + with mcp_peer() as obo_peer, mcp_peer() as math_peer: + stem: Final = "gh" + uuid.uuid4().hex[:6] + cased: Final = stem.capitalize() + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["mcp_servers"] = { + stem + "_obo": { + "alias": stem, + "transport": "http", + "url": obo_peer.url, + "auth_type": "oauth2_token_exchange", + "token_exchange_endpoint": "http://127.0.0.1:9/token", + "credentials": {"client_id": "obo-client", "client_secret": "obo-secret"}, + } + } + path: Final = tmp_path / "case-collision.yaml" + path.write_text(yaml.safe_dump(config)) + with owned_proxy(gateway, tmp_path, {}, config=path) as candidate, candidate.scenario() as scenario: + by_name: Final = register_mcp(scenario, math_peer, cased) + key: Final = scenario.key(object_permission={"mcp_servers": [by_name]}) + + init: Final = _rpc(candidate, f"/mcp/{cased}", key, {}) + assert init.status_code == 200, init.text + listed: Final = _rpc(candidate, f"/mcp/{cased}", key, {}, method="tools/list") + assert listed.status_code == 200, listed.text + tools: Final = json.loads(_sse_data(listed))["result"]["tools"] + add: Final = next(tool["name"] for tool in tools if tool["name"].endswith("add")) + called: Final = _rpc( + candidate, f"/mcp/{cased}", key, {}, method="tools/call", params={"name": add, "arguments": ADD} + ) + assert called.status_code == 200, called.text + assert json.loads(_sse_data(called))["result"]["content"][0]["text"] == "5", called.text + assert len(tool_calls(math_peer.drain())) == 1 + assert tool_calls(obo_peer.drain()) == () + assert _advertised(candidate, cased) == _advertised(candidate, by_name) + + same_name: Final = _rpc(candidate, f"/mcp/{stem}", key, {}) + assert same_name.status_code == 200, same_name.text + assert "www-authenticate" not in same_name.headers, same_name.headers + via_alias: Final = _rpc( + candidate, f"/mcp/{stem}", key, {}, method="tools/call", params={"name": add, "arguments": ADD} + ) + assert json.loads(_sse_data(via_alias))["result"]["content"][0]["text"] == "5", via_alias.text + assert len(tool_calls(math_peer.drain())) == 1 + assert tool_calls(obo_peer.drain()) == () + + challenged: Final = _rpc(candidate, f"/mcp/{stem}", scenario.key(), {}) + assert challenged.status_code == 401, challenged.text + assert ( + f'resource_metadata="{_origin(candidate)}/.well-known/oauth-protected-resource/mcp/{stem}"' + in challenged.headers.get("www-authenticate", "") + ) + assert _advertised(candidate, stem) == _advertised(candidate, stem + "_obo") + assert _advertised(candidate, stem) != _advertised(candidate, cased) + + +def test_name_of_a_server_hidden_from_an_external_ip_does_not_reroute_to_a_case_variant( + gateway: Gateway, tmp_path: Path +) -> None: + stem: Final = "gh" + uuid.uuid4().hex[:6] + cased: Final = stem.capitalize() + external: Final = {"X-Forwarded-For": "203.0.113.7"} + with mcp_peer() as hidden_peer, mcp_peer() as public_peer: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["general_settings"] = { + **config.get("general_settings", {}), + "use_x_forwarded_for": True, + "mcp_trusted_proxy_ranges": ["127.0.0.0/8"], + } + config["mcp_servers"] = { + stem: {"transport": "http", "url": hidden_peer.url, "available_on_public_internet": False} + } + path: Final = tmp_path / "external-ip.yaml" + path.write_text(yaml.safe_dump(config)) + with owned_proxy(gateway, tmp_path, {}, config=path) as candidate, candidate.scenario() as scenario: + public_id: Final = register_mcp(scenario, public_peer, cased) + key: Final = scenario.key(object_permission={"mcp_servers": [public_id]}) + + hidden_name: Final = _rpc(candidate, f"/mcp/{stem}", key, external) + assert hidden_name.status_code == 403, hidden_name.text + assert "www-authenticate" not in hidden_name.headers, hidden_name.headers + assert ( + candidate.client.get(f"/.well-known/oauth-protected-resource/mcp/{stem}", headers=external).status_code + == 404 + ) + assert _rpc(candidate, f"/mcp/{stem}", key, {}).status_code == 200 + + own_name: Final = _rpc(candidate, f"/mcp/{cased}", key, external) + assert own_name.status_code == 200, own_name.text + listed: Final = _rpc(candidate, f"/mcp/{cased}", key, external, method="tools/list") + add: Final = next( + tool["name"] + for tool in json.loads(_sse_data(listed))["result"]["tools"] + if tool["name"].endswith("add") + ) + called: Final = _rpc( + candidate, f"/mcp/{cased}", key, external, method="tools/call", params={"name": add, "arguments": ADD} + ) + assert called.status_code == 200, called.text + assert json.loads(_sse_data(called))["result"]["content"][0]["text"] == "5", called.text + assert len(tool_calls(public_peer.drain())) == 1 + assert tool_calls(hidden_peer.drain()) == () + + +def test_jwt_signer_verifies_the_bearer_that_admitted_the_call(gateway: Gateway, tmp_path: Path) -> None: + signer_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048) + jwk: Final = json.loads(jwt_algorithms.RSAAlgorithm.to_jwk(signer_key.public_key())) + jwk["kid"] = "idp" + holder: Final = [] + + def idp(request: Request) -> Reply: + issuer: Final = holder[0].url + if request.target.endswith("/.well-known/openid-configuration"): + return Reply(body=json.dumps({"issuer": issuer, "jwks_uri": issuer + "/jwks"}).encode()) + return Reply(body=json.dumps({"keys": [jwk]}).encode()) + + def introspect(request: Request) -> Reply: + return Reply(body=b'{"active": true, "sub": "subject-1"}') + + with ( + mcp_peer() as peer, + wire_server(idp) as verify_idp, + wire_server(introspect) as introspect_stub, + gateway.scenario() as scenario, + ): + holder.append(verify_idp) + alias: Final = "sig" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias) + config: Final = _sign_in_config( + { + "guardrail": "mcp_jwt_signer", + "mode": "pre_mcp_call", + "default_on": True, + "access_token_discovery_uri": verify_idp.url + "/.well-known/openid-configuration", + "token_introspection_endpoint": introspect_stub.url, + "required_claims": ["sub"], + }, + tmp_path / "signer.yaml", + ) + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as owned: + key: Final = owned.key(object_permission={"mcp_servers": [identity]}) + + def call(headers: Mapping[str, str]) -> httpx.Response: + return candidate.client.post( + "/mcp-rest/tools/call", + headers=dict(headers), + json={"name": "add", "arguments": ADD, "server_id": identity}, + ) + + admitted: Final = call({"Authorization": f"Bearer {key}"}) + assert admitted.status_code == 200, admitted.text + assert len(tool_calls(peer.drain())) == 1 + probes: Final = tuple(item for item in introspect_stub.drain() if item.body) + assert any(key.encode() in probe.body for probe in probes), probes + + split: Final = call({"x-litellm-api-key": key}) + assert split.status_code == 403, split.text + assert introspect_stub.drain() == () + + +AGENT_365_PARAMS: Final = { + "guardrail": "agent_365", + "mode": "pre_mcp_call", + "default_on": True, + "tenant_id": "00000000-0000-0000-0000-000000000000", + "client_id": "22222222-2222-2222-2222-222222222222", + "client_secret": "secret", +} +ENTRA_ISSUER: Final = "https://login.microsoftonline.com/00000000-0000-0000-0000-000000000000/v2.0" +GATEWAY_SCOPE: Final = "api://22222222-2222-2222-2222-222222222222/access_as_user" + + +def test_agent_365_gated_server_challenges_at_connect_and_advertises_entra(gateway: Gateway, tmp_path: Path) -> None: + config: Final = _sign_in_config(dict(AGENT_365_PARAMS), tmp_path / "agent365.yaml") + with ( + owned_proxy(gateway, tmp_path, {}, config=config) as candidate, + mcp_peer() as peer, + candidate.scenario() as scenario, + ): + alias: Final = "a365" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias) + granted: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + denied: Final = scenario.key(object_permission={"mcp_servers": ["no-mcp-servers"]}) + + challenged: Final = _rpc(candidate, f"/mcp/{alias}", granted, {}) + assert challenged.status_code == 401, challenged.text + authenticate: Final = challenged.headers.get("www-authenticate", "") + assert ( + f'resource_metadata="{_origin(candidate)}/.well-known/oauth-protected-resource/mcp/{alias}"' in authenticate + ) + assert 'error="invalid_token"' in authenticate + + opaque: Final = _rpc(candidate, f"/mcp/{alias}", granted, {"Authorization": "Bearer not-a-jws"}) + assert opaque.status_code == 401, opaque.text + assert ( + f'resource_metadata="{_origin(candidate)}/.well-known/oauth-protected-resource/mcp/{alias}"' + in opaque.headers.get("www-authenticate", "") + ) + + discovery: Final = candidate.client.get(f"/.well-known/oauth-protected-resource/mcp/{alias}") + assert discovery.status_code == 200, discovery.text + document: Final = discovery.json() + assert document["authorization_servers"] == [ENTRA_ISSUER] + assert document["scopes_supported"] == [GATEWAY_SCOPE] + + refused: Final = _rpc(candidate, f"/mcp/{alias}", denied, {}) + assert refused.status_code == 403, refused.text + assert "www-authenticate" not in refused.headers + assert tool_calls(peer.drain()) == () + + +def test_agent_365_prm_advertises_the_servers_configured_scopes(gateway: Gateway, tmp_path: Path) -> None: + config: Final = _sign_in_config(dict(AGENT_365_PARAMS), tmp_path / "agent365-scopes.yaml") + with ( + owned_proxy(gateway, tmp_path, {}, config=config) as candidate, + mcp_peer() as peer, + candidate.scenario() as scenario, + ): + scoped: Final = "a365" + uuid.uuid4().hex[:8] + register_mcp( + scenario, + peer, + scoped, + credentials={"scopes": ["https://example/mcp/scoped/access_as_user", "offline_access"]}, + ) + unscoped: Final = "a365" + uuid.uuid4().hex[:8] + register_mcp(scenario, peer, unscoped) + + assert _advertised(candidate, scoped) == ( + 200, + (ENTRA_ISSUER,), + ["https://example/mcp/scoped/access_as_user", "offline_access"], + ) + assert _advertised(candidate, unscoped) == (200, (ENTRA_ISSUER,), [GATEWAY_SCOPE]) + + +def test_challenge_names_the_server_first_route_the_client_connected_on(gateway: Gateway, tmp_path: Path) -> None: + config: Final = _sign_in_config(dict(AGENT_365_PARAMS), tmp_path / "agent365-route.yaml") + with ( + owned_proxy(gateway, tmp_path, {}, config=config) as candidate, + mcp_peer() as peer, + candidate.scenario() as scenario, + ): + alias: Final = "a365" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias) + granted: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + metadata_url: Final = f"{_origin(candidate)}/.well-known/oauth-protected-resource/{alias}/mcp" + + challenged: Final = _rpc(candidate, f"/{alias}/mcp", granted, {}) + assert challenged.status_code == 401, challenged.text + assert f'resource_metadata="{metadata_url}"' in challenged.headers.get("www-authenticate", "") + + document: Final = httpx.get(metadata_url, timeout=15).json() + assert document["resource"] == f"{_origin(candidate)}/{alias}/mcp", document + assert document["authorization_servers"] == [ENTRA_ISSUER] + assert tool_calls(peer.drain()) == () + + +def test_challenge_names_the_forwarded_origin_only_from_a_trusted_proxy(gateway: Gateway, tmp_path: Path) -> None: + config: Final = _sign_in_config( + dict(AGENT_365_PARAMS), + tmp_path / "agent365-forwarded.yaml", + {"use_x_forwarded_for": True, "mcp_trusted_proxy_ranges": ["127.0.0.0/8"]}, + ) + with ( + owned_proxy(gateway, tmp_path, {}, config=config) as candidate, + mcp_peer() as peer, + candidate.scenario() as scenario, + ): + alias: Final = "a365" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias) + granted: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + forwarded: Final = {"X-Forwarded-Proto": "https", "X-Forwarded-Host": "public.example"} + + challenged: Final = _rpc(candidate, f"/mcp/{alias}", granted, forwarded) + assert challenged.status_code == 401, challenged.text + assert ( + f'resource_metadata="https://public.example/.well-known/oauth-protected-resource/mcp/{alias}"' + in challenged.headers.get("www-authenticate", "") + ) + + plain: Final = _rpc(candidate, f"/mcp/{alias}", granted, {}) + assert plain.status_code == 401, plain.text + assert ( + f'resource_metadata="{_origin(candidate)}/.well-known/oauth-protected-resource/mcp/{alias}"' + in plain.headers.get("www-authenticate", "") + ) + + +def test_challenge_and_prm_resolve_the_connected_case_variant(gateway: Gateway, tmp_path: Path) -> None: + config: Final = _sign_in_config(dict(AGENT_365_PARAMS), tmp_path / "agent365-case.yaml") + with ( + owned_proxy(gateway, tmp_path, {}, config=config) as candidate, + mcp_peer() as peer, + candidate.scenario() as scenario, + ): + alias: Final = "a365" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias) + granted: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + connected_as: Final = alias.upper() + + challenged: Final = _rpc(candidate, f"/mcp/{connected_as}", granted, {}) + assert challenged.status_code == 401, challenged.text + authenticate: Final = challenged.headers.get("www-authenticate", "") + assert ( + f'resource_metadata="{_origin(candidate)}/.well-known/oauth-protected-resource/mcp/{connected_as}"' + in authenticate + ) + assert 'error="invalid_token"' in authenticate + + discovery: Final = candidate.client.get(f"/.well-known/oauth-protected-resource/mcp/{connected_as}") + assert discovery.status_code == 200, discovery.text + document: Final = discovery.json() + assert document["authorization_servers"] == [ENTRA_ISSUER] + assert document["scopes_supported"] == [GATEWAY_SCOPE] + assert tool_calls(peer.drain()) == () diff --git a/tests/integration/mcp/test_mcp_oauth_flows.py b/tests/integration/mcp/test_mcp_oauth_flows.py index 60563e7aacd..fdedd62abb1 100644 --- a/tests/integration/mcp/test_mcp_oauth_flows.py +++ b/tests/integration/mcp/test_mcp_oauth_flows.py @@ -1,6 +1,8 @@ import base64 import hashlib +import json import secrets +import socket import time import uuid from dataclasses import dataclass @@ -27,6 +29,7 @@ from integration._support.mcp import ( ) from integration._support.mcp_grants import create_toolset from integration._support.oauth_server import AuthorizationServer, oauth_server +from integration._support.wire import Reply, Request, wire_server ADD: Final = {"a": 2, "b": 3} CLIENT_REDIRECT: Final = "http://127.0.0.1:9/cb" @@ -203,6 +206,105 @@ def test_token_exchange_without_a_subject_token_is_rejected_before_any_upstream_ _assert_subject_token_challenge(as_subject, alias) +@pytest.mark.parametrize("status", (429, 408), ids=("throttled", "timed-out")) +def test_a_throttled_token_exchange_is_an_outage_not_a_sign_in_challenge(gateway: Gateway, status: int) -> None: + def shedding_idp(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/token", request + return Reply(status=status, body=json.dumps({"error": "temporarily_unavailable"}).encode()) + + with mcp_peer() as peer, wire_server(shedding_idp) as idp, gateway.scenario() as scenario: + alias: Final = "te" + uuid.uuid4().hex[:8] + identity: Final = register_mcp( + scenario, + peer, + alias, + auth_type="oauth2_token_exchange", + token_exchange_endpoint=idp.url + "/token", + credentials={"client_id": "te-client", "client_secret": "te-secret"}, + ) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + caller: Final = McpCaller( + gateway, key, "server_mcp", alias, headers={"Authorization": "Bearer subject-" + uuid.uuid4().hex} + ) + peer.drain() + response: Final = caller.rpc("tools/call", {"name": f"{alias}-add", "arguments": ADD}) + assert response.status_code == 503, (response.status_code, response.text, dict(response.headers)) + assert "www-authenticate" not in response.headers, dict(response.headers) + assert len(idp.drain()) == 1 + assert tool_calls(peer.drain()) == () + + +def test_a_malformed_caller_assertion_is_a_sign_in_challenge_not_an_outage(gateway: Gateway) -> None: + def entra_like_idp(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/token", request + body: Final = {"error": "invalid_client", "error_codes": [5002723], "error_description": "Invalid JWT token"} + return Reply(status=401, body=json.dumps(body).encode()) + + with mcp_peer() as peer, wire_server(entra_like_idp) as idp, gateway.scenario() as scenario: + alias: Final = "te" + uuid.uuid4().hex[:8] + identity: Final = register_mcp( + scenario, + peer, + alias, + auth_type="oauth2_token_exchange", + token_exchange_endpoint=idp.url + "/token", + credentials={"client_id": "te-client", "client_secret": "te-secret"}, + ) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + caller: Final = McpCaller( + gateway, + key, + "server_mcp", + alias, + headers={"Authorization": "Bearer eyJhbGciOiJSUzI1NiJ9.eyJhdWQiOiJ3cm9uZyJ9.c2ln"}, + ) + peer.drain() + response: Final = caller.rpc("tools/call", {"name": f"{alias}-add", "arguments": ADD}) + assert response.status_code == 401, (response.status_code, response.text, dict(response.headers)) + challenge: Final = response.headers["www-authenticate"] + origin: Final = str(gateway.client.base_url).rstrip("/") + assert f'resource_metadata="{origin}/.well-known/oauth-protected-resource/{alias}/mcp"' in challenge, challenge + assert 'error="invalid_token"' in challenge, challenge + assert len(idp.drain()) == 1 + assert tool_calls(peer.drain()) == () + + +def test_a_refused_token_exchange_answers_before_the_request_body_arrives(gateway: Gateway) -> None: + """The exchange preflight runs on the headers alone, so a client whose body is still in flight gets the + challenge at once instead of the gateway waiting for bytes it will never use.""" + + def entra_like_idp(request: Request) -> Reply: + body: Final = {"error": "invalid_client", "error_codes": [5002723], "error_description": "Invalid JWT token"} + return Reply(status=401, body=json.dumps(body).encode()) + + with mcp_peer() as peer, wire_server(entra_like_idp) as idp, gateway.scenario() as scenario: + alias: Final = "te" + uuid.uuid4().hex[:8] + identity: Final = register_mcp( + scenario, + peer, + alias, + auth_type="oauth2_token_exchange", + token_exchange_endpoint=idp.url + "/token", + credentials={"client_id": "te-client", "client_secret": "te-secret"}, + ) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + origin: Final = urlsplit(str(gateway.client.base_url)) + head: Final = ( + f"POST /mcp/{alias} HTTP/1.1\r\nHost: {origin.netloc}\r\nx-litellm-api-key: {key}\r\n" + "Authorization: Bearer eyJhbGciOiJSUzI1NiJ9.eyJhdWQiOiJ3cm9uZyJ9.c2ln\r\n" + "Accept: application/json, text/event-stream\r\nContent-Type: application/json\r\n" + "Content-Length: 4096\r\n\r\n" + ) + started: Final = time.monotonic() + with socket.create_connection((origin.hostname or "127.0.0.1", origin.port or 80), timeout=10) as raw: + raw.sendall(head.encode()) + status_line: Final = raw.recv(4096).split(b"\r\n", 1)[0] + assert status_line == b"HTTP/1.1 401 Unauthorized", status_line + assert time.monotonic() - started < 5 + assert len(idp.drain()) == 1 + assert tool_calls(peer.drain()) == () + + def _assert_subject_token_challenge(response: httpx.Response, alias: str) -> None: assert response.status_code == 401, response.text challenge: Final = response.headers["www-authenticate"] diff --git a/tests/unit/proxy/_experimental/mcp_server/outbound_credentials/test_adapter.py b/tests/unit/proxy/_experimental/mcp_server/outbound_credentials/test_adapter.py index 141260db700..ac1a23d0309 100644 --- a/tests/unit/proxy/_experimental/mcp_server/outbound_credentials/test_adapter.py +++ b/tests/unit/proxy/_experimental/mcp_server/outbound_credentials/test_adapter.py @@ -615,6 +615,21 @@ def test_raise_token_exchange_challenge_is_rfc9728_invalid_token(): assert "error_description=" in www +def test_raise_token_exchange_challenge_advertises_the_connected_route_metadata_url(): + from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import ( + raise_token_exchange_challenge, + ) + + connected: Final = "https://gw.example/.well-known/oauth-protected-resource/obo-srv/mcp" + with pytest.raises(HTTPException) as exc_info: + raise_token_exchange_challenge(_server(alias="obo-srv"), root_path="/", resource_metadata=connected) + assert exc_info.value.headers["WWW-Authenticate"] == ( + f'Bearer resource_metadata="{connected}", ' + 'error="invalid_token", ' + 'error_description="Missing or invalid subject token; authenticate with the IdP and retry"' + ) + + def test_raise_token_exchange_challenge_includes_server_root_path(monkeypatch): from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import ( raise_token_exchange_challenge, diff --git a/tests/unit/proxy/_experimental/mcp_server/outbound_credentials/test_token_exchange_provider.py b/tests/unit/proxy/_experimental/mcp_server/outbound_credentials/test_token_exchange_provider.py index cf15fdb3e26..3e714a73a98 100644 --- a/tests/unit/proxy/_experimental/mcp_server/outbound_credentials/test_token_exchange_provider.py +++ b/tests/unit/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, Ok, ServerSpec from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchange_provider import ( _post_exchange_endpoint, build_token_exchanger, @@ -17,24 +19,26 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchanger SubjectTokenRejected, TokenExchangeClientError, ) +from litellm.proxy._experimental.mcp_server.outbound_credentials.types import TokenExchangeConfig _HTTP_CLIENT = "litellm.llms.custom_httpx.http_handler.get_async_httpx_client" -def _client_raising_4xx(body: object): - """An httpx client whose POST returns a 4xx whose ``raise_for_status`` raises an HTTPStatusError - carrying ``body`` as its JSON, so the RFC 6749 error-code classification can be driven.""" +def _client_raising_status(status: int, body: object): + """An httpx client whose POST returns ``status`` whose ``raise_for_status`` raises an + HTTPStatusError carrying ``body`` as its JSON, so the RFC 6749 error-code classification can be + driven.""" import httpx request = httpx.Request("POST", "https://idp/token") - response = httpx.Response(400, json=body, request=request) + response = httpx.Response(status, json=body, request=request) class _Resp: def raise_for_status(self) -> None: raise httpx.HTTPStatusError("bad request", request=request, response=response) class _Client: - async def post(self, url, headers, data): + async def post(self, url, headers, data, timeout=None): return _Resp() return _Client() @@ -49,6 +53,25 @@ def test_build_gives_each_caller_an_independent_cache(): assert build_token_exchanger() is not build_token_exchanger() +@pytest.mark.asyncio +async def test_build_token_exchanger_drives_the_injected_http_edge(): + seen: list[tuple[str, dict[str, str]]] = [] + + async def post(url: str, form: dict[str, str], client_auth_headers: dict[str, str]) -> dict[str, object] | None: + seen.append((url, form)) + return {"access_token": "x", "expires_in": 60} + + config = TokenExchangeConfig( + token_exchange_endpoint="https://idp/token", client_id="cid", client_secret=SecretStr("csec") + ) + server = ServerSpec(server_id="srv", resource="https://up.example.com", config=config) + result = await build_token_exchanger(post=post).exchange("jwt", server, config) + assert isinstance(result, Ok) + assert result.ok.access_token == "x" + assert [url for url, _ in seen] == ["https://idp/token"] + assert seen[0][1]["subject_token"] == "jwt" + + @pytest.mark.asyncio async def test_post_returns_none_on_transport_error(): with patch(_HTTP_CLIENT, side_effect=RuntimeError("boom")): @@ -66,7 +89,7 @@ async def test_post_parses_json_body_on_success(): return {"access_token": "x", "expires_in": 60} class _Client: - async def post(self, url, headers, data): + async def post(self, url, headers, data, timeout=None): return _Resp() with patch(_HTTP_CLIENT, return_value=_Client()): @@ -80,7 +103,35 @@ async def test_post_parses_json_body_on_success(): ) async def test_post_maps_gateway_fault_4xx_to_client_error(code): # RFC 6749 5.2 gateway-fault codes must raise TokenExchangeClientError (-> 500), not the caller 401. - with patch(_HTTP_CLIENT, return_value=_client_raising_4xx({"error": code})): + with patch(_HTTP_CLIENT, return_value=_client_raising_status(400, {"error": code})): + with pytest.raises(TokenExchangeClientError): + await _post_exchange_endpoint("https://idp/token", {"grant_type": "x"}, {}) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("aadsts_code", [5002723, "5002710"], ids=["invalid_jwt", "no_kid_as_string"]) +async def test_post_maps_invalid_client_with_an_aadsts_50027xx_code_to_subject_rejected(aadsts_code): + # Entra reports a malformed or unverifiable caller assertion as invalid_client with an AADSTS50027xx + # sub-code (the same top-level error it uses for a bad gateway secret); that one is the caller's 401. + body = { + "error": "invalid_client", + "error_description": f"AADSTS{aadsts_code}: Invalid JWT token.", + "error_codes": [aadsts_code], + } + with patch(_HTTP_CLIENT, return_value=_client_raising_status(401, body)): + with pytest.raises(SubjectTokenRejected): + await _post_exchange_endpoint("https://idp/token", {"grant_type": "x"}, {}) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "error_codes", + [[7000215], [5002723.0], "5002723", None], + ids=["bad_secret_code", "float_code", "codes_not_a_list", "no_codes"], +) +async def test_post_keeps_invalid_client_without_an_assertion_code_as_client_error(error_codes): + body = {"error": "invalid_client", **({} if error_codes is None else {"error_codes": error_codes})} + with patch(_HTTP_CLIENT, return_value=_client_raising_status(401, body)): with pytest.raises(TokenExchangeClientError): await _post_exchange_endpoint("https://idp/token", {"grant_type": "x"}, {}) @@ -93,7 +144,7 @@ async def test_post_maps_gateway_fault_4xx_to_client_error(code): ) async def test_post_maps_subject_fault_4xx_to_subject_rejected(body): # A subject-fault code (or an unparseable/absent error) is the caller's problem -> SubjectTokenRejected (401). - with patch(_HTTP_CLIENT, return_value=_client_raising_4xx(body)): + with patch(_HTTP_CLIENT, return_value=_client_raising_status(400, body)): with pytest.raises(SubjectTokenRejected): await _post_exchange_endpoint("https://idp/token", {"grant_type": "x"}, {}) @@ -110,7 +161,7 @@ async def test_post_returns_none_on_non_object_json(payload): return payload class _Client: - async def post(self, url, headers, data): + async def post(self, url, headers, data, timeout=None): return _Resp() with patch(_HTTP_CLIENT, return_value=_Client()): @@ -128,7 +179,7 @@ async def test_post_threads_step_up_error_and_claims_into_subject_rejected(): "error_description": "AADSTS50079: the user must enroll MFA", "claims": claims, } - with patch(_HTTP_CLIENT, return_value=_client_raising_4xx(body)): + with patch(_HTTP_CLIENT, return_value=_client_raising_status(400, body)): with pytest.raises(SubjectTokenRejected) as exc_info: await _post_exchange_endpoint("https://idp/token", {"grant_type": "x"}, {}) assert exc_info.value.claims == claims @@ -137,7 +188,7 @@ async def test_post_threads_step_up_error_and_claims_into_subject_rejected(): @pytest.mark.asyncio async def test_post_subject_rejection_without_claims_carries_none_claims(): - with patch(_HTTP_CLIENT, return_value=_client_raising_4xx({"error": "invalid_grant"})): + with patch(_HTTP_CLIENT, return_value=_client_raising_status(400, {"error": "invalid_grant"})): with pytest.raises(SubjectTokenRejected) as exc_info: await _post_exchange_endpoint("https://idp/token", {"grant_type": "x"}, {}) assert exc_info.value.claims is None @@ -148,6 +199,26 @@ async def test_post_gateway_fault_still_wins_when_claims_are_present(): # A gateway-fault code stays a 500-class TokenExchangeClientError even if the body carries # claims; the caller cannot fix invalid_client by stepping up. body = {"error": "invalid_client", "claims": '{"access_token":{}}'} - with patch(_HTTP_CLIENT, return_value=_client_raising_4xx(body)): + with patch(_HTTP_CLIENT, return_value=_client_raising_status(400, body)): with pytest.raises(TokenExchangeClientError): await _post_exchange_endpoint("https://idp/token", {"grant_type": "x"}, {}) + + +_CONFIG = TokenExchangeConfig( + token_exchange_endpoint="https://idp.example.com/token", + client_id="cid", + client_secret=SecretStr("csec"), + scopes=("s1",), +) +_SERVER = ServerSpec(server_id="srv", resource="https://up.example.com", config=_CONFIG) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("status", [408, 429]) +async def test_exchange_maps_throttled_or_timed_out_4xx_to_upstream_unavailable(status): + # 408/429 are the IdP shedding load, not the caller presenting a bad subject: the exchange must + # surface upstream_unavailable (503-class, retryable) and never tell the caller to sign in again. + with patch(_HTTP_CLIENT, return_value=_client_raising_status(status, {"error": "temporarily_unavailable"})): + result = await OboTokenExchanger(_post_exchange_endpoint).exchange("caller-jwt", _SERVER, _CONFIG) + assert isinstance(result, Error) + assert result.error.tag == "upstream_unavailable" diff --git a/tests/unit/proxy/_experimental/mcp_server/test_caller_sign_in.py b/tests/unit/proxy/_experimental/mcp_server/test_caller_sign_in.py new file mode 100644 index 00000000000..edd29bf3ece --- /dev/null +++ b/tests/unit/proxy/_experimental/mcp_server/test_caller_sign_in.py @@ -0,0 +1,179 @@ +import asyncio +from collections.abc import Iterator, Mapping +from typing import Final + +import pytest + +import litellm +from litellm.integrations.custom_guardrail import CustomGuardrail +from litellm.proxy._experimental.mcp_server.caller_sign_in import ( + CallerSignIn, + CallerSignInPreflight, + CallerSignInProvider, + SignedIn, + caller_sign_in_for, + preflight_caller_sign_in, +) +from litellm.proxy._types import UserAPIKeyAuth +from litellm.types.mcp import MCPAuth, MCPTransport +from litellm.types.mcp_server.mcp_server_manager import MCPServer + + +class _SignInGuardrail(CustomGuardrail): + def __init__(self, issuer: str, scope: str, gated: bool = True) -> None: + super().__init__(guardrail_name=f"sign-in-{issuer}") + self.issuer: Final = issuer + self.scope: Final = scope + self.gated: Final = gated + self.seen: list[tuple[str, str | None]] = [] # mutable-ok: call recorder + + def caller_sign_in(self, server: MCPServer, user_api_key_auth: UserAPIKeyAuth | None) -> CallerSignIn | None: + self.seen.append((server.name, user_api_key_auth.user_id if user_api_key_auth else None)) + if not self.gated: + return None + return CallerSignIn(issuers=(self.issuer,), scopes=(self.scope,)) + + async def preflight_caller_sign_in( + self, server: MCPServer, user_api_key_auth: UserAPIKeyAuth | None, subject_token: str + ) -> CallerSignInPreflight: + return SignedIn() + + +class _PlainGuardrail(CustomGuardrail): + def __init__(self) -> None: + super().__init__(guardrail_name="plain") + + +def _server(auth_type: MCPAuth | None = None, scopes: list[str] | None = None) -> MCPServer: + return MCPServer( + server_id="tools-id", + name="tools", + server_name="tools", + transport=MCPTransport.http, + url="https://tools.test/mcp", + auth_type=auth_type, + scopes=scopes, + ) + + +@pytest.fixture +def registered() -> Iterator[tuple[_SignInGuardrail, _SignInGuardrail]]: + first: Final = _SignInGuardrail("https://idp-a.test", "scope-a") + second: Final = _SignInGuardrail("https://idp-b.test", "scope-b") + plain: Final = _PlainGuardrail() + for callback in (first, plain, second): + litellm.logging_callback_manager.add_litellm_callback(callback) + try: + yield first, second + finally: + for callback in (first, plain, second): + litellm.logging_callback_manager.remove_callback_from_list_by_object( + litellm.callbacks, callback, require_self=False + ) + + +def test_protocol_matches_only_guardrails_implementing_the_hook(): + assert isinstance(_SignInGuardrail("i", "s"), CallerSignInProvider) + assert not isinstance(_PlainGuardrail(), CallerSignInProvider) + + +def test_no_registered_provider_and_non_obo_advertises_nothing(): + assert caller_sign_in_for(_server(), None) is None + + +def test_registered_providers_merge_in_order_and_dedupe(registered): + first, second = registered + sign_in: Final = caller_sign_in_for(_server(), None) + + assert sign_in is not None + assert sign_in.issuers == ("https://idp-a.test", "https://idp-b.test") + assert sign_in.scopes == ("scope-a", "scope-b") + assert first.seen == [("tools", None)] + assert second.seen == [("tools", None)] + + +def test_provider_returning_none_contributes_nothing(registered): + ungated = _SignInGuardrail("https://idp-c.test", "scope-c", gated=False) + litellm.logging_callback_manager.add_litellm_callback(ungated) + try: + sign_in: Final = caller_sign_in_for(_server(), None) + assert sign_in is not None + assert "https://idp-c.test" not in sign_in.issuers + finally: + litellm.logging_callback_manager.remove_callback_from_list_by_object( + litellm.callbacks, ungated, require_self=False + ) + + +def test_obo_server_contributes_jwt_issuers_and_own_scopes(monkeypatch): + monkeypatch.setenv("JWT_ISSUER", "https://jwt-idp.test") + sign_in: Final = caller_sign_in_for(_server(auth_type=MCPAuth.oauth2_token_exchange, scopes=["read"]), None) + assert sign_in == CallerSignIn(issuers=("https://jwt-idp.test",), scopes=("read",)) + + +def test_obo_server_and_provider_merge_and_dedupe(monkeypatch, registered): + monkeypatch.setenv("JWT_ISSUER", "https://idp-a.test") + sign_in: Final = caller_sign_in_for( + _server(auth_type=MCPAuth.oauth2_token_exchange, scopes=["read", "scope-a"]), None + ) + assert sign_in is not None + assert sign_in.issuers == ("https://idp-a.test", "https://idp-b.test") + assert sign_in.scopes == ("read", "scope-a", "scope-b") + + +def test_obo_server_without_jwt_issuer_still_signs_in_when_a_provider_gates(registered): + sign_in: Final = caller_sign_in_for(_server(auth_type=MCPAuth.oauth2_token_exchange), None) + assert sign_in is not None + assert sign_in.issuers == ("https://idp-a.test", "https://idp-b.test") + + +def test_oauth_utils_strips_the_route_relative_root_path(): + """Regression: Starlette sets ``app_root_path`` to ``""`` on an unmounted app, so the strip must + fall back to ``root_path`` (which is where ``/mcp`` lands when the MCP app is mounted).""" + from litellm.proxy._experimental.mcp_server.oauth_utils import get_route_relative_request_path + + scope: Final[Mapping[str, object]] = { + "type": "http", + "path": "/mcp/catalog", + "root_path": "/mcp", + "app_root_path": "", + } + assert get_route_relative_request_path(scope) == "/catalog" # pyright: ignore[reportArgumentType] + + +@pytest.mark.asyncio +async def test_preflight_with_no_gating_provider_never_reads_the_request_body(monkeypatch): + """A plain OBO connect has nothing to pre-flight, so the exchange answers before the body arrives, as it + did before the sign-in seam; reading the body first would stall a client that sends its headers early.""" + monkeypatch.setenv("JWT_ISSUER", "https://jwt-idp.test") + body_read: Final = asyncio.Event() + + async def connecting() -> bool: + body_read.set() + return True + + await preflight_caller_sign_in( + _server(auth_type=MCPAuth.oauth2_token_exchange), + None, + "sub.ject.jws", + root_path="", + resource_metadata=None, + connecting=connecting, + ) + + assert not body_read.is_set() + + +@pytest.mark.asyncio +async def test_preflight_with_a_gating_provider_reads_the_body_to_tell_a_connect_apart(registered): + body_read: Final = asyncio.Event() + + async def connecting() -> bool: + body_read.set() + return True + + await preflight_caller_sign_in( + _server(), None, "sub.ject.jws", root_path="", resource_metadata=None, connecting=connecting + ) + + assert body_read.is_set() diff --git a/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 4e27ec134d4..b2753eebc14 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -7,6 +7,7 @@ from base64 import urlsafe_b64encode from datetime import datetime, timedelta, timezone from typing import TYPE_CHECKING, Final from unittest.mock import AsyncMock, MagicMock, patch +from urllib.parse import parse_qs, urlparse import pytest from fastapi import HTTPException @@ -392,6 +393,84 @@ def trust_xff(): yield +def _registered_gateway_oauth2_server(): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager + from litellm.proxy._types import MCPTransport + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + server: Final = MCPServer( + server_id="oid-7f3a", + name="gwx", + server_name="gwx", + alias="gwx", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow="authorization_code", + client_id="gw-client", + client_secret="gw-secret", + authorization_url="https://provider.com/oauth/authorize", + token_url="https://provider.com/oauth/token", + scopes=["read"], + ) + global_mcp_server_manager.registry.clear() + global_mcp_server_manager.registry[server.server_id] = server + return server + + +@pytest.mark.asyncio +@pytest.mark.parametrize("lookup", ["GWX", "oid-7f3a"], ids=["alias_case", "server_id"]) +async def test_authorization_server_doc_for_a_moved_lookup_matches_the_exact_name_doc(lookup): + """A case variant or server id now resolves like the exact name, so its AS metadata is the + exact-name doc with the requested spelling in the issuer and endpoint paths.""" + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + oauth_authorization_server_mcp_standard, + ) + + _registered_gateway_oauth2_server() + request: Final = _mock_callback_request("http://litellm.example.com/") + + exact: Final = await oauth_authorization_server_mcp_standard(request=request, mcp_server_name="gwx") + moved: Final = await oauth_authorization_server_mcp_standard(request=request, mcp_server_name=lookup) + + assert exact["issuer"] == "http://litellm.example.com/mcp/gwx" + assert moved == { + key: value.replace("/gwx", f"/{lookup}") if isinstance(value, str) else value for key, value in exact.items() + } + + +@pytest.mark.asyncio +@pytest.mark.parametrize("lookup", ["GWX", "oid-7f3a"], ids=["alias_case", "server_id"]) +async def test_authorize_relay_for_a_moved_lookup_redirects_like_the_exact_name(lookup): + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import authorize + + _registered_gateway_oauth2_server() + request: Final = _mock_callback_request("http://litellm.example.com/") + + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.encrypt_value_helper", return_value="sealed" + ): + exact: Final = await authorize( + request=request, mcp_server_name="gwx", redirect_uri="http://127.0.0.1:60108/callback", state="s1" + ) + moved: Final = await authorize( + request=request, mcp_server_name=lookup, redirect_uri="http://127.0.0.1:60108/callback", state="s1" + ) + + exact_target: Final = urlparse(exact.headers["location"]) + moved_target: Final = urlparse(moved.headers["location"]) + exact_query: Final = parse_qs(exact_target.query) + moved_query: Final = parse_qs(moved_target.query) + assert exact.status_code == 307 + assert exact_target._replace(query="") == urlparse("https://provider.com/oauth/authorize") + assert exact_query["client_id"] == ["gw-client"] + assert len(exact_query.pop("state")) == 1 and len(moved_query.pop("state")) == 1 + assert (moved.status_code, moved_target._replace(query=""), moved_query) == ( + exact.status_code, + exact_target._replace(query=""), + exact_query, + ) + + @pytest.mark.asyncio async def test_authorize_endpoint_includes_response_type(): """Test that authorize endpoint includes response_type=code parameter (fixes #15684)""" @@ -3642,8 +3721,8 @@ async def test_authorize_resolves_server_by_id_when_name_lookup_fails(): assert response.status_code == 307 assert "https://provider.com/oauth/authorize" in response.headers["location"] - by_name.assert_called_once_with(server.server_id, client_ip=None) - by_id.assert_called_once_with(server.server_id, client_ip=None) + by_name.assert_called_once_with(server.server_id) + by_id.assert_called_once_with(server.server_id) @pytest.mark.asyncio @@ -3679,8 +3758,8 @@ async def test_token_endpoint_resolves_server_by_id_when_name_lookup_fails(): ) assert json.loads(result.body)["access_token"] == "token" - by_name.assert_called_once_with(server.server_id, client_ip=None) - by_id.assert_called_once_with(server.server_id, client_ip=None) + by_name.assert_called_once_with(server.server_id) + by_id.assert_called_once_with(server.server_id) @pytest.mark.asyncio @@ -3711,8 +3790,8 @@ async def test_register_client_resolves_server_by_id_when_name_lookup_fails(): result = await discoverable_endpoints.register_client(request=request, mcp_server_name=server.server_id) assert json.loads(result.body)["client_id"] == "registered-client" - by_name.assert_called_once_with(server.server_id, client_ip=None) - by_id.assert_called_once_with(server.server_id, client_ip=None) + by_name.assert_called_once_with(server.server_id) + by_id.assert_called_once_with(server.server_id) @pytest.mark.asyncio @@ -3739,8 +3818,58 @@ async def test_protected_resource_metadata_resolves_server_by_id_when_name_looku assert result["authorization_servers"] == ["https://llm.example.com/mcp"] assert result["resource"] == f"https://llm.example.com/mcp/{server.server_id}" - by_name.assert_called_once_with(server.server_id, client_ip=None) - by_id.assert_called_once_with(server.server_id, client_ip=None) + by_name.assert_called_once_with(server.server_id) + by_id.assert_called_once_with(server.server_id) + + +@pytest.mark.asyncio +async def test_protected_resource_metadata_resolves_the_connected_case_variant(): + """The challenge points clients at the segment they connected with (``/mcp/CATALOG``), so the PRM + route must resolve that same segment through the answering-to fallback.""" + from fastapi import Request + + from litellm.proxy._experimental.mcp_server import discoverable_endpoints + from litellm.proxy._experimental.mcp_server.caller_sign_in import CallerSignIn + from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager + from litellm.types.mcp import MCPAuth, MCPTransport + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + server = MCPServer( + server_id="catalog-server-id-001", + name="catalog", + alias="catalog", + server_name="catalog", + url="https://catalog.test/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + mcp_info={"server_name": "catalog"}, + ) + sign_in: Final = CallerSignIn( + issuers=("https://login.microsoftonline.com/00000000-0000-0000-0000-000000000000/v2.0",), + scopes=("api://22222222-2222-2222-2222-222222222222/access_as_user",), + ) + request = MagicMock(spec=Request) + request.base_url = "https://llm.example.com/" + request.headers = {} + + with ( + patch.object(global_mcp_server_manager, "get_mcp_server_by_name", return_value=None), # test-quality-ok: resolver seam + patch.object(global_mcp_server_manager, "get_mcp_server_by_id", return_value=None), # test-quality-ok: resolver seam + patch.object( + global_mcp_server_manager, "get_filtered_registry", return_value={server.server_id: server} + ), # test-quality-ok: resolver seam + patch.object(discoverable_endpoints, "caller_sign_in_for", return_value=sign_in), # test-quality-ok: provider seam + ): + result = await discoverable_endpoints._build_oauth_protected_resource_response( + request=request, + mcp_server_name="CATALOG", + use_standard_pattern=True, + ) + + assert result["authorization_servers"] == ( + "https://login.microsoftonline.com/00000000-0000-0000-0000-000000000000/v2.0", + ) + assert result["resource"] == "https://llm.example.com/mcp/CATALOG" def test_authorization_server_metadata_resolves_server_by_id_when_name_lookup_fails(): @@ -3765,8 +3894,8 @@ def test_authorization_server_metadata_resolves_server_by_id_when_name_lookup_fa assert result["scopes_supported"] == server.scopes assert result["issuer"] == f"https://llm.example.com/{server.server_id}" - by_name.assert_called_once_with(server.server_id, client_ip=None) - by_id.assert_called_once_with(server.server_id, client_ip=None) + by_name.assert_called_once_with(server.server_id) + by_id.assert_called_once_with(server.server_id) @pytest.mark.asyncio @@ -7288,7 +7417,7 @@ async def test_token_exchange_persists_for_oauth2(): # ------------------------------------------------------------------- _OBO_RESOURCE = "https://litellm.example.com/mcp/obo_mcp" -_PATCH_ISSUERS = "litellm.proxy._experimental.mcp_server.discoverable_endpoints._jwt_auth_issuers" +_PATCH_ISSUERS = "litellm.proxy._experimental.mcp_server.caller_sign_in.jwt_auth_issuers" def _obo_server(scopes=None): @@ -7307,48 +7436,48 @@ def _obo_server(scopes=None): ) -def test_obo_protected_resource_response_names_jwt_issuers(): +def test_caller_sign_in_protected_resource_response_names_jwt_issuers(): """An OBO server's PRM points authorization_servers at the configured JWT issuers (the IdP that mints and validates the subject token), with the gateway resource echoed back.""" from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( - _obo_protected_resource_response, + _caller_sign_in_protected_resource_response, ) with patch(_PATCH_ISSUERS, return_value=["https://idp.example.com"]): - response = _obo_protected_resource_response(_obo_server(scopes=["read"]), _OBO_RESOURCE) + response = _caller_sign_in_protected_resource_response(_obo_server(scopes=["read"]), _OBO_RESOURCE) assert response == { - "authorization_servers": ["https://idp.example.com"], + "authorization_servers": ("https://idp.example.com",), "resource": _OBO_RESOURCE, - "scopes_supported": ["read"], + "scopes_supported": ("read",), } -def test_obo_protected_resource_response_scopes_default_empty(): - """A scopeless OBO server reports scopes_supported as [] rather than None.""" +def test_caller_sign_in_protected_resource_response_scopes_default_empty(): + """A scopeless OBO server reports scopes_supported as an empty array rather than None.""" from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( - _obo_protected_resource_response, + _caller_sign_in_protected_resource_response, ) with patch(_PATCH_ISSUERS, return_value=["https://idp.example.com"]): - response = _obo_protected_resource_response(_obo_server(scopes=None), _OBO_RESOURCE) - assert response["scopes_supported"] == [] + response = _caller_sign_in_protected_resource_response(_obo_server(scopes=None), _OBO_RESOURCE) + assert response["scopes_supported"] == () -def test_obo_protected_resource_response_falls_back_when_no_issuer(): +def test_caller_sign_in_protected_resource_response_falls_back_when_no_issuer(): """With no JWT issuer configured, the OBO branch returns None so the caller falls back to the gateway-default PRM (discovery still works, it just can't name the IdP).""" from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( - _obo_protected_resource_response, + _caller_sign_in_protected_resource_response, ) with patch(_PATCH_ISSUERS, return_value=[]): - assert _obo_protected_resource_response(_obo_server(), _OBO_RESOURCE) is None + assert _caller_sign_in_protected_resource_response(_obo_server(), _OBO_RESOURCE) is None -def test_obo_protected_resource_response_ignores_non_obo_server(): +def test_caller_sign_in_protected_resource_response_ignores_non_obo_server(): """Non-OBO servers are not handled by this branch (returns None -> gateway default).""" from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( - _obo_protected_resource_response, + _caller_sign_in_protected_resource_response, ) from litellm.proxy._types import MCPTransport from litellm.types.mcp import MCPAuth @@ -7360,7 +7489,7 @@ def test_obo_protected_resource_response_ignores_non_obo_server(): transport=MCPTransport.http, auth_type=MCPAuth.oauth2, ) - assert _obo_protected_resource_response(oauth2_server, _OBO_RESOURCE) is None + assert _caller_sign_in_protected_resource_response(oauth2_server, _OBO_RESOURCE) is None @pytest.mark.asyncio @@ -7390,7 +7519,7 @@ async def test_build_oauth_protected_resource_response_obo_end_to_end(): mcp_server_name="obo_mcp", use_standard_pattern=True, ) - assert response["authorization_servers"] == ["https://idp.example.com"] + assert response["authorization_servers"] == ("https://idp.example.com",) assert response["resource"] == "https://litellm.example.com/mcp/obo_mcp" finally: global_mcp_server_manager.registry.clear() diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py index 316988ef175..468359a819d 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py @@ -924,6 +924,7 @@ async def test_get_tools_from_mcp_servers(): mock_manager.get_mcp_server_by_id = lambda server_id: ( mock_server_1 if server_id == "server1_id" else mock_server_2 ) + mock_manager.get_mcp_server_answering_to = MCPServerManager().get_mcp_server_answering_to mock_manager._get_tools_from_server = AsyncMock(return_value=[mock_tool_1]) # Mock filter_server_ids_by_ip_with_info to return input unchanged (no IP filtering in test) mock_manager.filter_server_ids_by_ip_with_info = MagicMock( @@ -1004,6 +1005,7 @@ async def test_get_tools_from_mcp_servers(): if server_id == "server1_id" else (mock_server_2 if server_id == "server2_id" else mock_server_3) ) + mock_manager.get_mcp_server_answering_to = MagicMock(return_value=None) mock_manager._get_tools_from_server = AsyncMock(return_value=[mock_tool_1]) # Mock filter_server_ids_by_ip_with_info to return input unchanged (no IP filtering in test) mock_manager.filter_server_ids_by_ip_with_info = MagicMock( diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 0d2597975ae..b002e63c9ea 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -1637,6 +1637,19 @@ class TestMCPServerManager: assert server.uses_per_server_oauth_relay is True assert server.advertises_gateway_authorization_server is False + @pytest.mark.asyncio + async def test_load_servers_from_config_does_not_advertise_gateway_as_for_token_exchange(self): + # keeps_caller_authorization includes oauth2_token_exchange so a sign-in provider may gate it, + # but named discovery must still fall to the server's own PRM rather than the aggregate AS. + manager = MCPServerManager() + + with patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=None)): + await manager.load_servers_from_config(self._client_forwarded_config(MCPAuth.oauth2_token_exchange)) + + server = next(iter(manager.config_mcp_servers.values())) + assert server.keeps_caller_authorization is True + assert server.advertises_gateway_authorization_server is False + @pytest.mark.asyncio @pytest.mark.parametrize( "config", @@ -3090,6 +3103,34 @@ class TestMCPServerManager: www_authenticate = headers.get("WWW-Authenticate") or headers.get("www-authenticate") or "" assert "resource_metadata" in www_authenticate + @pytest.mark.asyncio + async def test_preflight_rejected_subject_challenge_names_the_connected_segment(self): + """A subject rejected on ``/mcp/`` must point resource_metadata at that same + segment, the way the sign-in preflight does, so the client's discovery fetch resolves.""" + from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Error + from litellm.proxy._experimental.mcp_server.outbound_credentials.types import CredError + + class _FakeProvider: + async def resolve_credentials(self, subject, server): + return Error(CredError.of_unauthorized("subject token rejected by the IdP")) + + manager = MCPServerManager(cred_provider=_FakeProvider()) + server = self._token_exchange_server("te-preflight-segment") + + with pytest.raises(HTTPException) as exc_info: + await manager.preflight_token_exchange( + server=server, + oauth2_headers={"Authorization": "Bearer rejected-subject"}, + user_api_key_auth=None, + resource_metadata=f"http://gw.test/.well-known/oauth-protected-resource/mcp/{server.server_id}", + ) + headers = exc_info.value.headers or {} + www_authenticate = headers.get("WWW-Authenticate") or headers.get("www-authenticate") or "" + assert ( + f'resource_metadata="http://gw.test/.well-known/oauth-protected-resource/mcp/{server.server_id}"' + in www_authenticate + ), www_authenticate + @pytest.mark.asyncio async def test_preflight_token_exchange_maps_gateway_fault_to_public_status(self): """A gateway-fault CredError (e.g. invalid_client) must surface its public status (500) @@ -6346,6 +6387,94 @@ class TestMCPServerManager: assert manager._resolve_mcp_server_for_tool_call("zapier", "create_zap") is zapier assert manager._resolve_mcp_server_for_tool_call("other", "create_zap") is other + @pytest.mark.parametrize("gh_first", [True, False], ids=["gh-listed-first", "gh-public-listed-first"]) + def test_answering_to_prefers_alias_over_earlier_prefix_match(self, gh_first): + manager = MCPServerManager() + gh = MCPServer( + server_id="gh-id", name="gh", server_name="gh", transport=MCPTransport.http, auth_type=MCPAuth.oauth2 + ) + gh_public = MCPServer( + server_id="gh-public-id", name="gh_public", server_name="gh_public", alias="gh", transport=MCPTransport.http + ) + manager.registry = ( + {"gh-id": gh, "gh-public-id": gh_public} if gh_first else {"gh-public-id": gh_public, "gh-id": gh} + ) + + assert manager.get_mcp_server_answering_to("gh") is gh_public + assert manager.get_mcp_server_answering_to("GH") is gh_public + assert manager.get_mcp_server_answering_to("Gh_Public") is gh_public + assert manager.get_mcp_server_answering_to("gh-public-id") is gh_public + assert manager.get_mcp_server_answering_to("GH_PUBLIC") is gh_public + assert manager.get_mcp_server_answering_to("gh-id") is gh + + @pytest.mark.parametrize("gh_first", [True, False], ids=["alias-listed-first", "server-name-listed-first"]) + def test_answering_to_agrees_with_exact_name_before_case_folding(self, gh_first): + manager = MCPServerManager() + by_alias = MCPServer( + server_id="a-id", + name="a", + server_name="a", + alias="gh", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + ) + by_server_name = MCPServer(server_id="b-id", name="b", server_name="Gh", transport=MCPTransport.http) + manager.registry = ( + {"a-id": by_alias, "b-id": by_server_name} if gh_first else {"b-id": by_server_name, "a-id": by_alias} + ) + + for name in ("Gh", "gh"): + assert manager.get_mcp_server_answering_to(name) is manager.get_mcp_server_by_name(name), name + assert manager.get_mcp_server_answering_to("Gh") is by_server_name + assert manager.get_mcp_server_answering_to("gh") is by_alias + assert manager.get_mcp_server_answering_to("GH") is by_alias + + @pytest.mark.parametrize("hidden_first", [True, False], ids=["hidden-listed-first", "public-listed-first"]) + def test_answering_to_never_reroutes_a_name_hidden_from_an_ip_to_a_case_variant(self, hidden_first): + manager = MCPServerManager() + hidden = MCPServer( + server_id="p-id", + name="gh", + server_name="gh", + transport=MCPTransport.http, + available_on_public_internet=False, + ) + public = MCPServer(server_id="u-id", name="u", server_name="u", alias="Gh", transport=MCPTransport.http) + manager.registry = {"p-id": hidden, "u-id": public} if hidden_first else {"u-id": public, "p-id": hidden} + + assert manager.get_mcp_server_answering_to("gh", client_ip="203.0.113.7") is None + assert manager.get_mcp_server_answering_to("p-id", client_ip="203.0.113.7") is None + assert manager.get_mcp_server_answering_to("gh", client_ip="10.0.0.7") is hidden + assert manager.get_mcp_server_answering_to("Gh", client_ip="203.0.113.7") is public + + @pytest.mark.parametrize("pinned_first", [True, False], ids=["pinned-id-listed-first", "alias-listed-first"]) + def test_answering_to_and_discovery_agree_on_a_pinned_id_that_another_alias_case_folds_to(self, pinned_first): + from litellm.proxy._experimental.mcp_server import discoverable_endpoints + from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager + + pinned = MCPServer(server_id="foo", name="pinned", server_name="pinned", transport=MCPTransport.http) + by_alias = MCPServer( + server_id="b-id", + name="b", + server_name="b", + alias="Foo", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + ) + registry = {"foo": pinned, "b-id": by_alias} if pinned_first else {"b-id": by_alias, "foo": pinned} + global_mcp_server_manager.registry.clear() + global_mcp_server_manager.registry.update(registry) + try: + for name in ("foo", "Foo", "FOO"): + connected = global_mcp_server_manager.get_mcp_server_answering_to(name) + discovered = discoverable_endpoints._resolve_mcp_server_by_name_or_id(name, client_ip=None) + assert discovered is connected, name + assert global_mcp_server_manager.get_mcp_server_answering_to("foo") is pinned + assert global_mcp_server_manager.get_mcp_server_answering_to("Foo") is by_alias + assert global_mcp_server_manager.get_mcp_server_answering_to("FOO") is by_alias + finally: + global_mcp_server_manager.registry.clear() + def test_remove_server_drops_only_its_own_tool_mapping_rows(self): manager = self._manager_with_deepwiki_and_huggingface() @@ -6973,6 +7102,185 @@ class TestMCPServerManager: assert exc_info.value.status_code == 403 + @pytest.mark.asyncio + @pytest.mark.parametrize( + ("raw_headers", "api_key", "expected_bearer", "expected_subject"), + [ + pytest.param( + {"authorization": "Bearer sk-1234"}, + "sk-1234", + "sk-1234", + None, + id="litellm-key-as-bearer-is-not-a-subject", + ), + pytest.param( + {"x-litellm-api-key": "sk-1234", "authorization": "Bearer eyJ.x.y"}, + "sk-1234", + "eyJ.x.y", + "eyJ.x.y", + id="key-admission-plus-idp-bearer-subject", + ), + pytest.param( + {"x-litellm-api-key": "sk-1234", "authorization": "bearer eyJ.x.y"}, + "sk-1234", + "eyJ.x.y", + "eyJ.x.y", + id="lowercase-bearer-scheme-is-the-subject", + ), + pytest.param( + {"x-litellm-api-key": "sk-1234", "authorization": "Basic a.b.c"}, + "sk-1234", + None, + None, + id="non-bearer-scheme-is-not-a-subject", + ), + pytest.param( + {"x-litellm-api-key": "sk-1234", "authorization": "Digest x.y.z"}, + "sk-1234", + None, + None, + id="digest-scheme-is-not-a-subject", + ), + pytest.param( + {"x-litellm-api-key": "sk-1234", "authorization": "eyJ.x.y"}, + "sk-1234", + None, + None, + id="scheme-less-value-is-not-a-subject", + ), + ], + ) + async def test_pre_call_tool_check_separates_raw_bearer_from_subject( + self, raw_headers, api_key, expected_bearer, expected_subject + ): + manager = MCPServerManager() + server = MCPServer( + server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv", allowed_tools=None + ) + proxy_logging = MagicMock() + proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={}) + proxy_logging.pre_call_hook = AsyncMock(return_value=None) + + await manager.pre_call_tool_check( + server_name="srv", + name="turn", + arguments={}, + user_api_key_auth=UserAPIKeyAuth(api_key=api_key, user_id="u"), + proxy_logging_obj=proxy_logging, + server=server, + raw_headers=raw_headers, + ) + + kwargs: Final = proxy_logging._create_mcp_request_object_from_kwargs.call_args.args[0] + assert kwargs["incoming_bearer_token"] == expected_bearer + assert kwargs["incoming_subject_token"] == expected_subject + + @pytest.mark.asyncio + @pytest.mark.parametrize( + ("raw_headers", "expected_subject"), + [ + pytest.param({"authorization": "Bearer gw.master.key"}, None, id="master-key-alone-is-not-a-subject"), + pytest.param( + {"x-litellm-api-key": "sk-1234", "authorization": "Bearer gw.master.key"}, + None, + id="master-key-next-to-a-virtual-key-is-not-a-subject", + ), + pytest.param( + {"x-litellm-api-key": "sk-1234", "authorization": "Bearer gw.other.jws"}, + "gw.other.jws", + id="a-dotted-bearer-that-is-not-the-master-key-is-the-subject", + ), + ], + ) + async def test_pre_call_tool_check_withholds_a_dotted_master_key_from_the_subject( + self, raw_headers, expected_subject + ): + """A master key is a LiteLLM credential whatever its shape, so even one with the two dots of a + compact JWS never becomes the sign-in subject a provider would send to its IdP.""" + manager = MCPServerManager() + server = MCPServer( + server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv", allowed_tools=None + ) + proxy_logging = MagicMock() + proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={}) + proxy_logging.pre_call_hook = AsyncMock(return_value=None) + + with patch("litellm.proxy.proxy_server.master_key", "gw.master.key"): + await manager.pre_call_tool_check( + server_name="srv", + name="turn", + arguments={}, + user_api_key_auth=UserAPIKeyAuth(api_key="sk-1234", user_id="u"), + proxy_logging_obj=proxy_logging, + server=server, + raw_headers=raw_headers, + ) + + kwargs: Final = proxy_logging._create_mcp_request_object_from_kwargs.call_args.args[0] + assert kwargs["incoming_bearer_token"] == raw_headers["authorization"].removeprefix("Bearer ") + assert kwargs["incoming_subject_token"] == expected_subject + + @pytest.mark.asyncio + @pytest.mark.parametrize( + ("raw_headers", "api_key", "custom_auth", "expected_subject"), + [ + pytest.param( + {"authorization": "Bearer eyJ.x.y"}, + "eyJ.x.y", + True, + "eyJ.x.y", + id="custom-auth-idp-bearer-is-the-subject", + ), + pytest.param( + {"authorization": "Bearer eyJ.x.y"}, + "eyJ.x.y", + False, + "eyJ.x.y", + id="built-in-oauth2-admission-bearer-is-the-subject", + ), + pytest.param({"authorization": "Bearer sk-1234"}, "sk-1234", True, None, id="virtual-key-is-not-a-subject"), + pytest.param( + {"x-litellm-api-key": "ca-key", "authorization": "Bearer ca-key"}, + "ca-key", + True, + None, + id="explicit-key-admission-repeated-in-authorization-is-not-a-subject", + ), + ], + ) + async def test_pre_call_tool_check_hands_sign_in_the_bearer_that_admitted_the_caller( + self, raw_headers, api_key, custom_auth, expected_subject + ): + """Custom auth and the built-in OAuth2 admission both admit the caller on its own IdP token in + ``Authorization`` with no ``x-litellm-api-key`` and record it as ``api_key``; that token is the + sign-in subject as the raw bearer was before the subject split.""" + manager = MCPServerManager() + server = MCPServer( + server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv", allowed_tools=None + ) + admitted = UserAPIKeyAuth(api_key=api_key, user_id="u") + admitted.authenticated_by_custom_auth = custom_auth + proxy_logging = MagicMock() + proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={}) + proxy_logging.pre_call_hook = AsyncMock(return_value=None) + + await manager.pre_call_tool_check( + server_name="srv", + name="turn", + arguments={}, + user_api_key_auth=admitted, + proxy_logging_obj=proxy_logging, + server=server, + raw_headers=raw_headers, + ) + + kwargs: Final = proxy_logging._create_mcp_request_object_from_kwargs.call_args.args[0] + assert kwargs["incoming_bearer_token"] == api_key + assert kwargs["incoming_subject_token"] == expected_subject + @pytest.mark.asyncio async def test_check_tool_permission_for_key_team_allows_permitted_tool(self): """ diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py index 545b2757ffd..7eae64bd160 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py @@ -1,4 +1,3 @@ -from litellm.proxy._experimental.mcp_server import operations as mcp_operations import asyncio import contextlib import contextvars @@ -7,7 +6,7 @@ import os from datetime import datetime, timedelta from types import SimpleNamespace from typing import Final -from unittest.mock import AsyncMock, MagicMock, patch +from unittest.mock import AsyncMock, MagicMock, call, patch import httpx import pytest @@ -29,8 +28,11 @@ 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._experimental.mcp_server.mcp_server_manager import ListedToolsCaller, MCPServerManager from litellm.proxy._types import ( LiteLLM_MCPServerTable, MCPTransport, @@ -51,6 +53,10 @@ def test_mcp_available_on_sdk2(): assert MCP_AVAILABLE is True +async def _connecting() -> bool: + return True + + def _rendered_log_message(call): message = str(call.args[0]) values = call.args[1:] @@ -1159,6 +1165,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" @@ -1276,6 +1283,7 @@ async def test_get_tools_from_mcp_servers_continues_when_one_server_fails(): mock_manager.get_mcp_server_by_id = lambda server_id: ( working_server if server_id == "working_server" else failing_server ) + mock_manager.get_mcp_server_answering_to = MCPServerManager().get_mcp_server_answering_to # Mock filter_server_ids_by_ip to return server_ids unchanged (no IP filtering) mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: ( server_ids, @@ -3750,6 +3758,521 @@ async def test_stateful_mcp_session_owner_mismatch_returns_403(): mcp_server._stateful_session_owners.pop(session_id, None) +@pytest.mark.asyncio +async def test_stateful_mcp_session_owner_mismatch_is_rejected_before_the_body_is_read(): + """A POST carrying another caller's mcp-session-id is refused before any body chunk is awaited, so a + slow sender cannot hold the request open past the owner check.""" + try: + from litellm.proxy._experimental.mcp_server import server as mcp_server + from litellm.proxy._experimental.mcp_server.server import ( + handle_streamable_http_mcp, + session_manager_stateful, + ) + except ImportError: + pytest.skip("MCP server not available") + + session_id = "owned-session-slow-body" + owner_auth = UserAPIKeyAuth(api_key="owner-key", user_id="owner") + intruder_auth = UserAPIKeyAuth(api_key="intruder-key", user_id="intruder") + mcp_server._stateful_session_auth_contexts[session_id] = MagicMock() + mcp_server._stateful_session_owners[session_id] = mcp_server._owner_fingerprint_for(owner_auth) + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp", + "headers": [ + (b"content-type", b"application/json"), + (b"authorization", b"Bearer intruder-key"), + (b"mcp-session-id", session_id.encode()), + ], + } + body_never_arrives = asyncio.Event() + + async def stalled_receive(): + await body_never_arrives.wait() + return {"type": "http.request", "body": b"", "more_body": False} + + sent_messages: list = [] + + async def capture_send(message): + sent_messages.append(message) + + try: + with ( + patch( + "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", + new_callable=AsyncMock, + return_value=(intruder_auth, None, None, None, None, None), + ), + patch("litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", True), + patch.object(session_manager_stateful, "handle_request", new_callable=AsyncMock), + patch.object(session_manager_stateful, "_server_instances", {session_id: MagicMock()}), + ): + await asyncio.wait_for(handle_streamable_http_mcp(scope, stalled_receive, capture_send), timeout=2) + finally: + body_never_arrives.set() + mcp_server._stateful_session_auth_contexts.pop(session_id, None) + mcp_server._stateful_session_owners.pop(session_id, None) + + statuses = [m["status"] for m in sent_messages if m.get("type") == "http.response.start"] + assert statuses == [403], sent_messages + + +@pytest.mark.asyncio +async def test_stateful_mcp_session_owner_mismatch_is_refused_before_the_body_arrives(): + """A POST naming another caller's session is refused with 403 while the sender is still withholding the body, + so a stalled body cannot delay the refusal or reach the stateful manager.""" + try: + from litellm.proxy._experimental.mcp_server import server as mcp_server + from litellm.proxy._experimental.mcp_server.server import ( + handle_streamable_http_mcp, + session_manager_stateful, + ) + except ImportError: + pytest.skip("MCP server not available") + + session_id = "owned-session-stalled-body" + owner_auth = UserAPIKeyAuth(api_key="owner-key", user_id="owner") + intruder_auth = UserAPIKeyAuth(api_key="intruder-key", user_id="intruder") + mcp_server._stateful_session_auth_contexts[session_id] = MagicMock() + mcp_server._stateful_session_owners[session_id] = mcp_server._owner_fingerprint_for(owner_auth) + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp", + "headers": [ + (b"content-type", b"application/json"), + (b"authorization", b"Bearer intruder-key"), + (b"mcp-session-id", session_id.encode()), + ], + } + body_never_sent = asyncio.Event() + + async def withheld_body() -> Message: + await body_never_sent.wait() + return {"type": "http.request", "body": b"", "more_body": False} + + sent_messages: list[Message] = [] + + async def capture_send(message: Message) -> None: + sent_messages.append(message) + + handle_request_mock = AsyncMock() + + try: + with ( + patch( + "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", + new_callable=AsyncMock, + return_value=(intruder_auth, None, None, None, None, None), + ), + patch( + "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", + True, + ), + patch.object( + session_manager_stateful, + "handle_request", + side_effect=handle_request_mock, + ), + patch.object( + session_manager_stateful, + "_server_instances", + {session_id: MagicMock()}, + ), + ): + await asyncio.wait_for(handle_streamable_http_mcp(scope, withheld_body, capture_send), timeout=2) + finally: + mcp_server._stateful_session_auth_contexts.pop(session_id, None) + mcp_server._stateful_session_owners.pop(session_id, None) + + handle_request_mock.assert_not_awaited() + statuses = [m["status"] for m in sent_messages if m.get("type") == "http.response.start"] + assert statuses == [403] + + +@pytest.mark.asyncio +async def test_owner_mismatch_on_a_torn_down_session_is_refused_before_the_body_under_a_fail_closed_outage(): + """While a session's transport is already gone but its owner binding is still recorded, a POST from another + caller is refused with 403 before its withheld body arrives even when the fail-closed sign-in gate would + otherwise peek at the body to tell an ``initialize`` apart.""" + try: + from litellm.proxy._experimental.mcp_server import server as mcp_server + from litellm.proxy._experimental.mcp_server.caller_sign_in import Unavailable + from litellm.proxy._experimental.mcp_server.server import ( + handle_streamable_http_mcp, + session_manager_stateful, + ) + except ImportError: + pytest.skip("MCP server not available") + + server = _catalog_server() + guardrail = _CallerSignInGuardrail( + guardrail_name="sign-in-stub", + preflight_result=Unavailable("the Entra token endpoint could not be reached", fail_open=False), + ) + session_id = "owned-session-being-torn-down" + owner_auth = UserAPIKeyAuth(api_key="owner-key", user_id="owner") + intruder_auth = UserAPIKeyAuth(api_key="intruder-key", user_id="intruder") + mcp_server._stateful_session_owners[session_id] = mcp_server._owner_fingerprint_for(owner_auth) + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/catalog", + "headers": [ + (b"content-type", b"application/json"), + (b"x-litellm-api-key", b"intruder-key"), + (b"authorization", b"Bearer entra.jwt.token"), + (b"mcp-session-id", session_id.encode()), + ], + } + body_never_sent = asyncio.Event() + + async def withheld_body() -> Message: + await body_never_sent.wait() + return {"type": "http.request", "body": b"", "more_body": False} + + sent_messages: list[Message] = [] + + async def capture_send(message: Message) -> None: + sent_messages.append(message) + + handle_request_mock = AsyncMock() + litellm.logging_callback_manager.add_litellm_callback(guardrail) + try: + with ( + patch( + "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", + new_callable=AsyncMock, + return_value=( + intruder_auth, + None, + ["catalog"], + None, + None, + {"x-litellm-api-key": "intruder-key", "authorization": "Bearer entra.jwt.token"}, + ), + ), + 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])), + patch.object(mcp_server, "_check_passthrough_upstream_auth", AsyncMock()), + patch.object(mcp_operations, "_raise_if_initialize_grants_no_mcp_servers", AsyncMock()), + patch.object(mcp_server, "_SESSION_MANAGERS_INITIALIZED", True), + patch.object(session_manager_stateful, "handle_request", side_effect=handle_request_mock), + patch.object(session_manager_stateful, "_server_instances", {}), + ): + await asyncio.wait_for(handle_streamable_http_mcp(scope, withheld_body, capture_send), timeout=2) + finally: + mcp_server._stateful_session_owners.pop(session_id, None) + litellm.logging_callback_manager.remove_callback_from_list_by_object( + litellm.callbacks, guardrail, require_self=False + ) + + assert guardrail.preflight_calls == [] + handle_request_mock.assert_not_awaited() + statuses = [m["status"] for m in sent_messages if m.get("type") == "http.response.start"] + assert statuses == [403] + + +@pytest.mark.asyncio +async def test_initialize_naming_a_stale_session_still_meets_the_fail_closed_connect_gate(): + """A client that retries ``initialize`` with a session id this worker no longer knows is connecting, so a + fail-closed sign-in outage answers it 503 at the gate instead of letting the stripped-header retry open a + session.""" + try: + from litellm.proxy._experimental.mcp_server import server as mcp_server + from litellm.proxy._experimental.mcp_server.caller_sign_in import Unavailable + from litellm.proxy._experimental.mcp_server.server import ( + handle_streamable_http_mcp, + session_manager_stateful, + ) + except ImportError: + pytest.skip("MCP server not available") + + server = _catalog_server() + guardrail = _CallerSignInGuardrail( + guardrail_name="sign-in-stub", + preflight_result=Unavailable("the Entra token endpoint could not be reached", fail_open=False), + ) + initialize = json.dumps({"jsonrpc": "2.0", "id": 1, "method": "initialize", "params": {}}).encode() + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/catalog", + "headers": [ + (b"content-type", b"application/json"), + (b"authorization", b"Bearer entra.jwt.token"), + (b"mcp-session-id", b"stale-session-from-a-restarted-worker"), + ], + } + incoming: asyncio.Queue[Message] = asyncio.Queue() + await incoming.put({"type": "http.request", "body": initialize, "more_body": False}) + sent_messages: list[Message] = [] + + async def capture_send(message: Message) -> None: + sent_messages.append(message) + + handle_request_mock = AsyncMock() + litellm.logging_callback_manager.add_litellm_callback(guardrail) + try: + with ( + patch( + "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", + new_callable=AsyncMock, + return_value=( + UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"), + None, + ["catalog"], + None, + None, + {"x-litellm-api-key": "sk-litellm-virtual-key", "authorization": "Bearer entra.jwt.token"}, + ), + ), + 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])), + patch.object(mcp_server, "_check_passthrough_upstream_auth", AsyncMock()), + patch.object(mcp_operations, "_raise_if_initialize_grants_no_mcp_servers", AsyncMock()), + patch.object(mcp_server, "_SESSION_MANAGERS_INITIALIZED", True), + patch.object(session_manager_stateful, "handle_request", side_effect=handle_request_mock), + patch.object(session_manager_stateful, "_server_instances", {}), + pytest.raises(HTTPException) as exc, + ): + await asyncio.wait_for(handle_streamable_http_mcp(scope, incoming.get, capture_send), timeout=2) + finally: + litellm.logging_callback_manager.remove_callback_from_list_by_object( + litellm.callbacks, guardrail, require_self=False + ) + + assert exc.value.status_code == 503 + assert guardrail.preflight_calls == ["entra.jwt.token"] + handle_request_mock.assert_not_awaited() + assert sent_messages == [] + + +@pytest.mark.asyncio +async def test_connect_challenge_answers_before_the_body_is_read(monkeypatch): + """A session-less ``POST`` to a gated server without a subject token gets its RFC 9728 challenge straight + away: the gate must not wait for the body it never needs, so a client that withholds it still sees 401.""" + try: + from litellm.proxy._experimental.mcp_server import server as mcp_server + from litellm.proxy._experimental.mcp_server.server import ( + handle_streamable_http_mcp, + session_manager_stateful, + ) + except ImportError: + pytest.skip("MCP server not available") + + monkeypatch.delenv("SERVER_ROOT_PATH", raising=False) + server = _make_obo_server("obo_server") + mcp_operations.global_mcp_server_manager.registry.update({server.server_id: server}) + scope = { + "type": "http", + "method": "POST", + "scheme": "http", + "path": "/mcp/obo_server", + "root_path": "", + "query_string": b"", + "server": ("gw.example", 4000), + "client": ("10.0.0.7", 51000), + "headers": [ + (b"host", b"gw.example:4000"), + (b"content-type", b"application/json"), + (b"x-litellm-api-key", b"sk-litellm-virtual-key"), + ], + } + withheld_body: asyncio.Queue[Message] = asyncio.Queue() + sent_messages: list[Message] = [] + + async def capture_send(message: Message) -> None: + sent_messages.append(message) + + handle_request_mock = AsyncMock() + with ( + patch( + "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", + new_callable=AsyncMock, + return_value=( + UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"), + None, + ["obo_server"], + None, + None, + {"x-litellm-api-key": "sk-litellm-virtual-key"}, + ), + ), + patch.object( # test-quality-ok: the key's grant list lives in the DB; the real router runs on it + mcp_operations.global_mcp_server_manager, + "get_allowed_mcp_servers", + AsyncMock(return_value=[server.server_id]), + ), + patch.object(mcp_server, "_SESSION_MANAGERS_INITIALIZED", True), + patch.object(session_manager_stateful, "handle_request", side_effect=handle_request_mock), + patch.object(session_manager_stateful, "_server_instances", {}), + pytest.raises(HTTPException) as exc, + ): + await asyncio.wait_for(handle_streamable_http_mcp(scope, withheld_body.get, capture_send), timeout=2) + + assert exc.value.status_code == 401 + assert (exc.value.headers or {})["WWW-Authenticate"].startswith( + 'Bearer resource_metadata="http://gw.example:4000/.well-known/oauth-protected-resource/mcp/obo_server"' + ), exc.value.headers + handle_request_mock.assert_not_awaited() + assert sent_messages == [] + + +@pytest.mark.asyncio +async def test_initialize_naming_a_session_whose_transport_is_already_gone_still_meets_the_fail_closed_connect_gate(): + """While idle purge or cap eviction is still terminating a transport, its owner entry outlives the transport; + an ``initialize`` retried with that id is still a new connection and meets the fail-closed gate.""" + try: + from litellm.proxy._experimental.mcp_server import server as mcp_server + from litellm.proxy._experimental.mcp_server.caller_sign_in import Unavailable + from litellm.proxy._experimental.mcp_server.server import ( + handle_streamable_http_mcp, + session_manager_stateful, + ) + except ImportError: + pytest.skip("MCP server not available") + + server = _catalog_server() + guardrail = _CallerSignInGuardrail( + guardrail_name="sign-in-stub", + preflight_result=Unavailable("the Entra token endpoint could not be reached", fail_open=False), + ) + initialize = json.dumps({"jsonrpc": "2.0", "id": 1, "method": "initialize", "params": {}}).encode() + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/catalog", + "headers": [ + (b"content-type", b"application/json"), + (b"authorization", b"Bearer entra.jwt.token"), + (b"mcp-session-id", b"session-mid-termination"), + ], + } + incoming: asyncio.Queue[Message] = asyncio.Queue() + await incoming.put({"type": "http.request", "body": initialize, "more_body": False}) + sent_messages: list[Message] = [] + + async def capture_send(message: Message) -> None: + sent_messages.append(message) + + handle_request_mock = AsyncMock() + litellm.logging_callback_manager.add_litellm_callback(guardrail) + try: + with ( + patch( + "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", + new_callable=AsyncMock, + return_value=( + UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"), + None, + ["catalog"], + None, + None, + {"x-litellm-api-key": "sk-litellm-virtual-key", "authorization": "Bearer entra.jwt.token"}, + ), + ), + 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])), + patch.object(mcp_server, "_check_passthrough_upstream_auth", AsyncMock()), + patch.object(mcp_operations, "_raise_if_initialize_grants_no_mcp_servers", AsyncMock()), + patch.object(mcp_server, "_SESSION_MANAGERS_INITIALIZED", True), + patch.object(session_manager_stateful, "handle_request", side_effect=handle_request_mock), + patch.object(session_manager_stateful, "_server_instances", {}), + patch.object(mcp_server, "_owner_fingerprint_for", return_value="owner-fingerprint"), + patch.dict( + mcp_server._stateful_session_owners, {"session-mid-termination": "owner-fingerprint"}, clear=True + ), + pytest.raises(HTTPException) as exc, + ): + await asyncio.wait_for(handle_streamable_http_mcp(scope, incoming.get, capture_send), timeout=2) + finally: + litellm.logging_callback_manager.remove_callback_from_list_by_object( + litellm.callbacks, guardrail, require_self=False + ) + + assert exc.value.status_code == 503 + assert guardrail.preflight_calls == ["entra.jwt.token"] + handle_request_mock.assert_not_awaited() + assert sent_messages == [] + + +@pytest.mark.asyncio +async def test_stateful_mcp_session_post_hands_the_whole_body_to_the_session_manager(): + """The owner's follow-up POST on a live session reaches the stateful manager with every body byte intact + after the routing peek.""" + try: + from litellm.proxy._experimental.mcp_server import server as mcp_server + from litellm.proxy._experimental.mcp_server.server import ( + handle_streamable_http_mcp, + session_manager_stateful, + ) + except ImportError: + pytest.skip("MCP server not available") + + session_id = "owned-session-replay" + owner_auth = UserAPIKeyAuth(api_key="owner-key", user_id="owner") + mcp_server._stateful_session_auth_contexts[session_id] = MagicMock() + mcp_server._stateful_session_owners[session_id] = mcp_server._owner_fingerprint_for(owner_auth) + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp", + "headers": [ + (b"content-type", b"application/json"), + (b"authorization", b"Bearer owner-key"), + (b"mcp-session-id", session_id.encode()), + ], + } + chunks = [ + {"type": "http.request", "body": b'{"jsonrpc":"2.0","id":7,"method":"tools/list",', "more_body": True}, + {"type": "http.request", "body": b'"params":{}}', "more_body": False}, + ] + receive = AsyncMock(side_effect=list(chunks)) + delivered: list[bytes] = [] + + async def drain_body(scope_, receive_, send_): + while True: + message = await receive_() + delivered.append(message.get("body", b"")) + if not message.get("more_body", False): + return + + try: + with ( + patch( + "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", + new_callable=AsyncMock, + return_value=(owner_auth, None, None, None, None, None), + ), + patch("litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", True), + patch.object(session_manager_stateful, "handle_request", side_effect=drain_body), + patch.object(session_manager_stateful, "_server_instances", {session_id: MagicMock()}), + ): + await handle_streamable_http_mcp(scope, receive, AsyncMock()) + finally: + mcp_server._stateful_session_auth_contexts.pop(session_id, None) + mcp_server._stateful_session_owners.pop(session_id, None) + + assert b"".join(delivered) == b'{"jsonrpc":"2.0","id":7,"method":"tools/list","params":{}}' + + @pytest.mark.asyncio async def test_stateful_mcp_session_serializes_concurrent_requests(): """ @@ -6761,6 +7284,7 @@ async def test_list_tools_with_legacy_db_m2m_server_resolves_oauth2_flow(): ): mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["legacy-m2m-id"]) mock_manager.get_mcp_server_by_id = MagicMock(return_value=legacy_server) + mock_manager.get_mcp_server_answering_to = MCPServerManager().get_mcp_server_answering_to mock_manager.filter_server_ids_by_ip_with_info = MagicMock(return_value=(["legacy-m2m-id"], 0)) mock_manager._get_tools_from_server = AsyncMock(side_effect=capture_extra_headers) @@ -8682,6 +9206,458 @@ async def test_get_allowed_mcp_servers_from_mcp_server_names_known_alias_returns assert [s.server_id for s in result] == ["id-a"] +@pytest.mark.asyncio +@pytest.mark.parametrize("alias_server_first", [True, False], ids=["alias-granted-first", "server-name-granted-first"]) +async def test_scoped_router_selects_the_server_the_connect_preflight_resolves(alias_server_first): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager + from litellm.proxy._experimental.mcp_server.server import _get_allowed_mcp_servers_from_mcp_server_names + + by_alias = MCPServer(server_id="a-id", name="a", server_name="a", alias="gh", transport=MCPTransport.http) + by_server_name = MCPServer(server_id="b-id", name="b", server_name="Gh", transport=MCPTransport.http) + global_mcp_server_manager.registry.clear() + global_mcp_server_manager.registry.update({"a-id": by_alias, "b-id": by_server_name}) + granted_both = [by_alias, by_server_name] if alias_server_first else [by_server_name, by_alias] + try: + with patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp." + "MCPRequestHandler._get_mcp_servers_from_access_groups", + new_callable=AsyncMock, + return_value=[], + ): + for name in ("Gh", "gh", "GH"): + expected = global_mcp_server_manager.get_mcp_server_answering_to(name) + selected = await _get_allowed_mcp_servers_from_mcp_server_names( + mcp_servers=[name], allowed_mcp_servers=granted_both + ) + assert [s.server_id for s in selected] == [expected.server_id], name + only_b = await _get_allowed_mcp_servers_from_mcp_server_names( + mcp_servers=["gh"], allowed_mcp_servers=[by_server_name] + ) + finally: + global_mcp_server_manager.registry.clear() + + assert [s.server_id for s in only_b] == ["b-id"], ( + "the granted server answering to the name wins over the registry's ungranted alias holder" + ) + + +@pytest.mark.asyncio +async def test_scoped_name_of_an_ungranted_server_is_retried_as_an_access_group_the_key_holds(): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager + from litellm.proxy._experimental.mcp_server.server import _get_allowed_mcp_servers_from_mcp_server_names + + shadow = MCPServer(server_id="s-id", name="docs", server_name="docs", transport=MCPTransport.http) + member = MCPServer(server_id="m-id", name="m", server_name="m", transport=MCPTransport.http) + global_mcp_server_manager.registry.clear() + global_mcp_server_manager.registry.update({"s-id": shadow, "m-id": member}) + group_members = {"docs": ["m-id"]} + try: + with patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp." + "MCPRequestHandler._get_mcp_servers_from_access_groups", + new_callable=AsyncMock, + side_effect=lambda names: [sid for name in names for sid in group_members.get(name, [])], + ) as groups: + collided = await _get_allowed_mcp_servers_from_mcp_server_names( + mcp_servers=["docs"], allowed_mcp_servers=[member] + ) + unmatched = await _get_allowed_mcp_servers_from_mcp_server_names( + mcp_servers=["team"], allowed_mcp_servers=[member] + ) + finally: + global_mcp_server_manager.registry.clear() + + assert [s.server_id for s in collided] == ["m-id"], ( + "a name owned by an ungranted server must still resolve to the access group of that name the key holds" + ) + assert unmatched == [], "a name matching neither a granted server nor a granted access group stays denied" + assert groups.await_args_list == [call(["docs"]), call(["team"])] + + +@pytest.mark.asyncio +async def test_scoped_router_hides_a_private_server_from_an_external_ip_like_the_connect_preflight(): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager + from litellm.proxy._experimental.mcp_server.server import _get_allowed_mcp_servers_from_mcp_server_names + + private = MCPServer( + server_id="p-id", + name="p", + server_name="p", + alias="gh", + transport=MCPTransport.http, + available_on_public_internet=False, + ) + public = MCPServer(server_id="u-id", name="u", server_name="u", alias="Gh", transport=MCPTransport.http) + global_mcp_server_manager.registry.clear() + global_mcp_server_manager.registry.update({"p-id": private, "u-id": public}) + try: + with patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp." + "MCPRequestHandler._get_mcp_servers_from_access_groups", + new_callable=AsyncMock, + return_value=[], + ): + external = await _get_allowed_mcp_servers_from_mcp_server_names( + mcp_servers=["gh"], allowed_mcp_servers=[public], client_ip="203.0.113.7" + ) + internal = await _get_allowed_mcp_servers_from_mcp_server_names( + mcp_servers=["gh"], allowed_mcp_servers=[public], client_ip=None + ) + by_own_name = await _get_allowed_mcp_servers_from_mcp_server_names( + mcp_servers=["Gh"], allowed_mcp_servers=[public], client_ip="203.0.113.7" + ) + assert global_mcp_server_manager.get_mcp_server_answering_to("gh", client_ip="203.0.113.7") is None + assert global_mcp_server_manager.get_mcp_server_answering_to("gh", client_ip=None) is private + finally: + global_mcp_server_manager.registry.clear() + + assert external == [], "a name the preflight hides from this IP must not reroute to a case variant" + assert [s.server_id for s in internal] == ["u-id"], "with no IP hiding in play the granted case variant wins" + assert [s.server_id for s in by_own_name] == ["u-id"] + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("registry", "scope", "granted", "expected"), + [ + pytest.param(("a-id", "d-id"), "docs", ("d-id",), ["d-id"], id="alias-collision-alias-holder-listed-first"), + pytest.param(("d-id", "a-id"), "docs", ("d-id",), ["d-id"], id="alias-collision-exact-name-listed-first"), + pytest.param(("g1", "g2"), "GITHUB", ("g2",), ["g2"], id="case-variant-collision"), + ], +) +async def test_scoped_router_prefers_the_granted_server_answering_to_the_name(registry, scope, granted, expected): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager + from litellm.proxy._experimental.mcp_server.server import _get_allowed_mcp_servers_from_mcp_server_names + + servers: Final = { + "a-id": MCPServer( + server_id="a-id", name="a_docs", server_name="a_docs", alias="docs", transport=MCPTransport.http + ), + "d-id": MCPServer(server_id="d-id", name="docs", server_name="docs", transport=MCPTransport.http), + "g1": MCPServer(server_id="g1", name="GitHub", server_name="GitHub", transport=MCPTransport.http), + "g2": MCPServer(server_id="g2", name="github", server_name="github", transport=MCPTransport.http), + } + global_mcp_server_manager.registry.clear() + global_mcp_server_manager.registry.update({server_id: servers[server_id] for server_id in registry}) + try: + with patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp." + "MCPRequestHandler._get_mcp_servers_from_access_groups", + new_callable=AsyncMock, + return_value=[], + ) as groups: + selected: Final = await _get_allowed_mcp_servers_from_mcp_server_names( + mcp_servers=[scope], allowed_mcp_servers=[servers[server_id] for server_id in granted] + ) + finally: + global_mcp_server_manager.registry.clear() + + assert [s.server_id for s in selected] == expected + groups.assert_not_awaited() + + +def test_get_mcp_server_answering_to_among_applies_the_registry_pass_order_and_ip_hiding(): + manager: Final = MCPServerManager() + by_alias: Final = MCPServer(server_id="a-id", name="a", server_name="a", alias="svc", transport=MCPTransport.http) + by_server_name: Final = MCPServer(server_id="b-id", name="b", server_name="svc", transport=MCPTransport.http) + by_name: Final = MCPServer(server_id="c-id", name="svc", server_name="c", transport=MCPTransport.http) + by_id: Final = MCPServer(server_id="svc", name="d", server_name="d", transport=MCPTransport.http) + folded: Final = MCPServer(server_id="e-id", name="e", server_name="SVC", transport=MCPTransport.http) + hidden: Final = MCPServer( + server_id="h-id", + name="h", + server_name="h", + alias="svc", + transport=MCPTransport.http, + available_on_public_internet=False, + ) + + assert manager.get_mcp_server_answering_to("svc", among=[by_name, by_server_name, by_alias]) is by_alias + assert manager.get_mcp_server_answering_to("svc", among=[by_name, by_server_name]) is by_server_name + assert manager.get_mcp_server_answering_to("svc", among=[by_id, by_name]) is by_name + assert manager.get_mcp_server_answering_to("svc", among=[folded, by_id]) is by_id + assert manager.get_mcp_server_answering_to("svc", among=[folded]) is folded + assert manager.get_mcp_server_answering_to("E-ID", among=[folded]) is folded + assert manager.get_mcp_server_answering_to("svc", among=()) is None + assert manager.get_mcp_server_answering_to("svc", client_ip="203.0.113.7", among=[hidden, folded]) is None + assert manager.get_mcp_server_answering_to("svc", client_ip="10.0.0.7", among=[hidden, folded]) is hidden + assert manager.get_mcp_server_answering_to("svc", client_ip="203.0.113.7", among=[folded]) is folded + assert manager.get_mcp_server_answering_to("svc") is None + + manager.registry = {"e-id": folded, "b-id": by_server_name, "a-id": by_alias} + + assert manager.get_mcp_server_answering_to("svc") is by_alias + assert manager.get_mcp_server_answering_to("svc") is manager.get_mcp_server_answering_to( + "svc", among=tuple(manager.registry.values()) + ) + + +class _GrantedServerSignInGuardrail(CustomGuardrail): + """Requires caller sign-in on one server only and records every server it is asked about.""" + + def __init__(self, *args, gated_server_id: str, **kwargs): + super().__init__(*args, **kwargs) + self._gated_server_id = gated_server_id + self.asked_about = [] # mutable-ok: call recorder + + def caller_sign_in(self, server, user_api_key_auth): + from litellm.proxy._experimental.mcp_server.caller_sign_in import CallerSignIn + + self.asked_about.append(server.server_id) + if server.server_id != self._gated_server_id: + return None + return CallerSignIn(issuers=("https://idp.test",), scopes=("scope-a",)) + + async def preflight_caller_sign_in(self, server, user_api_key_auth, subject_token): + from litellm.proxy._experimental.mcp_server.caller_sign_in import SignedIn + + return SignedIn() + + +class TestConnectPreflightRoutesLikeTheScopedRouter: + """A granted key connecting to ``/mcp/{name}`` is pre-flighted for the server the scoped router routes + it to, so an ungranted server holding the name as an alias neither hides the granted server's sign-in + challenge nor skips its connect-time exchange.""" + + @staticmethod + def _register_alias_collision() -> None: + alias_holder: Final = MCPServer( + server_id="a-id", + name="a_docs", + server_name="a_docs", + alias="docs", + url="https://a.test/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + ) + granted: Final = MCPServer( + server_id="d-id", + name="docs", + server_name="docs", + url="https://d.test/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + ) + mcp_operations.global_mcp_server_manager.registry.update({"a-id": alias_holder, "d-id": granted}) + + @staticmethod + async def _connect_to_docs() -> None: + from litellm.proxy._experimental.mcp_server import server as server_module + + await server_module._raise_preemptive_401_for_unauthenticated_servers( + scope={"type": "http", "method": "POST", "path": "/mcp/docs", "headers": []}, + mcp_servers=["docs"], + oauth2_headers=None, + mcp_server_auth_headers=None, + user_api_key_auth=UserAPIKeyAuth(api_key="sk-granted-docs", user_id="u-1"), + client_ip=None, + raw_headers={"x-litellm-api-key": "sk-granted-docs"}, + connecting=_connecting, + ) + + @pytest.mark.asyncio + async def test_exchange_runs_for_the_granted_server_not_the_alias_holder(self): + self._register_alias_collision() + + async def refuse_exchange(server, **kwargs): + raise HTTPException( + status_code=401, detail=f"exchange refused for {server.server_id} at {kwargs['resource_metadata']}" + ) + + with ( + patch.object( # test-quality-ok: the key's grant list lives in the DB; the real router runs on it + mcp_operations.global_mcp_server_manager, "get_allowed_mcp_servers", AsyncMock(return_value=["d-id"]) + ), + patch.object( # test-quality-ok: a real exchanger would call an IdP; this one reports which server it ran for + mcp_operations.global_mcp_server_manager, "preflight_token_exchange", refuse_exchange + ), + pytest.raises(HTTPException) as exc, + ): + await self._connect_to_docs() + + assert exc.value.status_code == 401 + assert exc.value.detail == "exchange refused for d-id at /.well-known/oauth-protected-resource/mcp/docs" + + @pytest.mark.asyncio + async def test_sign_in_challenge_names_the_granted_server_not_the_alias_holder(self, monkeypatch): + monkeypatch.delenv("SERVER_ROOT_PATH", raising=False) + self._register_alias_collision() + guardrail: Final = _GrantedServerSignInGuardrail(guardrail_name="sign-in-stub", gated_server_id="d-id") + litellm.logging_callback_manager.add_litellm_callback(guardrail) + try: + with ( + patch.object( # test-quality-ok: the key's grant list lives in the DB; the real router runs on it + mcp_operations.global_mcp_server_manager, + "get_allowed_mcp_servers", + AsyncMock(return_value=["d-id"]), + ), + pytest.raises(HTTPException) as exc, + ): + await self._connect_to_docs() + finally: + litellm.logging_callback_manager.remove_callback_from_list_by_object( + litellm.callbacks, guardrail, require_self=False + ) + + assert exc.value.status_code == 401 + assert ((exc.value.headers or {}).get("WWW-Authenticate") or "").startswith( + 'Bearer resource_metadata="/.well-known/oauth-protected-resource/mcp/docs"' + ) + assert guardrail.asked_about == ["d-id"] + + @staticmethod + def _register_obo_alias_collision() -> None: + obo: Final = MCPServer( + server_id="o-id", + name="obo_server", + server_name="obo_server", + alias="obo", + url="https://obo.test/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2_token_exchange, + token_exchange_endpoint="https://idp.test/token", + client_id="cid", + client_secret="csecret", + ) + plain: Final = MCPServer( + server_id="p-id", + name="obo", + server_name="obo", + url="https://plain.test/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + ) + mcp_operations.global_mcp_server_manager.registry.update({"o-id": obo, "p-id": plain}) + + @staticmethod + async def _connect_to(route_name: str) -> None: + from litellm.proxy._experimental.mcp_server import server as server_module + + 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-collision", user_id="u-1"), + client_ip=None, + raw_headers={"x-litellm-api-key": "sk-collision"}, + connecting=_connecting, + ) + + @pytest.mark.asyncio + @pytest.mark.parametrize("route_name", ["obo", "OBO"], ids=["exact_alias", "case_variant"]) + async def test_granted_plain_server_connects_past_an_ungranted_obo_alias_holder(self, route_name): + self._register_obo_alias_collision() + + async def report_exchange(server, **kwargs): + raise HTTPException( + status_code=401, detail=f"exchange ran for {server.server_id} at {kwargs['resource_metadata']}" + ) + + with ( + patch.object( # test-quality-ok: the key's grant list lives in the DB; the real router runs on it + mcp_operations.global_mcp_server_manager, "get_allowed_mcp_servers", AsyncMock(return_value=["p-id"]) + ), + patch.object( # test-quality-ok: a real exchanger would call an IdP; this one reports which server it ran for + mcp_operations.global_mcp_server_manager, "preflight_token_exchange", report_exchange + ), + pytest.raises(HTTPException) as exc, + ): + await self._connect_to(route_name) + + assert exc.value.status_code == 401 + assert exc.value.detail == f"exchange ran for p-id at /.well-known/oauth-protected-resource/mcp/{route_name}" + + @pytest.mark.asyncio + @pytest.mark.parametrize("route_name", ["obo", "obo_server"], ids=["alias", "server_name"]) + async def test_granted_obo_server_still_challenges_without_a_subject(self, route_name, monkeypatch): + monkeypatch.delenv("SERVER_ROOT_PATH", raising=False) + self._register_obo_alias_collision() + with ( + patch.object( # test-quality-ok: the key's grant list lives in the DB; the real router runs on it + mcp_operations.global_mcp_server_manager, "get_allowed_mcp_servers", AsyncMock(return_value=["o-id"]) + ), + pytest.raises(HTTPException) as exc, + ): + await self._connect_to(route_name) + + assert exc.value.status_code == 401 + assert ((exc.value.headers or {}).get("WWW-Authenticate") or "").startswith( + f'Bearer resource_metadata="/.well-known/oauth-protected-resource/mcp/{route_name}"' + ) + + @pytest.mark.asyncio + @pytest.mark.parametrize("route_name", ["obo", "OBO"], ids=["exact_alias", "case_variant"]) + async def test_key_granting_neither_server_is_not_preflighted_for_the_plain_one(self, route_name): + self._register_obo_alias_collision() + preflight: Final = AsyncMock() + with ( + patch.object( # test-quality-ok: the key's grant list lives in the DB; the real router runs on it + mcp_operations.global_mcp_server_manager, "get_allowed_mcp_servers", AsyncMock(return_value=[]) + ), + patch.object( # test-quality-ok: a real exchanger would call an IdP; this one records which server ran + mcp_operations.global_mcp_server_manager, "preflight_token_exchange", preflight + ), + pytest.raises(HTTPException) as exc, + ): + await self._connect_to(route_name) + + assert exc.value.status_code == 401 + assert preflight.await_count == 0 + + @pytest.mark.asyncio + async def test_access_group_named_like_an_ungranted_server_connects_like_the_merge_base(self): + from litellm.proxy._experimental.mcp_server import server as server_module + + member: Final = MCPServer( + server_id="w-id", + name="wiki_obo", + server_name="wiki_obo", + url="https://wiki-obo.test/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2_token_exchange, + token_exchange_endpoint="https://idp.test/token", + client_id="cid", + client_secret="csecret", + ) + shadow: Final = MCPServer( + server_id="s-id", + name="wiki", + server_name="wiki", + url="https://wiki.test/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + ) + mcp_operations.global_mcp_server_manager.registry.update({"w-id": member, "s-id": shadow}) + group_members: Final = {"wiki": ["w-id"]} + with ( + patch.object( # test-quality-ok: the key's grant list lives in the DB; the real router runs on it + mcp_operations.global_mcp_server_manager, "get_allowed_mcp_servers", AsyncMock(return_value=["w-id"]) + ), + patch( # test-quality-ok: access group membership lives in the DB; the real scoped router runs on it + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp." + "MCPRequestHandler._get_mcp_servers_from_access_groups", + new_callable=AsyncMock, + side_effect=lambda names: [sid for name in names for sid in group_members.get(name, [])], + ), + ): + outcome = await server_module._raise_preemptive_401_for_unauthenticated_servers( + scope={"type": "http", "method": "POST", "path": "/mcp/wiki", "headers": [(b"host", b"testserver")]}, + mcp_servers=["wiki"], + oauth2_headers=None, + mcp_server_auth_headers=None, + user_api_key_auth=UserAPIKeyAuth(api_key="sk-group", user_id="u-1"), + client_ip=None, + raw_headers={"x-litellm-api-key": "sk-group"}, + connecting=_connecting, + ) + + assert outcome is None, ( + "a key granted only the access group named like an ungranted plain server must connect without a " + "sign-in challenge, as at the merge base; the challenge would advertise the plain server's metadata" + ) + + @pytest.mark.asyncio async def test_get_allowed_mcp_servers_from_mcp_server_names_mixed_known_and_unknown(): """ @@ -9529,7 +10505,7 @@ async def test_call_tool_with_legacy_db_m2m_server_resolves_oauth2_flow(): ), patch( "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers_from_mcp_server_names", - new=AsyncMock(side_effect=lambda mcp_servers, allowed_mcp_servers: allowed_mcp_servers), + new=AsyncMock(side_effect=lambda mcp_servers, allowed_mcp_servers, client_ip=None: allowed_mcp_servers), ), ): mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["legacy-m2m-id"]) @@ -9823,6 +10799,7 @@ async def test_aggregate_listing_reports_per_server_outcomes(): mock_manager.get_mcp_server_by_id = lambda server_id: ( working_server if server_id == "working_server" else broken_server ) + mock_manager.get_mcp_server_answering_to = MCPServerManager().get_mcp_server_answering_to mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (server_ids, 0) async def mock_get_tools_from_server(server, **kwargs): @@ -9978,7 +10955,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( @@ -9995,8 +10972,116 @@ class TestPreemptive401ModeAware: mcp_server_auth_headers=None, user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key"), client_ip=None, + connecting=_connecting, ) + async def _connect_with_a_grant(self, server, requested: str, path: str) -> HTTPException: + from litellm.proxy._experimental.mcp_server import server as server_module + + manager = mcp_operations.global_mcp_server_manager + manager.registry.clear() + manager.registry[server.server_id] = server + with ( + patch.object( # test-quality-ok: the key's grant list lives in the DB; the real router runs on it + manager, "get_allowed_mcp_servers", AsyncMock(return_value=[server.server_id]) + ), + patch.object(manager, "has_user_oauth_token", new_callable=AsyncMock, return_value=False), + pytest.raises(HTTPException) as exc, + ): + await server_module._raise_preemptive_401_for_unauthenticated_servers( + scope={"type": "http", "method": "POST", "path": path, "headers": [(b"host", b"testserver")]}, + mcp_servers=[requested], + oauth2_headers=None, + mcp_server_auth_headers=None, + user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key"), + client_ip=None, + connecting=_connecting, + ) + return exc.value + + @pytest.mark.asyncio + @pytest.mark.parametrize("delegate", [True, False], ids=["oauth_delegate", "gateway_interactive"]) + @pytest.mark.parametrize("shape", ["alias_case", "server_id", "x_mcp_servers"]) + async def test_moved_connect_shapes_get_the_exact_name_routes_challenge(self, delegate, shape, monkeypatch): + """Alias-case, server-id and x-mcp-servers connects resolve the same server the router serves, so + they answer the exact-name route's 401 with the requested spelling in the route segment.""" + monkeypatch.delenv("SERVER_ROOT_PATH", raising=False) + server = _make_oauth2_server("gwx", oauth2_flow="authorization_code", delegate_auth_to_upstream=delegate) + requested, path, exact_path = { + "alias_case": ("GWX", "/mcp/GWX", "/mcp/gwx"), + "server_id": (server.server_id, f"/mcp/{server.server_id}", "/mcp/gwx"), + "x_mcp_servers": ("GWX", "/mcp", "/mcp"), + }[shape] + + exact = await self._connect_with_a_grant(server, "gwx", exact_path) + moved = await self._connect_with_a_grant(server, requested, path) + + assert exact.status_code == 401 + assert (moved.status_code, moved.detail) == (exact.status_code, exact.detail) + exact_header = {k.lower(): v for k, v in (exact.headers or {}).items()}["www-authenticate"] + moved_header = {k.lower(): v for k, v in (moved.headers or {}).items()}["www-authenticate"] + assert "/gwx" in exact_header + assert moved_header == exact_header.replace("/gwx", f"/{requested}") + + @pytest.mark.asyncio + @pytest.mark.parametrize("kind", ["plain_obo", "oauth_passthrough"]) + @pytest.mark.parametrize("shape", ["alias_case", "server_id", "x_mcp_servers"]) + async def test_moved_obo_and_passthrough_shapes_get_the_exact_name_routes_challenge(self, kind, shape, monkeypatch): + monkeypatch.delenv("SERVER_ROOT_PATH", raising=False) + server = ( + _make_obo_server("obx") + if kind == "plain_obo" + else MCPServer( + server_id="id-obx", + name="obx", + alias="obx", + server_name="obx", + url="https://obx.test/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + extra_headers=["Authorization"], + oauth_passthrough=True, + mcp_info={"server_name": "obx"}, + ) + ) + requested, path, exact_path = { + "alias_case": ("OBX", "/mcp/OBX", "/mcp/obx"), + "server_id": (server.server_id, f"/mcp/{server.server_id}", "/mcp/obx"), + "x_mcp_servers": ("OBX", "/mcp", "/mcp"), + }[shape] + + exact = await self._connect_with_a_grant(server, "obx", exact_path) + moved = await self._connect_with_a_grant(server, requested, path) + + assert exact.status_code == 401 + assert (moved.status_code, moved.detail) == (exact.status_code, exact.detail) + exact_header = {k.lower(): v for k, v in (exact.headers or {}).items()}["www-authenticate"] + moved_header = {k.lower(): v for k, v in (moved.headers or {}).items()}["www-authenticate"] + assert "/obx" in exact_header + assert moved_header == exact_header.replace("/obx", f"/{requested}") + + @pytest.mark.asyncio + async def test_aggregate_connect_without_a_server_selection_is_not_challenged(self): + from litellm.proxy._experimental.mcp_server import server as server_module + + manager = mcp_operations.global_mcp_server_manager + manager.registry.clear() + for alias, delegate in (("gwx", False), ("relay", True)): + server = _make_oauth2_server(alias, oauth2_flow="authorization_code", delegate_auth_to_upstream=delegate) + manager.registry[server.server_id] = server + + outcome = await server_module._raise_preemptive_401_for_unauthenticated_servers( + scope={"type": "http", "method": "POST", "path": "/mcp", "headers": [(b"host", b"testserver")]}, + mcp_servers=None, + oauth2_headers=None, + mcp_server_auth_headers=None, + user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key"), + client_ip=None, + connecting=_connecting, + ) + + assert outcome is None, "an unselected aggregate connect must pass without a sign-in challenge" + @pytest.mark.asyncio async def test_deferred_discovery_runs_before_delegate_challenge(self): from litellm.proxy._experimental.mcp_server import server as server_module @@ -10077,7 +11162,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( @@ -10095,6 +11180,7 @@ class TestPreemptive401ModeAware: mcp_server_auth_headers=None, user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key"), client_ip=None, + connecting=_connecting, ) assert exc.value.status_code == 401 @@ -10184,7 +11270,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 @@ -10203,6 +11289,7 @@ class TestSingleServerPreflightReachesIdJag: mcp_server_auth_headers=None, user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"), client_ip=None, + connecting=_connecting, ) @pytest.mark.asyncio @@ -10244,7 +11331,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 @@ -10261,6 +11348,7 @@ class TestSingleServerPreflightReachesIdJag: mcp_server_auth_headers=None, user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"), client_ip=None, + connecting=_connecting, ) assert exc.value.status_code == 401 @@ -10309,7 +11397,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 @@ -10325,6 +11413,11 @@ 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"], + }, + connecting=_connecting, ) return allowed_lookup, preflight @@ -10333,9 +11426,7 @@ class TestOboPreflightScopedToAllowedServers: requested = _make_obo_server("obo_tools") key = UserAPIKeyAuth(api_key="sk-plain-only") - allowed_lookup, preflight = await self._run( - requested, allowed=[_make_obo_server("plain_tools")], user_api_key_auth=key - ) + allowed_lookup, preflight = await self._run(requested, allowed=[], user_api_key_auth=key) preflight.assert_not_awaited() allowed_lookup.assert_awaited_once_with( @@ -10350,10 +11441,62 @@ 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"], + }, + resource_metadata=f"/.well-known/oauth-protected-resource/mcp/{requested.alias}", ) +class TestOboChallengeGateKeepsBaseConnectRules: + """An OBO connect carrying a bearer in oauth2_headers is challenged by the exchange path, not the + preemptive gate, so a multi-server connect or a single-server connect with any bearer at all must + not be refused before the session opens.""" + + LITELLM_KEY_BEARER = {"Authorization": "Bearer sk-1234"} + + async def _run(self, servers: list[MCPServer], mcp_servers: list[str]) -> None: + from litellm.proxy._experimental.mcp_server import server as server_module + + with ( + patch.object( # test-quality-ok: route wiring must use the manager's configured server + mcp_operations.global_mcp_server_manager, + "get_mcp_server_answering_to", + return_value=servers[0], + ), + patch.object( # test-quality-ok: allowed-set resolution needs the DB; the test controls its answer + mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=servers) + ), + ): + await server_module._raise_preemptive_401_for_unauthenticated_servers( + scope={"type": "http", "method": "POST", "path": "/mcp/obo", "headers": []}, + mcp_servers=mcp_servers, + oauth2_headers=self.LITELLM_KEY_BEARER, + mcp_server_auth_headers=None, + user_api_key_auth=UserAPIKeyAuth(api_key="sk-1234", user_id="u-1"), + client_ip=None, + raw_headers={"x-litellm-api-key": "sk-1234", "authorization": "Bearer sk-1234"}, + connecting=_connecting, + ) + + @pytest.mark.asyncio + async def test_multi_server_connect_with_any_bearer_is_not_preemptively_challenged(self): + obo = _make_obo_server("obo") + catalog = MCPServer(server_id="id-catalog", name="catalog", alias="catalog", transport=MCPTransport.http) + + await self._run([obo, catalog], ["obo", "catalog"]) + + @pytest.mark.asyncio + async def test_single_obo_connect_with_litellm_key_bearer_still_challenges(self): + with pytest.raises(HTTPException) as exc: + await self._run([_make_obo_server("obo")], ["obo"]) + assert exc.value.status_code == 401 + + @pytest.mark.asyncio async def test_post_mcp_call_guardrails_return_the_rewritten_result(): """The result a post_mcp_call guardrail rewrote must be what the caller sends back.""" @@ -10871,8 +12014,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") @@ -10908,8 +12051,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() @@ -11067,6 +12211,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()) @@ -11085,3 +12230,663 @@ 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 _connect_scope( + path: str, *, headers: list[tuple[bytes, bytes]] | None = None, client_ip: str = "10.0.0.7" +) -> dict[str, object]: + return { + "type": "http", + "method": "POST", + "scheme": "http", + "path": path, + "root_path": "", + "query_string": b"", + "server": ("gw.example", 4000), + "client": (client_ip, 51000), + "headers": [(b"host", b"gw.example:4000"), *(headers or [])], + } + + +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 __init__(self, *args, preflight_result=None, **kwargs): + super().__init__(*args, **kwargs) + self._preflight_result = preflight_result + self.preflight_calls = [] # mutable-ok: call recorder + + 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",)) + + async def preflight_caller_sign_in(self, server, user_api_key_auth, subject_token): + from litellm.proxy._experimental.mcp_server.caller_sign_in import SignedIn + + self.preflight_calls.append(subject_token) + return self._preflight_result if self._preflight_result is not None else SignedIn() + + +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, + connecting=_connecting, + ) + 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, + connecting=_connecting, + ) + 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 + @pytest.mark.parametrize( + ("path", "route_name", "metadata_url"), + [ + ("/mcp/obo", "obo", "http://gw.example:4000/.well-known/oauth-protected-resource/mcp/obo"), + ( + "/mcp/obo_server", + "obo_server", + "http://gw.example:4000/.well-known/oauth-protected-resource/mcp/obo_server", + ), + ( + "/obo_server/mcp", + "obo_server", + "http://gw.example:4000/.well-known/oauth-protected-resource/obo_server/mcp", + ), + ], + ids=["alias_route", "server_name_route", "server_first_route"], + ) + async def test_obo_challenge_names_the_absolute_metadata_url_of_the_connected_route( + self, monkeypatch, path, route_name, metadata_url + ): + """RFC 9728 5.1 makes resource_metadata a URL and the MCP SDK fetches it verbatim, then refuses a + document whose ``resource`` does not prefix-match the URL it connected to. The OBO challenge must + therefore advertise the absolute metadata URL of the route the client used, not the alias's path.""" + from litellm.proxy._experimental.mcp_server import server as server_module + + monkeypatch.delenv("SERVER_ROOT_PATH", raising=False) + monkeypatch.delenv("PROXY_BASE_URL", raising=False) + obo = _make_obo_server("obo").model_copy(update={"name": "obo_server", "server_name": "obo_server"}) + 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=_connect_scope(path), + 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, + connecting=_connecting, + ) + + assert exc.value.status_code == 401 + assert (exc.value.headers or {}).get("WWW-Authenticate") == ( + f'Bearer resource_metadata="{metadata_url}", ' + 'error="invalid_token", ' + 'error_description="Missing or invalid subject token; authenticate with the IdP and retry"' + ) + + @pytest.mark.asyncio + @pytest.mark.parametrize( + ("path", "general_settings", "client_ip", "metadata_url"), + [ + ("/mcp/catalog", {}, "10.0.0.7", "http://gw.example:4000/.well-known/oauth-protected-resource/mcp/catalog"), + ( + "/catalog/mcp", + {}, + "10.0.0.7", + "http://gw.example:4000/.well-known/oauth-protected-resource/catalog/mcp", + ), + ( + "/catalog/mcp", + {"use_x_forwarded_for": True, "mcp_trusted_proxy_ranges": ["10.0.0.0/8"]}, + "10.0.0.7", + "https://public.example/.well-known/oauth-protected-resource/catalog/mcp", + ), + ( + "/catalog/mcp", + {"use_x_forwarded_for": True, "mcp_trusted_proxy_ranges": ["10.0.0.0/8"]}, + "203.0.113.9", + "http://gw.example:4000/.well-known/oauth-protected-resource/catalog/mcp", + ), + ], + ids=["mcp_first", "server_first", "forwarded_from_trusted_proxy", "forwarded_from_untrusted_client"], + ) + async def test_provider_challenge_names_the_absolute_metadata_url_of_the_connected_route( + self, monkeypatch, path, general_settings, client_ip, metadata_url + ): + """The sign-in challenge must point at the metadata document of the route the client used + (``/catalog/mcp`` and ``/mcp/catalog`` are distinct documents with distinct ``resource`` values), + built on the public origin only when the forwarded headers come from a configured trusted proxy.""" + from litellm.proxy._experimental.mcp_server import server as server_module + + monkeypatch.delenv("SERVER_ROOT_PATH", raising=False) + monkeypatch.delenv("PROXY_BASE_URL", raising=False) + server = _catalog_server() + guardrail = _CallerSignInGuardrail(guardrail_name="sign-in-stub") + litellm.logging_callback_manager.add_litellm_callback(guardrail) + forwarded = [(b"x-forwarded-proto", b"https"), (b"x-forwarded-host", b"public.example")] + try: + with ( + patch("litellm.proxy.proxy_server.general_settings", general_settings, create=True), + 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=_connect_scope(path, headers=forwarded, client_ip=client_ip), + mcp_servers=["catalog"], + 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=client_ip, + connecting=_connecting, + ) + finally: + litellm.logging_callback_manager.remove_callback_from_list_by_object( + litellm.callbacks, guardrail, require_self=False + ) + + assert exc.value.status_code == 401 + assert ( + (exc.value.headers or {}) + .get("WWW-Authenticate", "") + .startswith(f'Bearer resource_metadata="{metadata_url}"') + ) + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "route_name", + ["catalog-server-id-001", "CATALOG"], + ids=["server_id", "uppercase_name"], + ) + async def test_challenge_resource_metadata_names_the_connected_segment(self, route_name): + """The PRM path in the challenge must name the segment the client connected with, or the + client's follow-up metadata fetch 404s against the route it was pointed at.""" + 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) + 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, + connecting=_connecting, + ) + finally: + litellm.logging_callback_manager.remove_callback_from_list_by_object( + litellm.callbacks, guardrail, require_self=False + ) + + authenticate: Final = (exc.value.headers or {}).get("WWW-Authenticate") or "" + assert authenticate.startswith( + f'Bearer resource_metadata="/.well-known/oauth-protected-resource/mcp/{route_name}"' + ) + + +class TestConnectSignInPreflight: + """A subject token the provider rejects must fail at connect with the RFC 9728 challenge, not as a + JSON-RPC error on every tools/call where ``WWW-Authenticate`` is lost.""" + + async def _connect( + self, route_names, guardrail, allowed, raw_headers=None, connecting=_connecting, user_api_key_auth=None + ): + from litellm.proxy._experimental.mcp_server import server as server_module + + server = _catalog_server() + 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=allowed), + ), + ): + await server_module._raise_preemptive_401_for_unauthenticated_servers( + scope={"type": "http", "method": "POST", "path": f"/mcp/{route_names[0]}", "headers": []}, + mcp_servers=list(route_names), + oauth2_headers=None, + mcp_server_auth_headers=None, + user_api_key_auth=user_api_key_auth or UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"), + client_ip=None, + raw_headers=raw_headers + or {"x-litellm-api-key": "sk-litellm-virtual-key", "authorization": "Bearer entra.jwt.token"}, + connecting=connecting, + ) + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "authorization", + ["Basic a.b.c", "Digest x.y.z", "eyJ.pay.sig"], + ids=["basic_scheme", "digest_scheme", "scheme_less"], + ) + async def test_authorization_without_a_bearer_scheme_is_challenged_not_pre_flighted(self, authorization): + """Only a ``Bearer`` credential is a sign-in subject; a Basic or Digest value, or a bare string that + merely has two dots, is challenged locally and never handed to a provider's IdP.""" + server = _catalog_server() + guardrail = _CallerSignInGuardrail(guardrail_name="sign-in-stub") + litellm.logging_callback_manager.add_litellm_callback(guardrail) + try: + with pytest.raises(HTTPException) as exc: + await self._connect( + ["catalog"], + guardrail, + [server], + raw_headers={"x-litellm-api-key": "sk-litellm-virtual-key", "authorization": authorization}, + ) + finally: + litellm.logging_callback_manager.remove_callback_from_list_by_object( + litellm.callbacks, guardrail, require_self=False + ) + + assert exc.value.status_code == 401 + assert 'error="invalid_token"' in ((exc.value.headers or {}).get("WWW-Authenticate") or "") + assert guardrail.preflight_calls == [] + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "raw_headers", + [ + {"authorization": "Bearer gw.master.key"}, + {"x-litellm-api-key": "sk-litellm-virtual-key", "authorization": "Bearer gw.master.key"}, + ], + ids=["master_key_alone", "master_key_next_to_a_virtual_key"], + ) + async def test_dotted_master_key_bearer_is_challenged_never_pre_flighted(self, raw_headers): + """The master key is a LiteLLM credential even when it has the two dots of a compact JWS, so the + connect answers the local challenge instead of sending the key to a provider's IdP.""" + server = _catalog_server() + guardrail = _CallerSignInGuardrail(guardrail_name="sign-in-stub") + litellm.logging_callback_manager.add_litellm_callback(guardrail) + try: + with ( + patch("litellm.proxy.proxy_server.master_key", "gw.master.key"), + pytest.raises(HTTPException) as exc, + ): + await self._connect(["catalog"], guardrail, [server], raw_headers=raw_headers) + finally: + litellm.logging_callback_manager.remove_callback_from_list_by_object( + litellm.callbacks, guardrail, require_self=False + ) + + assert exc.value.status_code == 401 + assert 'error="invalid_token"' in ((exc.value.headers or {}).get("WWW-Authenticate") or "") + assert guardrail.preflight_calls == [] + + @pytest.mark.asyncio + async def test_custom_auth_admitted_bearer_is_pre_flighted_not_challenged(self): + """Custom auth admits the caller on its own IdP token in ``Authorization`` with no + ``x-litellm-api-key``; that token is the sign-in subject, so connect pre-flights it as a tool + call forwarded it before the subject split, instead of challenging for a missing subject.""" + server = _catalog_server() + guardrail = _CallerSignInGuardrail(guardrail_name="sign-in-stub") + admitted = UserAPIKeyAuth(api_key="entra.jwt.token", user_id="u-1") + admitted.authenticated_by_custom_auth = True + litellm.logging_callback_manager.add_litellm_callback(guardrail) + try: + await self._connect( + ["catalog"], + guardrail, + [server], + raw_headers={"authorization": "Bearer entra.jwt.token"}, + user_api_key_auth=admitted, + ) + finally: + litellm.logging_callback_manager.remove_callback_from_list_by_object( + litellm.callbacks, guardrail, require_self=False + ) + + assert guardrail.preflight_calls == ["entra.jwt.token"] + + @pytest.mark.asyncio + async def test_custom_auth_admitted_virtual_key_bearer_is_still_challenged(self): + server = _catalog_server() + guardrail = _CallerSignInGuardrail(guardrail_name="sign-in-stub") + admitted = UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1") + admitted.authenticated_by_custom_auth = True + litellm.logging_callback_manager.add_litellm_callback(guardrail) + try: + with pytest.raises(HTTPException) as exc: + await self._connect( + ["catalog"], + guardrail, + [server], + raw_headers={"authorization": "Bearer sk-litellm-virtual-key"}, + user_api_key_auth=admitted, + ) + finally: + litellm.logging_callback_manager.remove_callback_from_list_by_object( + litellm.callbacks, guardrail, require_self=False + ) + + assert exc.value.status_code == 401 + assert 'error="invalid_token"' in ((exc.value.headers or {}).get("WWW-Authenticate") or "") + assert guardrail.preflight_calls == [] + + @pytest.mark.asyncio + async def test_built_in_oauth2_admitted_bearer_is_pre_flighted_not_challenged(self): + """The built-in OAuth2 admission records the caller's token as ``api_key`` without the custom-auth + marker; that token is still the sign-in subject, so connect pre-flights it instead of challenging.""" + server = _catalog_server() + guardrail = _CallerSignInGuardrail(guardrail_name="sign-in-stub") + litellm.logging_callback_manager.add_litellm_callback(guardrail) + try: + await self._connect( + ["catalog"], + guardrail, + [server], + raw_headers={"authorization": "Bearer entra.jwt.token"}, + user_api_key_auth=UserAPIKeyAuth(api_key="entra.jwt.token", user_id="u-1"), + ) + finally: + litellm.logging_callback_manager.remove_callback_from_list_by_object( + litellm.callbacks, guardrail, require_self=False + ) + + assert guardrail.preflight_calls == ["entra.jwt.token"] + + @pytest.mark.asyncio + async def test_rejected_subject_on_an_open_session_is_left_to_the_tool_call_hook(self): + """Only the connect pre-flights the subject. On an open session the gate must not exchange at all, so + the tool-call hook runs the one exchange and answers a rejection inside the JSON-RPC envelope with its + guardrail Logs row, as base did.""" + from litellm.proxy._experimental.mcp_server.caller_sign_in import Rejected + + async def _open_session() -> bool: + return False + + server = _catalog_server() + guardrail = _CallerSignInGuardrail( + guardrail_name="sign-in-stub", + preflight_result=Rejected("the Entra OBO exchange was rejected (AADSTS5002723)"), + ) + litellm.logging_callback_manager.add_litellm_callback(guardrail) + try: + await self._connect(["catalog"], guardrail, [server], connecting=_open_session) + finally: + litellm.logging_callback_manager.remove_callback_from_list_by_object( + litellm.callbacks, guardrail, require_self=False + ) + + assert guardrail.preflight_calls == [] + + @pytest.mark.asyncio + async def test_rejected_subject_challenges_at_connect(self): + from litellm.proxy._experimental.mcp_server.caller_sign_in import Rejected + + server = _catalog_server() + guardrail = _CallerSignInGuardrail( + guardrail_name="sign-in-stub", + preflight_result=Rejected("the Entra OBO exchange was rejected (AADSTS70002)"), + ) + litellm.logging_callback_manager.add_litellm_callback(guardrail) + try: + with pytest.raises(HTTPException) as exc: + await self._connect(["catalog"], guardrail, [server]) + finally: + litellm.logging_callback_manager.remove_callback_from_list_by_object( + litellm.callbacks, guardrail, require_self=False + ) + + assert exc.value.status_code == 401 + authenticate: Final = (exc.value.headers or {}).get("WWW-Authenticate") or "" + assert 'resource_metadata="/.well-known/oauth-protected-resource/mcp/catalog"' in authenticate + assert 'error="invalid_token"' in authenticate + assert guardrail.preflight_calls == ["entra.jwt.token"] + + @pytest.mark.asyncio + async def test_unavailable_fail_closed_answers_503(self): + from litellm.proxy._experimental.mcp_server.caller_sign_in import Unavailable + + server = _catalog_server() + guardrail = _CallerSignInGuardrail( + guardrail_name="sign-in-stub", + preflight_result=Unavailable("the Entra token endpoint could not be reached", fail_open=False), + ) + litellm.logging_callback_manager.add_litellm_callback(guardrail) + try: + with pytest.raises(HTTPException) as exc: + await self._connect(["catalog"], guardrail, [server]) + finally: + litellm.logging_callback_manager.remove_callback_from_list_by_object( + litellm.callbacks, guardrail, require_self=False + ) + + assert exc.value.status_code == 503 + assert exc.value.detail == "the Entra token endpoint could not be reached" + + @pytest.mark.asyncio + @pytest.mark.parametrize( + ("rpc_method", "session_id", "reaches"), + [ + ("initialize", None, None), + ("tools/call", "open-session-1", "stateful"), + ("tools/call", None, "stateless"), + ], + ids=["connect_503", "open_session_tools_call", "stateless_tools_call"], + ) + async def test_fail_closed_outage_is_the_connects_503_only_on_initialize(self, rpc_method, session_id, reaches): + """Only the ``initialize`` POST turns a fail-closed provider outage into the connect's 503. Every other + JSON-RPC POST on the gated route must reach the session manager with its body intact and without a gate + exchange, so the tools/call hook runs the one exchange, answers the outage inside the result envelope and + writes the guardrail Logs row, as base did.""" + from litellm.proxy._experimental.mcp_server import server as server_module + from litellm.proxy._experimental.mcp_server.caller_sign_in import Unavailable + + server = _catalog_server() + guardrail = _CallerSignInGuardrail( + guardrail_name="sign-in-stub", + preflight_result=Unavailable("Entra rejected the gateway's own Agent 365 credentials", fail_open=False), + ) + body = json.dumps({"jsonrpc": "2.0", "id": 7, "method": rpc_method, "params": {}}).encode() + session_headers = [(b"mcp-session-id", session_id.encode())] if session_id else [] + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/catalog", + "headers": [(b"content-type", b"application/json"), (b"authorization", b"Bearer entra.jwt.token")] + + session_headers, + } + receive = AsyncMock(return_value={"type": "http.request", "body": body, "more_body": False}) + send = AsyncMock() + delivered: dict[str, bytes] = {} # mutable-ok: records which manager saw the replayed body + + def _recorder(manager: str): + async def handle(_scope, replayed_receive, _send): + delivered[manager] = (await replayed_receive())["body"] + + return handle + + litellm.logging_callback_manager.add_litellm_callback(guardrail) + try: + with ( + patch.object( + server_module, + "extract_mcp_auth_context", + AsyncMock( + return_value=( + UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"), + None, + ["catalog"], + None, + None, + {"x-litellm-api-key": "sk-litellm-virtual-key", "authorization": "Bearer entra.jwt.token"}, + ) + ), + ), + 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])), + patch.object(server_module, "_SESSION_MANAGERS_INITIALIZED", True), + patch.object(server_module.session_manager_stateful, "handle_request", _recorder("stateful")), + patch.object(server_module.session_manager_stateless, "handle_request", _recorder("stateless")), + patch.object( + server_module.session_manager_stateful, + "_server_instances", + {session_id: MagicMock()} if session_id else {}, + ), + ): + if reaches is None: + with pytest.raises(HTTPException) as exc: + await server_module.handle_streamable_http_mcp(scope, receive, send) + assert exc.value.status_code == 503 + assert exc.value.detail == "Entra rejected the gateway's own Agent 365 credentials" + else: + await server_module.handle_streamable_http_mcp(scope, receive, send) + finally: + litellm.logging_callback_manager.remove_callback_from_list_by_object( + litellm.callbacks, guardrail, require_self=False + ) + if session_id: + server_module._remove_stateful_session_tracking(session_id) + + assert guardrail.preflight_calls == (["entra.jwt.token"] if reaches is None else []) + assert delivered == ({} if reaches is None else {reaches: body}) + send.assert_not_awaited() + + @pytest.mark.asyncio + async def test_multi_server_connect_never_awaits_the_preflight(self): + server = _catalog_server() + guardrail = _CallerSignInGuardrail(guardrail_name="sign-in-stub") + litellm.logging_callback_manager.add_litellm_callback(guardrail) + try: + await self._connect(["catalog", "other"], guardrail, [server]) + finally: + litellm.logging_callback_manager.remove_callback_from_list_by_object( + litellm.callbacks, guardrail, require_self=False + ) + + assert guardrail.preflight_calls == [] diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_stale_session.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_stale_session.py index eb1d8573ee3..4cbf8a75b74 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_stale_session.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_stale_session.py @@ -650,6 +650,10 @@ async def test_per_user_oauth_missing_stored_token_returns_preemptive_401(): "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_by_name", return_value=oauth_server, ), + patch( + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_answering_to", + return_value=oauth_server, + ), patch.object( session_manager_stateless, "handle_request", @@ -738,6 +742,10 @@ async def test_admitted_subject_missing_stored_token_challenged_with_resource_me "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_by_name", return_value=oauth_server, ), + patch( + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_answering_to", + return_value=oauth_server, + ), patch.object( session_manager_stateless, "handle_request", @@ -1033,6 +1041,10 @@ async def test_per_user_oauth_with_stored_token_skips_preemptive_401(): "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_by_name", return_value=oauth_server, ), + patch( + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_answering_to", + return_value=oauth_server, + ), patch.object( session_manager_stateless, "handle_request", @@ -1136,6 +1148,10 @@ async def test_handle_streamable_http_mcp_delegated_server_without_token_returns "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_by_name", return_value=delegated_server, ), + patch( + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_answering_to", + return_value=delegated_server, + ), patch.object( session_manager_stateful, "handle_request", @@ -1224,6 +1240,10 @@ async def test_handle_streamable_http_mcp_token_exchange_without_subject_returns "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_by_name", return_value=obo_server, ), + patch( + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_answering_to", + return_value=obo_server, + ), patch.object( session_manager_stateful, "handle_request", @@ -1323,6 +1343,10 @@ async def test_handle_streamable_http_mcp_oauth_delegate_without_token_returns_g "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_by_name", return_value=od_server, ), + patch( + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_answering_to", + return_value=od_server, + ), patch.object( session_manager_stateful, "handle_request", @@ -1580,6 +1604,10 @@ async def test_handle_streamable_http_mcp_true_passthrough_without_token_surface "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_by_name", return_value=tp_server, ), + patch( + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_answering_to", + return_value=tp_server, + ), patch.object( session_manager_stateful, "handle_request", @@ -1648,6 +1676,10 @@ async def test_handle_streamable_http_mcp_true_passthrough_dcr_bridge_challenges "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_by_name", return_value=bridge_server, ), + patch( + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_answering_to", + return_value=bridge_server, + ), patch.object( session_manager_stateful, "handle_request", diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/test_agent_365.py b/tests/unit/proxy/guardrails/guardrail_hooks/test_agent_365.py index e72b716665c..a225be3cfda 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/test_agent_365.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/test_agent_365.py @@ -14,6 +14,20 @@ 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, + Rejected, + SignedIn, + Unavailable, + 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, @@ -25,11 +39,15 @@ from litellm.proxy.utils import ProxyLogging from litellm.types.guardrails import ( GuardrailEventHooks, LitellmParams, + Mode, 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 +63,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 +177,10 @@ class FakeHandler: def _make_guardrail( handler: FakeHandler, *, + exchanger: StubTokenExchanger | None = None, unreachable_fallback: str = "fail_closed", + default_on: bool = True, + event_hook: str | Mode = "pre_mcp_call", ) -> Agent365Guardrail: return Agent365Guardrail( guardrail_name="agent-365-guard", @@ -130,29 +189,61 @@ def _make_guardrail( client_secret="secret-123", unreachable_fallback=unreachable_fallback, async_handler=handler, - event_hook="pre_mcp_call", - default_on=True, + token_exchanger=exchanger if exchanger is not None else StubTokenExchanger(_obo_ok()), + event_hook=event_hook, + default_on=default_on, ) -def _default_fallback_guardrail(handler: FakeHandler) -> Agent365Guardrail: +def _default_fallback_guardrail(handler: FakeHandler, exchanger: StubTokenExchanger | None = None) -> Agent365Guardrail: return Agent365Guardrail( guardrail_name="agent-365-guard", tenant_id="tenant-abc", client_id="client-xyz", client_secret="secret-123", async_handler=handler, + token_exchanger=exchanger if exchanger is not None else StubTokenExchanger(_obo_ok()), event_hook="pre_mcp_call", default_on=True, ) +def _entra_driven_guardrail( + handler: FakeHandler, *, unreachable_fallback: str = "fail_closed", request_timeout: float = 10.0 +) -> Agent365Guardrail: + """A guardrail whose Entra exchange runs through the real exchanger and the guardrail's own HTTP edge, so + ``handler`` answers the token POST first and the evaluate POST after it.""" + return Agent365Guardrail( + guardrail_name="agent-365-guard", + tenant_id="tenant-abc", + client_id="client-xyz", + client_secret="secret-123", + unreachable_fallback=unreachable_fallback, + request_timeout=request_timeout, + async_handler=handler, + event_hook="pre_mcp_call", + default_on=True, + ) + + +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: data: Final[dict] = { "mcp_tool_name": "send_email", "mcp_arguments": {"to": "user@example.com", "body": "hello"}, "mcp_server_name": "outlook_mcp", - "incoming_bearer_token": FAKE_ASSERTION, + "incoming_subject_token": FAKE_ASSERTION, "metadata": {"headers": {"mcp-session-id": "sess-123"}}, } data.update(overrides) @@ -257,14 +348,18 @@ 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()]) + exchanger: Final = StubTokenExchanger(_obo_ok()) with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): - guardrail: Final = initialize_guardrail(params, {"guardrail_name": "a365-stale"}, async_handler=handler) + guardrail: Final = initialize_guardrail( + params, {"guardrail_name": "a365-stale"}, async_handler=handler, token_exchanger=exchanger + ) assert "ignoring api_base, resource_app_id, agent_id" in caplog.text 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" + _, server, config = exchanger.calls[0] + assert server.resource == AGENT_365_PROD_API_BASE + assert config.scopes == (f"{AGENT_365_PROD_RESOURCE_APP_ID}/{AGENT_365_SCOPE_NAME}",) + evaluate_call: Final = handler.calls[0] assert evaluate_call.url == EVALUATE_URL assert evaluate_call.json["agentId"] == "my-agent-key" @@ -305,7 +400,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 +414,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 +445,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"] == { + assert handler.calls[0].json["tool"] == { "name": "send_email", "description": "Send an email", "inputSchema": schema, @@ -364,10 +461,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"} + 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 +489,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 +516,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 +562,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 +571,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 +587,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 +604,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 +618,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 +629,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 +638,18 @@ 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 +657,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 +674,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 +683,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 +694,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 +703,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 +712,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 +723,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()) @@ -628,23 +735,35 @@ class TestUnreachableFallback: handler: Final = FakeHandler([]) guardrail: Final = _make_guardrail(handler) with pytest.raises(HTTPException) as exc_info: - await _run(guardrail, _mcp_data(incoming_bearer_token=None)) + await _run(guardrail, _mcp_data(incoming_subject_token=None)) assert exc_info.value.status_code == 401 assert handler.calls == [] @pytest.mark.asyncio - async def test_non_jwt_bearer_token_fail_closed(self): + async def test_raw_bearer_without_subject_token_is_no_bearer(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")) + await _run(guardrail, _mcp_data(incoming_subject_token=None, incoming_bearer_token=FAKE_ASSERTION)) assert exc_info.value.status_code == 401 + assert exchanger.calls == [] + + @pytest.mark.asyncio + async def test_non_jwt_bearer_token_fail_closed(self): + exchanger: Final = StubTokenExchanger() + handler: Final = FakeHandler([]) + guardrail: Final = _make_guardrail(handler, exchanger=exchanger) + with pytest.raises(HTTPException) as exc_info: + await _run(guardrail, _mcp_data(incoming_subject_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): handler: Final = FakeHandler([]) guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open") - data: Final = _mcp_data(incoming_bearer_token=None) + data: Final = _mcp_data(incoming_subject_token=None) with pytest.raises(HTTPException) as exc_info: await _run(guardrail, data) assert exc_info.value.status_code == 401 @@ -655,18 +774,51 @@ 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_malformed_assertion_reported_as_invalid_client_blocks_even_fail_open(self): + """Entra answers a garbled or unverifiable caller assertion with invalid_client AADSTS5002723, the + same top-level code as a wrong gateway secret. The sub-code makes it the caller's 401 challenge, + never the fail-open Unscanned pass and never a 503 that blames the gateway credentials.""" + handler: Final = FakeHandler( + [ + _response( + 400, + { + "error": "invalid_client", + "error_description": "AADSTS5002723: Invalid JWT token. Token is not well formed.", + "error_codes": [5002723], + }, + ) + ] + ) + guardrail: Final = _entra_driven_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 == 401 + assert "On-Behalf-Of token exchange was rejected" in exc_info.value.detail["message"] + info: Final = _guardrail_info(data) + assert info["guardrail_status"] == "guardrail_intervened" + assert info["guardrail_response"]["verdict"] == "Rejected" + assert ( + info["guardrail_response"]["reason"] + == "the Entra On-Behalf-Of token exchange was rejected (invalid_client)" + ) + assert [call.url for call in handler.calls] == [TOKEN_URL] @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,53 +832,56 @@ 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"] + ("entra", "error_code"), + [ + (_response(400, {"error": "invalid_scope"}), "invalid_scope"), + (_response(401, {"error": "invalid_client"}), "invalid_client"), + (_response(400, {"error": "invalid_client", "error_codes": [7000215]}), "invalid_client"), + (_response(400, {"error": "unauthorized_client"}), "unauthorized_client"), + ], + ids=["invalid_scope", "invalid_client_401", "invalid_client_wrong_secret", "unauthorized_client"], ) - 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) + async def test_gateway_credential_rejection_is_unavailable_not_a_caller_401( + self, entra: httpx.Response, error_code: str + ): + """The verdict reason carries the OAuth error code Entra answered with, the text the Logs row and the + 503 detail show an admin, not a generic exchanger summary.""" + handler: Final = FakeHandler([entra]) + guardrail: Final = _entra_driven_guardrail(handler) data: Final = _mcp_data() with pytest.raises(HTTPException) as exc_info: await _run(guardrail, data) assert exc_info.value.status_code == 503 assert exc_info.value.headers is None or "WWW-Authenticate" not in exc_info.value.headers + assert f"({error_code})" in exc_info.value.detail["message"] info: Final = _guardrail_info(data) assert info["guardrail_status"] == "guardrail_failed_to_respond" assert info["guardrail_response"]["verdict"] == "Unavailable" - assert error_code in info["guardrail_response"]["reason"] - assert "client_secret" in info["guardrail_response"]["reason"] + assert info["guardrail_response"]["reason"] == ( + f"Entra rejected the gateway's own Agent 365 credentials ({error_code}); " + "check the guardrail's client_id and client_secret" + ) + assert [call.url for call in handler.calls] == [TOKEN_URL] @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,22 +890,56 @@ 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") + handler: Final = FakeHandler([_response(401, {"error": "invalid_client"})]) + guardrail: Final = _entra_driven_guardrail(handler, 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 "invalid_client" in info["guardrail_response"]["reason"] + assert "(invalid_client)" in info["guardrail_response"]["reason"] + assert [call.url for call in handler.calls] == [TOKEN_URL] + + @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 @@ -759,50 +948,122 @@ class TestUnreachableFallback: assert info["guardrail_response"]["verdict"] == "Unscanned" +class TestEntraTokenEndpointReasons: + """Every way the Entra token endpoint can fail keeps its own verdict reason, since that text is what the + guardrail Logs row and the 503 detail carry.""" + + @pytest.mark.asyncio + @pytest.mark.parametrize( + ("entra", "reason"), + [ + ( + _response(500, text="gateway"), + "the Entra token endpoint could not be reached (HTTPStatusError)", + ), + (httpx.ConnectError("refused"), "the Entra token endpoint could not be reached (ConnectError)"), + (_response(200, text="waf page"), "the Entra token endpoint returned a non-JSON body"), + (_response(200, payload=["x"]), "the Entra token endpoint returned a non-object JSON body"), + (_response(200, payload={"token_type": "Bearer"}), "the Entra token endpoint returned no access_token"), + ( + _response(200, payload={"access_token": 7}), + "the Entra token endpoint returned a non-string access_token", + ), + ], + ids=["http_500", "connect_error", "non_json", "non_object", "no_access_token", "non_string_access_token"], + ) + async def test_unavailable_reason_names_the_fault(self, entra: object, reason: str): + handler: Final = FakeHandler([entra]) + guardrail: Final = _entra_driven_guardrail(handler) + data: Final = _mcp_data() + with pytest.raises(HTTPException) as exc_info: + await _run(guardrail, data) + assert exc_info.value.status_code == 503 + assert reason in exc_info.value.detail["message"] + info: Final = _guardrail_info(data) + assert info["guardrail_status"] == "guardrail_failed_to_respond" + assert info["guardrail_response"]["verdict"] == "Unavailable" + assert info["guardrail_response"]["reason"] == reason + assert [call.url for call in handler.calls] == [TOKEN_URL] + + @pytest.mark.asyncio + @pytest.mark.parametrize("error_code", ["invalid_grant", "interaction_required", "invalid_resource"]) + async def test_caller_rejection_reason_names_the_oauth_code(self, error_code: str): + handler: Final = FakeHandler([_response(400, {"error": error_code, "error_codes": [700082]})]) + guardrail: Final = _entra_driven_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 == 401 + assert _guardrail_info(data)["guardrail_response"]["reason"] == ( + f"the Entra On-Behalf-Of token exchange was rejected ({error_code})" + ) + + @pytest.mark.asyncio + @pytest.mark.parametrize("unreachable_fallback", ["fail_closed", "fail_open"]) + async def test_throttled_token_endpoint_blocks_regardless_of_fallback(self, unreachable_fallback: str): + handler: Final = FakeHandler([_response(429, text="slow down"), _response(429, text="slow down")]) + guardrail: Final = _entra_driven_guardrail(handler, unreachable_fallback=unreachable_fallback) + 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_response"]["verdict"] == "Throttled" + assert info["guardrail_response"]["reason"] == "the Entra token endpoint returned HTTP 429" + connect: Final = await guardrail.preflight_caller_sign_in(_server(), _user(), FAKE_ASSERTION) + assert connect == Unavailable(detail="the Entra token endpoint returned HTTP 429", fail_open=False) + + @pytest.mark.asyncio + async def test_preflight_gateway_fault_detail_names_the_oauth_code(self): + handler: Final = FakeHandler([_response(400, {"error": "invalid_scope"})]) + guardrail: Final = _entra_driven_guardrail(handler) + verdict: Final = await guardrail.preflight_caller_sign_in(_server(), _user(), FAKE_ASSERTION) + assert verdict == Unavailable( + detail=( + "Entra rejected the gateway's own Agent 365 credentials (invalid_scope); " + "check the guardrail's client_id and client_secret" + ), + fail_open=False, + ) + + @pytest.mark.asyncio + async def test_preflight_unavailable_detail_names_the_fault(self): + handler: Final = FakeHandler([_response(200, text="waf page")]) + guardrail: Final = _entra_driven_guardrail(handler, unreachable_fallback="fail_open") + verdict: Final = await guardrail.preflight_caller_sign_in(_server(), _user(), FAKE_ASSERTION) + assert verdict == Unavailable(detail="the Entra token endpoint returned a non-JSON body", fail_open=True) + + 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" + await _run(guardrail, _mcp_data(incoming_subject_token=other_assertion)) + 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 +1096,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 +1125,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 +1133,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 +1146,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 +1162,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 +1171,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 +1194,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 +1219,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 +1244,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 +1252,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) @@ -1091,6 +1301,34 @@ class _ArgumentMasker(CustomGuardrail): return data +class TestMcpBridgeHandsOverTheSubjectToken: + """The MCP manager separates the raw ``Authorization`` bearer from the caller's subject token (the + bearer minus LiteLLM's own admission credentials). The bridge that turns the manager's kwargs into + the guardrail's data dict has to carry the subject token, or every tool call looks anonymous.""" + + @pytest.mark.asyncio + async def test_manager_kwargs_reach_the_obo_exchange(self): + exchanger: Final = StubTokenExchanger([_ok_exchange()]) + handler: Final = FakeHandler([_allow_response()]) + guardrail: Final = _make_guardrail(handler, exchanger=exchanger) + proxy_logging: Final = ProxyLogging(user_api_key_cache=DualCache()) + manager_kwargs: Final = { + "name": "send_email", + "arguments": {"to": "user@example.com"}, + "server_name": "outlook_mcp", + "user_api_key_auth": _user(), + "incoming_bearer_token": "sk-1234", + "incoming_subject_token": FAKE_ASSERTION, + "headers": {"mcp-session-id": "sess-123"}, + } + data: Final = proxy_logging._convert_mcp_to_llm_format( + proxy_logging._create_mcp_request_object_from_kwargs(manager_kwargs), manager_kwargs + ) + await _run(guardrail, data) + assert [call[0] for call in exchanger.calls] == [FAKE_ASSERTION] + assert handler.calls[0].json["tool"]["name"] == "send_email" + + class TestFinalArgumentsEvaluated: """Agent 365 must judge the arguments that reach the upstream tool. A sibling guardrail that rewrites them must not be able to slip a different argument state past the verdict, whichever way the two @@ -1099,7 +1337,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 +1354,182 @@ 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_configured_server_scopes_replace_the_gateway_scope(self): + guardrail: Final = _make_guardrail(FakeHandler([])) + server: Final = _server(scopes=["https://example/mcp/scoped/access_as_user", "offline_access"]) + sign_in: Final = guardrail.caller_sign_in(server, None) + assert sign_in == CallerSignIn( + issuers=("https://login.microsoftonline.com/tenant-abc/v2.0",), + scopes=("https://example/mcp/scoped/access_as_user", "offline_access"), + ) + assert guardrail.caller_sign_in(_server(scopes=[]), None) == 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_or_team_does_not_gate(self): + guardrail: Final = _make_guardrail(FakeHandler([])) + opted_out_key: Final = UserAPIKeyAuth( + api_key="k", user_id="u-1", metadata={"opted_out_global_guardrails": [guardrail.guardrail_name]} + ) + opted_out_team: Final = UserAPIKeyAuth( + api_key="k", user_id="u-1", team_metadata={"opted_out_global_guardrails": [guardrail.guardrail_name]} + ) + assert guardrail.caller_sign_in(_server(), opted_out_key) is None + assert guardrail.caller_sign_in(_server(), opted_out_team) is None + assert guardrail.caller_sign_in(_server(), UserAPIKeyAuth(api_key="k", user_id="u-1")) is not None + assert guardrail.caller_sign_in(_server(), None) is not None + + def test_tag_mode_advertises_entra_only_when_it_gates_a_tagless_connect(self, monkeypatch): + monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True) + plain_key: Final = UserAPIKeyAuth(api_key="k", user_id="u-1") + tag_only: Final = _make_guardrail(FakeHandler([]), event_hook=Mode(tags={"a365": "pre_mcp_call"})) + assert tag_only.caller_sign_in(_server(), plain_key) is None + assert tag_only.caller_sign_in(_server(), None) is None + with_default: Final = _make_guardrail( + FakeHandler([]), event_hook=Mode(tags={"a365": "pre_mcp_call"}, default="pre_mcp_call") + ) + expected: Final = CallerSignIn( + issuers=("https://login.microsoftonline.com/tenant-abc/v2.0",), + scopes=("api://client-xyz/access_as_user",), + ) + assert with_default.caller_sign_in(_server(), plain_key) == expected + assert with_default.caller_sign_in(_server(), None) == expected + + def test_obo_server_with_provider_advertises_both_issuers_and_the_server_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",) + + +class TestPreflightCallerSignIn: + """The connect-time check must give the connect gate a verdict it can challenge on: a rejected + subject becomes the RFC 9728 challenge, an unreachable endpoint the guardrail's fallback policy.""" + + @pytest.mark.asyncio + async def test_ok_exchange_signs_in(self): + exchanger: Final = StubTokenExchanger(_obo_ok()) + guardrail: Final = _make_guardrail(FakeHandler([]), exchanger=exchanger) + + verdict: Final = await guardrail.preflight_caller_sign_in(_server(), _user(), FAKE_ASSERTION) + + assert verdict == SignedIn() + assert [call[0] for call in exchanger.calls] == [FAKE_ASSERTION] + + @pytest.mark.asyncio + async def test_unauthorized_error_rejects_with_the_idp_detail(self): + exchanger: Final = StubTokenExchanger( + [Error(CredError.of_unauthorized("the provided assertion has expired", claims="step-up"))] + ) + guardrail: Final = _make_guardrail(FakeHandler([]), exchanger=exchanger) + + verdict: Final = await guardrail.preflight_caller_sign_in(_server(), _user(), FAKE_ASSERTION) + + assert verdict == Rejected(detail="the provided assertion has expired", claims="step-up") + + @pytest.mark.asyncio + async def test_configured_timeout_bounds_the_entra_exchange_leg(self): + handler: Final = FakeHandler([_response(200, {"access_token": "exchanged", "expires_in": 3600})]) + guardrail: Final = _entra_driven_guardrail(handler, request_timeout=0.5) + + verdict: Final = await guardrail.preflight_caller_sign_in(_server(), _user(), FAKE_ASSERTION) + + assert verdict == SignedIn() + assert [(call.url, call.timeout) for call in handler.calls] == [(TOKEN_URL, 0.5)], ( + "the Entra token POST must carry the guardrail's own request_timeout" + ) + + @pytest.mark.asyncio + async def test_malformed_assertion_is_rejected_at_connect_even_fail_open(self): + handler: Final = FakeHandler([_response(401, {"error": "invalid_client", "error_codes": [5002723]})]) + guardrail: Final = _entra_driven_guardrail(handler, unreachable_fallback="fail_open") + + verdict: Final = await guardrail.preflight_caller_sign_in(_server(), _user(), FAKE_ASSERTION) + + assert verdict == Rejected(detail="invalid_client", claims=None) + + @pytest.mark.asyncio + async def test_misconfigured_fail_closed_is_unavailable(self): + exchanger: Final = StubTokenExchanger([Error(CredError.of_misconfigured("bad client_secret"))]) + guardrail: Final = _make_guardrail(FakeHandler([]), exchanger=exchanger, unreachable_fallback="fail_closed") + + verdict: Final = await guardrail.preflight_caller_sign_in(_server(), _user(), FAKE_ASSERTION) + + assert isinstance(verdict, Unavailable) + assert verdict.fail_open is False + assert "bad client_secret" in verdict.detail + assert "resource_app_id" not in verdict.detail, ( + "the field was removed from the config; do not tell admins to check it" + ) + + @pytest.mark.asyncio + async def test_token_endpoint_failure_is_unavailable_under_the_fallback_policy(self): + exchanger: Final = StubTokenExchanger([Error(CredError.of_upstream_unavailable("token endpoint 503"))]) + guardrail: Final = _make_guardrail(FakeHandler([]), exchanger=exchanger, unreachable_fallback="fail_open") + + verdict: Final = await guardrail.preflight_caller_sign_in(_server(), _user(), FAKE_ASSERTION) + + assert verdict == Unavailable( + detail="the Entra token exchange failed (upstream unavailable: token endpoint 503)", fail_open=True + ) + + @pytest.mark.asyncio + async def test_endpoint_unreachable_fail_open_is_unavailable(self): + exchanger: Final = StubTokenExchanger( + [httpx.ConnectError("refused", request=httpx.Request("POST", "https://example.test"))] + ) + guardrail: Final = _make_guardrail(FakeHandler([]), exchanger=exchanger, unreachable_fallback="fail_open") + + verdict: Final = await guardrail.preflight_caller_sign_in(_server(), _user(), FAKE_ASSERTION) + + assert isinstance(verdict, Unavailable) + assert verdict.fail_open is True + + @pytest.mark.asyncio + async def test_non_assertion_subject_is_rejected_without_exchanging(self): + exchanger: Final = StubTokenExchanger() + guardrail: Final = _make_guardrail(FakeHandler([]), exchanger=exchanger) + + verdict: Final = await guardrail.preflight_caller_sign_in(_server(), _user(), "opaque-bearer") + + assert isinstance(verdict, Rejected) + assert verdict.claims is None + assert exchanger.calls == [] diff --git a/tests/unit/proxy/utils/proxy_logging/test_mcp_bridging.py b/tests/unit/proxy/utils/proxy_logging/test_mcp_bridging.py index 25be3b5de6b..c7fb0848f76 100644 --- a/tests/unit/proxy/utils/proxy_logging/test_mcp_bridging.py +++ b/tests/unit/proxy/utils/proxy_logging/test_mcp_bridging.py @@ -39,6 +39,7 @@ def test_convert_mcp_to_llm_format_returns_synthetic_data(proxy_logging, make_mc "user_api_key_hash": "hash", "user_api_key_request_route": "/mcp", "incoming_bearer_token": "tok", + "incoming_subject_token": "a.b.c", }, ) snapshot = { @@ -47,6 +48,7 @@ def test_convert_mcp_to_llm_format_returns_synthetic_data(proxy_logging, make_mc "mcp_tool_name": out["mcp_tool_name"], "mcp_arguments": out["mcp_arguments"], "incoming_bearer_token": out["incoming_bearer_token"], + "incoming_subject_token": out["incoming_subject_token"], "message_role": out["messages"][0]["role"], } assert snapshot == { @@ -55,6 +57,7 @@ def test_convert_mcp_to_llm_format_returns_synthetic_data(proxy_logging, make_mc "mcp_tool_name": "search", "mcp_arguments": {"q": "hello"}, "incoming_bearer_token": "tok", + "incoming_subject_token": "a.b.c", "message_role": "user", } @@ -66,12 +69,14 @@ def test_convert_mcp_to_llm_format_defaults_model(proxy_logging, make_mcp_reques "model": out["model"], "mcp_tool_name": out["mcp_tool_name"], "incoming_bearer_token": out["incoming_bearer_token"], + "incoming_subject_token": out["incoming_subject_token"], "user_id": out["user_api_key_user_id"], } assert snapshot == { "model": "mcp-tool-call", "mcp_tool_name": "calculator", "incoming_bearer_token": None, + "incoming_subject_token": None, "user_id": None, }