From 4cadb402e5ab71650d9ca64c13984b2aeded49b8 Mon Sep 17 00:00:00 2001 From: yucheng Date: Sat, 26 Sep 2026 09:32:46 +0000 Subject: [PATCH 01/51] feat(mcp): fold gateway sign-in into a caller sign-in contract on the challenge path Replace the parallel gateway sign-in provider registry with CallerSignInProvider, merged into the existing oauth2_token_exchange challenge: gates key off caller_sign_in_for() returning non-None, the resolver answers by server_id, case-insensitive name, and short prefix like the router, keeps_caller_authorization covers OBO servers, and the Agent 365 guardrail exchanges through an injected TokenExchanger. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/caller_sign_in.py | 124 ++++ .../mcp_server/discoverable_endpoints.py | 58 +- .../mcp_server/mcp_server_manager.py | 28 +- .../_experimental/mcp_server/operations.py | 10 +- .../token_exchange_provider.py | 5 + .../proxy/_experimental/mcp_server/server.py | 40 +- .../proxy/_experimental/mcp_server/utils.py | 8 + .../guardrail_hooks/agent_365/agent_365.py | 229 +++---- .../types/mcp_server/mcp_server_manager.py | 23 +- .../test_token_exchange_provider.py | 42 +- .../mcp_server/test_caller_sign_in.py | 134 +++++ .../mcp_server/test_discoverable_endpoints.py | 26 +- .../mcp_server/test_mcp_server.py | 181 +++++- .../mcp_server/test_mcp_server_manager.py | 13 + .../guardrail_hooks/test_agent_365.py | 561 ++++++++++-------- 15 files changed, 1005 insertions(+), 477 deletions(-) create mode 100644 litellm/proxy/_experimental/mcp_server/caller_sign_in.py create mode 100644 tests/test_litellm/proxy/_experimental/mcp_server/test_caller_sign_in.py diff --git a/litellm/proxy/_experimental/mcp_server/caller_sign_in.py b/litellm/proxy/_experimental/mcp_server/caller_sign_in.py new file mode 100644 index 00000000000..9d9b8b7297e --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/caller_sign_in.py @@ -0,0 +1,124 @@ +"""Caller-side sign-in requirements for MCP connects. + +A guardrail that evaluates tool calls in the caller's own identity (an On-Behalf-Of exchange of the caller's +bearer) needs the caller signed in with its issuer before the first tool call, and a tool call's JSON-RPC +error cannot carry ``WWW-Authenticate``. ``token_exchange`` (OBO) servers have the same need: the caller must +present a subject token the gateway can exchange. Both cases share one contract: a connect that carries no +usable subject answers 401 with the RFC 9728 challenge, and the protected-resource metadata advertises the +issuers and scopes the caller signs in for. + +Guardrails implement :class:`CallerSignInProvider`; :func:`caller_sign_in_for` merges the OBO server's own +requirement with every registered provider's so the challenge and the metadata always agree. The MCP package +never imports a concrete provider. +""" + +from __future__ import annotations + +from collections.abc import Mapping +from dataclasses import dataclass +from typing import TYPE_CHECKING, Final, Protocol, cast, runtime_checkable + +from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError + +import litellm +from litellm.integrations.custom_guardrail import CustomGuardrail +from litellm.types.mcp import MCPAuth + +if TYPE_CHECKING: + from litellm.proxy._types import UserAPIKeyAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + +@dataclass(frozen=True, slots=True) +class CallerSignIn: + """The issuers a caller signs in with and the scopes it requests before calling a server.""" + + issuers: tuple[str, ...] + scopes: tuple[str, ...] + + +@runtime_checkable +class CallerSignInProvider(Protocol): + def caller_sign_in(self, server: MCPServer, user_api_key_auth: UserAPIKeyAuth | None) -> CallerSignIn | None: + """The sign-in this provider requires of callers hitting ``server``; ``None`` when it does not gate + the server for this caller (``user_api_key_auth=None`` is the anonymous metadata fetch that follows + a challenge).""" + ... + + +class _JwtIssuerEntry(BaseModel): + model_config = ConfigDict(extra="ignore") + + issuer: str | None = None + + +class _JwtAuthConfig(BaseModel): + model_config = ConfigDict(extra="ignore") + + issuers: list[_JwtIssuerEntry] = [] # mutable-ok: pydantic copies the default per instance + + +_JWT_AUTH_ADAPTER: Final = TypeAdapter(_JwtAuthConfig) + + +def _jwt_auth_issuer_entries(jwtauth: object) -> tuple[_JwtIssuerEntry, ...]: + try: + return tuple(_JWT_AUTH_ADAPTER.validate_python(jwtauth, from_attributes=True).issuers) + except ValidationError: + return () + + +def _providers() -> tuple[CallerSignInProvider, ...]: + return tuple( + callback + for callback in litellm.logging_callback_manager.get_custom_loggers_for_type(CustomGuardrail) + if isinstance(callback, CallerSignInProvider) + ) + + +def jwt_auth_issuers() -> tuple[str, ...]: + """The OAuth issuer identifier(s) LiteLLM's JWT auth trusts, for the OBO PRM authorization_servers. + + In token_exchange the IdP that issues the subject JWT is the same one LiteLLM validates it + against, so OBO discovery points clients at the JWT-auth issuer to obtain a subject token. + Sourced from ``JWT_ISSUER`` and any configured ``litellm_jwtauth.issuers``. + """ + import os # noqa: PLC0415 + + from litellm.proxy.proxy_server import ( # noqa: PLC0415 # lazy: proxy_server pulls the whole proxy graph + general_settings, # pyright: ignore[reportUnknownVariableType] # proxy_server.general_settings is a raw untyped dict + ) + + env_issuer: Final = os.getenv("JWT_ISSUER") + env: Final[tuple[str, ...]] = (env_issuer,) if env_issuer else () + + settings: Final[Mapping[str, object]] = cast(Mapping[str, object], general_settings) + configured: Final = tuple( + entry.issuer for entry in _jwt_auth_issuer_entries(settings.get("litellm_jwtauth")) if entry.issuer + ) + return tuple(dict.fromkeys((*env, *configured))) + + +def caller_sign_in_for(server: MCPServer, user_api_key_auth: UserAPIKeyAuth | None) -> CallerSignIn | None: + """The merged sign-in requirement for ``server``: the OBO server's own issuer/scopes plus every + registered provider's contribution. ``None`` when nothing requires sign-in, which is also the gate the + connect-time challenge branches on.""" + contributions: Final = [ + contribution + for contribution in ( + *( + ( + CallerSignIn(issuers=jwt_auth_issuers(), scopes=tuple(server.scopes or ())), + ) + if server.auth_type == MCPAuth.oauth2_token_exchange + else () + ), + *(provider.caller_sign_in(server, user_api_key_auth) for provider in _providers()), + ) + if contribution is not None + ] + if not contributions: + return None + issuers: Final = tuple(dict.fromkeys(issuer for contribution in contributions for issuer in contribution.issuers)) + scopes: Final = tuple(dict.fromkeys(scope for contribution in contributions for scope in contribution.scopes)) + return CallerSignIn(issuers=issuers, scopes=scopes) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 7a0f59c3c2b..6940e8d4f02 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -37,6 +37,7 @@ from litellm.proxy._experimental.mcp_server.bridge_token_flow import ( can_store_oauth_credential, oauth_authorization_uses_gateway_credential, ) +from litellm.proxy._experimental.mcp_server.caller_sign_in import caller_sign_in_for from litellm.proxy._experimental.mcp_server.faults import ( CallerRejected, CredentialSource, @@ -2580,9 +2581,9 @@ async def _build_oauth_protected_resource_response( detail=(f"Upstream oauth-protected-resource metadata unavailable for MCP server {mcp_server.name!r}"), ) - obo_response: Final = _obo_protected_resource_response(mcp_server, resource_url) - if obo_response is not None: - return obo_response + sign_in_response: Final = _caller_sign_in_protected_resource_response(mcp_server, resource_url) + if sign_in_response is not None: + return sign_in_response if mcp_server is not None and mcp_server.advertises_gateway_authorization_server: return { @@ -2603,51 +2604,30 @@ async def _build_oauth_protected_resource_response( } -def _obo_protected_resource_response(mcp_server: MCPServer | None, resource_url: str) -> dict | None: - """The OBO (token_exchange) PRM, or None when this server is not OBO / no issuer is configured. +def _caller_sign_in_protected_resource_response( + mcp_server: MCPServer | None, resource_url: str +) -> dict[str, object] | None: + """The caller sign-in PRM: the OBO issuer(s) LiteLLM trusts merged with every registered + ``CallerSignInProvider``'s contribution, or None when no sign-in gates this server. - The client SSOs with the IdP to obtain a subject token, which LiteLLM then exchanges, so discovery - points at the JWT-auth issuer(s) LiteLLM trusts (the same IdP that issues and validates the - subject), not the gateway. None falls the caller back to the gateway default so discovery still - returns metadata; it just can't name the IdP. + The client SSOs with the IdP to obtain a subject token, which LiteLLM then exchanges (or a + guardrail consumes directly), so discovery points at the issuer(s), not the gateway. None falls + the caller back to the gateway default so discovery still returns metadata; it just can't name + the IdP. The anonymous metadata fetch passes ``user_api_key_auth=None`` because it cannot see + which key selected a provider. """ - if mcp_server is None or mcp_server.auth_type != MCPAuth.oauth2_token_exchange: + if mcp_server is None: return None - issuers: Final = _jwt_auth_issuers() - if not issuers: + sign_in: Final = caller_sign_in_for(mcp_server, None) + if sign_in is None or not sign_in.issuers: return None return { - "authorization_servers": issuers, + "authorization_servers": list(sign_in.issuers), "resource": resource_url, - "scopes_supported": (mcp_server.scopes if mcp_server.scopes else []), + "scopes_supported": list(sign_in.scopes), } -def _jwt_auth_issuers() -> list: - """The OAuth issuer identifier(s) LiteLLM's JWT auth trusts, for the OBO PRM authorization_servers. - - In token_exchange the IdP that issues the subject JWT is the same one LiteLLM validates it - against, so OBO discovery points clients at the JWT-auth issuer to obtain a subject token. - Sourced from ``JWT_ISSUER`` and any configured ``litellm_jwtauth.issuers``. - """ - import os # noqa: PLC0415 - - from litellm.proxy.proxy_server import general_settings # noqa: PLC0415 - - issuers: Final[list] = [] - env_issuer: Final = os.getenv("JWT_ISSUER") - if env_issuer: - issuers.append(env_issuer) - - jwtauth: Final = general_settings.get("litellm_jwtauth") if isinstance(general_settings, Mapping) else None - raw_issuers: Final = jwtauth.get("issuers") if isinstance(jwtauth, dict) else getattr(jwtauth, "issuers", None) - for cfg in raw_issuers or []: - issuer = cfg.get("issuer") if isinstance(cfg, dict) else getattr(cfg, "issuer", None) - if issuer and issuer not in issuers: - issuers.append(issuer) - return issuers - - @router.get("/.well-known/oauth-protected-resource") def oauth_protected_resource_root(request: Request) -> dict[str, str | tuple[str, ...]]: request_base_url: Final = get_request_base_url(request) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 09714b65f86..c38d156dd46 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -168,6 +168,7 @@ from litellm.proxy._experimental.mcp_server.utils import ( normalize_server_name, openapi_tool_name, parse_admin_env_vars, + server_answers_to_name, strip_known_server_prefix, validate_mcp_server_name, ) @@ -4241,6 +4242,12 @@ class MCPServerManager: if subject_token is not None: return case _: + from litellm.proxy._experimental.mcp_server.caller_sign_in import ( # noqa: PLC0415 # lazy: provider discovery pulls the guardrail registry + caller_sign_in_for, + ) + + if subject_token is None and caller_sign_in_for(server, user_api_key_auth) is not None: + raise_token_exchange_challenge(server, root_path=get_request_root_path()) return resolved_server: Final = await self.ensure_oauth_metadata_discovered(server) spec: Final = _to_server_spec_fail_closed(resolved_server) @@ -5964,13 +5971,8 @@ class MCPServerManager: if proxy_logging_obj is None: return hook_result - # Extract incoming Bearer token from raw request headers so - # guardrails like MCPJWTSigner can verify + re-sign it (FR-5). - normalized_raw: Final = {k.lower(): v for k, v in (raw_headers or {}).items()} - incoming_bearer_token: str | None = None - auth_hdr: Final = normalized_raw.get("authorization", "") - if auth_hdr.lower().startswith("bearer "): - incoming_bearer_token = auth_hdr[len("bearer ") :] + # Admission credentials are never handed to guardrails as the caller's assertion. + incoming_bearer_token: Final = self._extract_subject_token(None, raw_headers, user_api_key_auth) pre_hook_kwargs: Final = { "guardrail_context": guardrail_context, @@ -7189,6 +7191,18 @@ class MCPServerManager: return server return None + def get_mcp_server_answering_to(self, name: str, client_ip: str | None = None) -> MCPServer | None: + """The server a scoped ``/mcp/{name}`` connect resolves to, matched the way the router matches + it: case-insensitive over server_id, name and every published prefix form.""" + return next( + ( + server + for server in self.get_filtered_registry(client_ip).values() + if server_answers_to_name(server, name) + ), + None, + ) + def get_filtered_registry(self, client_ip: str | None = None) -> dict[str, MCPServer]: """ Get registry filtered by client IP access control. diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index 891c37f4fe7..abc7c52f5a7 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -120,6 +120,7 @@ from litellm.proxy._experimental.mcp_server.utils import ( logging_safe_mcp_headers, match_known_tool_name, normalize_server_name, + server_answers_to_name, split_server_prefix_from_name, strip_known_server_prefix, ) @@ -496,8 +497,7 @@ def _http_detail_message(detail: object) -> str: def _server_answers_to(server: MCPServer, name: str) -> bool: - requested: Final = name.lower() - return any(requested == known.lower() for known in iter_known_server_prefixes(server) if known) + return server_answers_to_name(server, name) async def raise_denied_scoped_mcp_access( @@ -1768,7 +1768,11 @@ def _challenge_missing_token_exchange_subject( warm path already raises. Gated to servers the key may reach so an unauthorized caller learns nothing about the catalog. """ - if server is None or server.auth_type != MCPAuth.oauth2_token_exchange: + from litellm.proxy._experimental.mcp_server.caller_sign_in import ( + caller_sign_in_for, # noqa: PLC0415 # lazy: caller_sign_in pulls the proxy graph + ) + + if server is None or caller_sign_in_for(server, user_api_key_auth) is None: return if requested_server is not None and requested_server.server_id != server.server_id: return diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/token_exchange_provider.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/token_exchange_provider.py index 2501f751d8e..7f5c0c99145 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/token_exchange_provider.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/token_exchange_provider.py @@ -79,6 +79,11 @@ async def _post_exchange_endpoint( parsed: Final[object] = response.json() # pyright: ignore except httpx.HTTPStatusError as status_err: status_code: Final = status_err.response.status_code + if status_code in (408, 429): + # Retry hints, not subject rejections: the IdP is shedding load, so a 401 would tell the + # caller to sign in again for nothing; surface it like a transport failure. + verbose_logger.warning("MCP token exchange throttled or timed out (HTTP %d)", status_code) + return None if 400 <= status_code < 500: oauth_error, claims = _oauth_error_fields(status_err.response) if oauth_error in _GATEWAY_FAULT_OAUTH_ERRORS: diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 2412e83b9d9..b73d183db86 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -36,6 +36,7 @@ from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( MCPRequestHandler, _is_mcp_admitted_user_subject, ) +from litellm.proxy._experimental.mcp_server.caller_sign_in import caller_sign_in_for from litellm.proxy._experimental.mcp_server.client_allowlist import ( MCPClientAllowlist, check_mcp_client_allowed, @@ -1580,6 +1581,21 @@ if MCP_AVAILABLE: ) return user_api_key_auth.model_copy(update={"object_permission": updated_op, "mcp_toolset_id": toolset_id}) + async def _key_granted_single_server( + server: MCPServer, + mcp_servers: Sequence[str] | None, + user_api_key_auth: UserAPIKeyAuth | None, + client_ip: str | None, + ) -> bool: + """Sign-in challenges are issued only on a single-server connect the key's grant admits, so a key + without access gets the grant's 403 instead of a sign-in it could not use.""" + if len(mcp_servers or []) != 1: + return False + allowed: Final = await operations._get_allowed_mcp_servers( + user_api_key_auth=user_api_key_auth, mcp_servers=mcp_servers, client_ip=client_ip + ) + return any(granted.server_id == server.server_id for granted in allowed) + async def _raise_preemptive_401_for_unauthenticated_servers( scope: Scope, mcp_servers: list[str] | None, @@ -1602,7 +1618,7 @@ if MCP_AVAILABLE: a server it will be 403'd on immediately after authentication. """ for server_name in mcp_servers or []: - server = operations.global_mcp_server_manager.get_mcp_server_by_name(server_name, client_ip=client_ip) + server = operations.global_mcp_server_manager.get_mcp_server_answering_to(server_name, client_ip=client_ip) if server is not None and allowed_server_ids is not None and server.server_id not in allowed_server_ids: # Caller's narrowed scope excludes this server — skip the # preemptive challenge and let downstream authorization @@ -1698,12 +1714,22 @@ if MCP_AVAILABLE: # reaches the token_exchange / pass-through blocks below. continue - # token_exchange (OBO): the caller supplied no subject token. Challenge at connect - # (transport level, where WWW-Authenticate survives) with the RFC 9728 resource_metadata - # so the client discovers the IdP, SSOs, and retries with a subject token, which LiteLLM - # then exchanges. A tool-call-time 401 would be wrapped into a JSON-RPC error and the - # header lost, so the discovery flow needs this pre-emptive challenge. - if server and server.auth_type == MCPAuth.oauth2_token_exchange and not oauth2_headers: + # Caller sign-in: challenge at connect because a tool-call-time 401 is wrapped into a + # JSON-RPC error and the WWW-Authenticate header is lost. Non-OBO gates fire only on a + # single-server connect the key's grant admits. + sign_in = caller_sign_in_for(server, user_api_key_auth) if server is not None else None + if ( + server + and sign_in is not None + and operations.global_mcp_server_manager._extract_subject_token( # pyright: ignore[reportPrivateUsage] # the manager owns the subject/admission filter shared with the preflight + oauth2_headers, raw_headers, user_api_key_auth + ) + is None + and ( + server.auth_type == MCPAuth.oauth2_token_exchange + or await _key_granted_single_server(server, mcp_servers, user_api_key_auth, client_ip) + ) + ): from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import ( # noqa: PLC0415 # lazy: adapter pulls MCP subgraph raise_token_exchange_challenge, ) diff --git a/litellm/proxy/_experimental/mcp_server/utils.py b/litellm/proxy/_experimental/mcp_server/utils.py index 7411dc5c4f0..8234a62c4ef 100644 --- a/litellm/proxy/_experimental/mcp_server/utils.py +++ b/litellm/proxy/_experimental/mcp_server/utils.py @@ -370,6 +370,14 @@ def iter_known_server_prefixes(server: _McpServerLike) -> Iterator[str]: yield from _emit(server_id) +def server_answers_to_name(server: _McpServerLike, name: str) -> bool: + """Whether a scoped ``/mcp/{name}`` connect resolves to ``server``: case-insensitive over every prefix + form routing accepts (alias, server_name, server_id, short prefix), the same match + ``_server_answers_to`` applies when the router scopes a request.""" + requested: Final = name.lower() + return any(requested == known.lower() for known in iter_known_server_prefixes(server) if known) + + def iter_known_tool_name_spellings(tool_name: str, server: MCPServer) -> Iterator[str]: """Yield every name that denotes the bare ``tool_name`` on ``server``: the bare name, then its wire spelling under each prefix ``iter_known_server_prefixes`` accepts. diff --git a/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py b/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py index 94de0b2df7b..e5397a7269d 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py +++ b/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py @@ -9,17 +9,14 @@ exchanged for a delegated Agent 365 token, so Defender evaluates and audits as the signed-in user. """ -import hashlib -import threading import time import uuid -from collections import OrderedDict from collections.abc import Mapping from typing import TYPE_CHECKING, ClassVar, Final, Literal, NoReturn import httpx from fastapi import HTTPException -from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError +from pydantic import BaseModel, ConfigDict, Field, SecretStr, TypeAdapter, ValidationError from typing_extensions import ReadOnly, TypedDict from litellm._logging import verbose_proxy_logger @@ -34,7 +31,18 @@ from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, ) +from litellm.proxy._experimental.mcp_server.caller_sign_in import CallerSignIn +from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Error, Ok +from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchange_provider import ( + build_token_exchanger, +) +from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchanger import TokenExchanger +from litellm.proxy._experimental.mcp_server.outbound_credentials.types import ( + ServerSpec, + TokenExchangeConfig, +) from litellm.types.guardrails import GuardrailEventHooks +from litellm.types.mcp_server.mcp_server_manager import MCPServer from litellm.types.proxy.guardrails.guardrail_hooks.agent_365 import ( AGENT_365_PROD_API_BASE, AGENT_365_PROD_RESOURCE_APP_ID, @@ -49,38 +57,14 @@ if TYPE_CHECKING: from litellm.types.utils import GuardrailStatus TOKEN_ENDPOINT_TEMPLATE: Final = "https://login.microsoftonline.com/{tenant_id}/oauth2/v2.0/token" +ENTRA_ISSUER_TEMPLATE: Final = "https://login.microsoftonline.com/{tenant_id}/v2.0" EVALUATE_URL: Final = f"{AGENT_365_PROD_API_BASE}/agents/tool-evaluation/evaluate" OBO_SCOPE: Final = f"{AGENT_365_PROD_RESOURCE_APP_ID}/{AGENT_365_SCOPE_NAME}" MCP_SESSION_ID_HEADER: Final = "mcp-session-id" DEFENDER_STATUS_EVALUATED: Final = "Evaluated" -_GATEWAY_OWNED_TOKEN_ERRORS: Final = frozenset( - {"invalid_client", "unauthorized_client", "invalid_scope", "invalid_resource"} -) -# Entra reports a malformed or unverifiable assertion as ``invalid_client`` too; only its AADSTS50027xx -# (InvalidJwtToken) sub-codes tell that apart from a bad gateway secret. -_INVALID_ASSERTION_AADSTS_PREFIX: Final = "50027" -_AADSTS_CODES_ADAPTER: Final = TypeAdapter(tuple[int, ...]) +GATEWAY_SCOPE_TEMPLATE: Final = "api://{client_id}/access_as_user" _MCP_CALL_TYPES: Final[tuple[str, ...]] = ("mcp_call", "call_mcp_tool") _TOOL_INPUT_SCHEMA_ADAPTER: Final = TypeAdapter(dict[str, object]) -_OBO_CACHE_MAX_ENTRIES: Final = 1000 -_DEFAULT_TOKEN_TTL_SECONDS: Final = 3599.0 -_TOKEN_EXPIRY_SLACK_SECONDS: Final = 60.0 - - -def _parse_expires_in(raw: object) -> float: - if not isinstance(raw, (int, float, str)): - return _DEFAULT_TOKEN_TTL_SECONDS - try: - return float(raw) - except ValueError: - return _DEFAULT_TOKEN_TTL_SECONDS - - -def _parse_aadsts_codes(raw: object) -> tuple[int, ...]: - try: - return _AADSTS_CODES_ADAPTER.validate_python(raw) - except ValidationError: - return () def _parse_tool_input_schema(raw: object) -> Mapping[str, object] | None: @@ -129,27 +113,6 @@ class _BlockedDetail(TypedDict): correlation_id: ReadOnly[str | None] -class Agent365TokenExchangeError(Exception): - def __init__(self, status_code: int, error_code: str, description: str, aadsts_codes: tuple[int, ...] = ()) -> None: - super().__init__(f"{error_code}: {description}") - self.status_code = status_code - self.error_code = error_code - self.description = description - self.aadsts_codes = aadsts_codes - - @property - def gateway_owned(self) -> bool: - """Whether the gateway's own client credentials, scope or resource were refused, as opposed to the - caller's assertion. The caller cannot fix a gateway-owned rejection by signing in again.""" - if self.error_code not in _GATEWAY_OWNED_TOKEN_ERRORS: - return False - return not any(str(code).startswith(_INVALID_ASSERTION_AADSTS_PREFIX) for code in self.aadsts_codes) - - -class Agent365MalformedResponseError(Exception): - pass - - class Agent365ThrottledError(Exception): def __init__(self, status_code: int) -> None: super().__init__(f"HTTP {status_code}") @@ -173,6 +136,7 @@ class Agent365Guardrail(CustomGuardrail): request_timeout: float = 10.0, unreachable_fallback: Literal["fail_closed", "fail_open"] = "fail_closed", async_handler: AsyncHTTPHandler | None = None, + token_exchanger: TokenExchanger | None = None, **kwargs, # noqa: ANN003 # kwargs-ok: forwarded verbatim to CustomGuardrail (event_hook, default_on) ) -> None: super().__init__( @@ -192,8 +156,19 @@ class Agent365Guardrail(CustomGuardrail): self.async_handler = async_handler or get_async_httpx_client( llm_provider=httpxSpecialProvider.GuardrailCallback ) - self._obo_token_cache: OrderedDict[str, tuple[str, float]] = OrderedDict() # mutable-ok: lock-guarded LRU - self._obo_cache_lock = threading.Lock() + self._exchange_config: Final = TokenExchangeConfig( + profile="entra_obo", + token_exchange_endpoint=TOKEN_ENDPOINT_TEMPLATE.format(tenant_id=tenant_id), + client_id=client_id, + client_secret=SecretStr(client_secret), + scopes=(OBO_SCOPE,), + ) + self._exchange_server: Final = ServerSpec( + server_id=f"agent-365:{tenant_id}", + resource=AGENT_365_PROD_API_BASE, + config=self._exchange_config, + ) + self._token_exchanger: Final = token_exchanger if token_exchanger is not None else build_token_exchanger() verbose_proxy_logger.info("Initialized Microsoft Agent 365 guardrail: %s", guardrail_name) @staticmethod @@ -233,42 +208,40 @@ class Agent365Guardrail(CustomGuardrail): ) try: - obo_token: Final = await self._get_obo_token(assertion) - except Agent365TokenExchangeError as exc: - if exc.gateway_owned: - return self._handle_unavailable( - data=data, - tool_name=tool_name, - reason=( - f"Entra rejected the gateway's own Agent 365 credentials ({exc.error_code}); " - "check the guardrail's client_id and client_secret" - ), - ) - self._handle_caller_fault( - data=data, - tool_name=tool_name, - status_code=401, - reason=f"the Entra On-Behalf-Of token exchange was rejected ({exc.error_code})", - ) - except Agent365ThrottledError as exc: - self._handle_throttled( - data=data, - tool_name=tool_name, - reason=f"the Entra token endpoint returned HTTP {exc.status_code}", - latency_ms=None, - ) + exchange_result: Final = await self._exchange_caller_assertion(assertion) except (httpx.HTTPError, LitellmTimeout, TimeoutError) as exc: return self._handle_unavailable( data=data, tool_name=tool_name, reason=f"the Entra token endpoint could not be reached ({type(exc).__name__})", ) - except Agent365MalformedResponseError as exc: - return self._handle_unavailable( - data=data, - tool_name=tool_name, - reason=str(exc), - ) + match exchange_result: + case Ok(token): + obo_token: Final = token.access_token + case Error(error): + match error.tag: + case "unauthorized": + self._handle_caller_fault( + data=data, + tool_name=tool_name, + status_code=401, + reason=f"the Entra On-Behalf-Of token exchange was rejected ({error.unauthorized.detail})", + ) + case "misconfigured": + return self._handle_unavailable( + data=data, + tool_name=tool_name, + reason=( + f"Entra rejected the gateway's own Agent 365 credentials ({error.misconfigured}); " + "check the guardrail's client_id and client_secret" + ), + ) + case _: + return self._handle_unavailable( + data=data, + tool_name=tool_name, + reason=f"the Entra token exchange failed ({error.summary})", + ) start: Final = time.perf_counter() try: @@ -284,14 +257,14 @@ class Agent365Guardrail(CustomGuardrail): reason=f"the Agent 365 endpoint could not be reached ({type(exc).__name__})", ) latency_ms: Final = (time.perf_counter() - start) * 1000.0 - fallback: Final = self._handle_evaluate_error( + fallback: Final = await self._handle_evaluate_error( data=data, tool_name=tool_name, assertion=assertion, response=response, latency_ms=latency_ms ) if fallback is not None: return fallback return self._enforce_verdict(data=data, tool_name=tool_name, response=response, latency_ms=latency_ms) - def _handle_evaluate_error( + async def _handle_evaluate_error( self, data: dict, # mutable-ok: guardrail logging appends into the request metadata in place tool_name: str, @@ -308,7 +281,9 @@ class Agent365Guardrail(CustomGuardrail): ) if 400 <= response.status_code < 500: if response.status_code == 401: - self._evict_obo_token(assertion) + await self._token_exchanger.invalidate( + assertion, self._exchange_server, self._exchange_config, tenant_id=self.tenant_id + ) self._record_verdict( data=data, verdict="Rejected", @@ -455,62 +430,29 @@ class Agent365Guardrail(CustomGuardrail): return call_id return str(uuid.uuid4()) - async def _get_obo_token(self, assertion: str) -> str: - cache_key: Final = hashlib.sha256(assertion.encode("utf-8")).hexdigest() - now: Final = time.time() - with self._obo_cache_lock: - cached: Final = self._obo_token_cache.get(cache_key) - if cached and cached[1] > now + _TOKEN_EXPIRY_SLACK_SECONDS: - self._obo_token_cache.move_to_end(cache_key) - return cached[0] - - response: Final = await self._post_allowing_error_status( - url=TOKEN_ENDPOINT_TEMPLATE.format(tenant_id=self.tenant_id), - data={ # mutable-ok: OAuth form body; AsyncHTTPHandler.post requires dict - "grant_type": "urn:ietf:params:oauth:grant-type:jwt-bearer", - "client_id": self.client_id, - "client_secret": self.client_secret, - "assertion": assertion, - "scope": OBO_SCOPE, - "requested_token_use": "on_behalf_of", - }, - headers={"Content-Type": "application/x-www-form-urlencoded"}, # mutable-ok: httpx header dict + def caller_sign_in(self, server: MCPServer, user_api_key_auth: "UserAPIKeyAuth | None") -> CallerSignIn | None: + """The Entra sign-in this guardrail requires of callers: only a ``default_on`` guardrail the caller's + key or team has not opted out of, because the anonymous metadata fetch that follows a challenge cannot + see which key selected a guardrail and would advertise the wrong issuer. Only servers that leave the + caller's top-level ``Authorization`` with the gateway qualify: a forwarded API-key header travels + upstream in its own slot and does not displace the Entra assertion.""" + if not (self.default_on and server.keeps_caller_authorization): + return None + if user_api_key_auth is not None: + probe: Final[dict[str, Mapping[str, object]]] = { # pyright: ignore[reportUnknownVariableType] # UserAPIKeyAuth metadata dicts are untyped + "metadata": { + "user_api_key_metadata": user_api_key_auth.metadata, # pyright: ignore[reportUnknownMemberType] # UserAPIKeyAuth.metadata is a raw dict + "user_api_key_team_metadata": user_api_key_auth.team_metadata, # pyright: ignore[reportUnknownMemberType] # UserAPIKeyAuth.team_metadata is a raw dict + } + } + if self.should_run_guardrail( # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # should_run_guardrail takes an untyped data dict + data=probe, event_type=GuardrailEventHooks.pre_mcp_call + ) is not True: + return None + return CallerSignIn( + issuers=(ENTRA_ISSUER_TEMPLATE.format(tenant_id=self.tenant_id),), + scopes=(GATEWAY_SCOPE_TEMPLATE.format(client_id=self.client_id),), ) - if response.status_code in (408, 429): - raise Agent365ThrottledError(status_code=response.status_code) - if response.status_code >= 500: - raise httpx.HTTPStatusError( - f"Entra token endpoint returned {response.status_code}", - request=response.request, - response=response, - ) - try: - parsed_body: Final = response.json() - except ValueError as exc: - raise Agent365MalformedResponseError("the Entra token endpoint returned a non-JSON body") from exc - if not isinstance(parsed_body, dict): - raise Agent365MalformedResponseError("the Entra token endpoint returned a non-object JSON body") - body: Final = parsed_body - if response.status_code >= 400: - raise Agent365TokenExchangeError( - status_code=response.status_code, - error_code=str(body.get("error", "invalid_grant")), - description=str(body.get("error_description", ""))[:512], - aadsts_codes=_parse_aadsts_codes(body.get("error_codes")), - ) - if "access_token" not in body: - raise Agent365MalformedResponseError("the Entra token endpoint returned no access_token") - raw_access_token: Final = body.get("access_token") - if not isinstance(raw_access_token, str) or not raw_access_token: - raise Agent365MalformedResponseError("the Entra token endpoint returned a non-string access_token") - access_token: Final = raw_access_token - expires_at: Final = time.time() + _parse_expires_in(body.get("expires_in", 3599)) - with self._obo_cache_lock: - self._obo_token_cache[cache_key] = (access_token, expires_at) - self._obo_token_cache.move_to_end(cache_key) - while len(self._obo_token_cache) > _OBO_CACHE_MAX_ENTRIES: - self._obo_token_cache.popitem(last=False) - return access_token async def _post_allowing_error_status( self, @@ -577,11 +519,6 @@ class Agent365Guardrail(CustomGuardrail): } raise HTTPException(status_code=503, detail=throttled_detail) - def _evict_obo_token(self, assertion: str) -> None: - cache_key: Final = hashlib.sha256(assertion.encode("utf-8")).hexdigest() - with self._obo_cache_lock: - self._obo_token_cache.pop(cache_key, None) - def _handle_unavailable( self, data: dict, # mutable-ok: guardrail logging appends into the request metadata in place diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index 91ae95eff48..b7dd62559f0 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -313,10 +313,10 @@ class MCPServer(BaseModel): return self.per_server_oauth_discovery and self.auth_type == MCPAuth.oauth2 and not self.has_client_credentials @property - def advertises_gateway_authorization_server(self) -> bool: - """Whether named discovery should advertise the aggregate gateway authorization server.""" - if self.auth_type == MCPAuth.oauth2: - return self.is_gateway_managed_oauth2 and not self.uses_per_server_oauth_relay + def keeps_caller_authorization(self) -> bool: + """Whether the caller's top-level ``Authorization`` stays with the gateway: the server neither relays + it upstream nor runs an OAuth mode that fills that slot itself, so a gateway guardrail may consume it + as the caller's own assertion. Forwarding a separate API-key header leaves the slot untouched.""" if self.auth_type not in ( None, MCPAuth.none, @@ -326,11 +326,20 @@ class MCPServer(BaseModel): MCPAuth.authorization, MCPAuth.token, MCPAuth.aws_sigv4, + MCPAuth.oauth2_token_exchange, ): return False - return not any( - header.lower() in ("authorization", "x-api-key", "api-key", "apikey") - for header in (self.extra_headers or ()) + return not any(header.lower() == "authorization" for header in (self.extra_headers or ())) + + @property + def advertises_gateway_authorization_server(self) -> bool: + """Whether named discovery should advertise the aggregate gateway authorization server.""" + if self.auth_type == MCPAuth.oauth2: + return self.is_gateway_managed_oauth2 and not self.uses_per_server_oauth_relay + if self.auth_type == MCPAuth.oauth2_token_exchange: + return False + return self.keeps_caller_authorization and not any( + header.lower() in ("x-api-key", "api-key", "apikey") for header in (self.extra_headers or ()) ) @property diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_token_exchange_provider.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_token_exchange_provider.py index cf15fdb3e26..002e70da9d2 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_token_exchange_provider.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_token_exchange_provider.py @@ -7,7 +7,9 @@ the I/O edge that maps any transport/HTTP failure to None and parses a JSON body from unittest.mock import patch import pytest +from pydantic import SecretStr +from litellm.proxy._experimental.mcp_server.outbound_credentials import Error, ServerSpec from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchange_provider import ( _post_exchange_endpoint, build_token_exchanger, @@ -17,17 +19,19 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchanger SubjectTokenRejected, TokenExchangeClientError, ) +from litellm.proxy._experimental.mcp_server.outbound_credentials.types import TokenExchangeConfig _HTTP_CLIENT = "litellm.llms.custom_httpx.http_handler.get_async_httpx_client" -def _client_raising_4xx(body: object): - """An httpx client whose POST returns a 4xx whose ``raise_for_status`` raises an HTTPStatusError - carrying ``body`` as its JSON, so the RFC 6749 error-code classification can be driven.""" +def _client_raising_status(status: int, body: object): + """An httpx client whose POST returns ``status`` whose ``raise_for_status`` raises an + HTTPStatusError carrying ``body`` as its JSON, so the RFC 6749 error-code classification can be + driven.""" import httpx request = httpx.Request("POST", "https://idp/token") - response = httpx.Response(400, json=body, request=request) + response = httpx.Response(status, json=body, request=request) class _Resp: def raise_for_status(self) -> None: @@ -80,7 +84,7 @@ async def test_post_parses_json_body_on_success(): ) async def test_post_maps_gateway_fault_4xx_to_client_error(code): # RFC 6749 5.2 gateway-fault codes must raise TokenExchangeClientError (-> 500), not the caller 401. - with patch(_HTTP_CLIENT, return_value=_client_raising_4xx({"error": code})): + with patch(_HTTP_CLIENT, return_value=_client_raising_status(400, {"error": code})): with pytest.raises(TokenExchangeClientError): await _post_exchange_endpoint("https://idp/token", {"grant_type": "x"}, {}) @@ -93,7 +97,7 @@ async def test_post_maps_gateway_fault_4xx_to_client_error(code): ) async def test_post_maps_subject_fault_4xx_to_subject_rejected(body): # A subject-fault code (or an unparseable/absent error) is the caller's problem -> SubjectTokenRejected (401). - with patch(_HTTP_CLIENT, return_value=_client_raising_4xx(body)): + with patch(_HTTP_CLIENT, return_value=_client_raising_status(400, body)): with pytest.raises(SubjectTokenRejected): await _post_exchange_endpoint("https://idp/token", {"grant_type": "x"}, {}) @@ -128,7 +132,7 @@ async def test_post_threads_step_up_error_and_claims_into_subject_rejected(): "error_description": "AADSTS50079: the user must enroll MFA", "claims": claims, } - with patch(_HTTP_CLIENT, return_value=_client_raising_4xx(body)): + with patch(_HTTP_CLIENT, return_value=_client_raising_status(400, body)): with pytest.raises(SubjectTokenRejected) as exc_info: await _post_exchange_endpoint("https://idp/token", {"grant_type": "x"}, {}) assert exc_info.value.claims == claims @@ -137,7 +141,7 @@ async def test_post_threads_step_up_error_and_claims_into_subject_rejected(): @pytest.mark.asyncio async def test_post_subject_rejection_without_claims_carries_none_claims(): - with patch(_HTTP_CLIENT, return_value=_client_raising_4xx({"error": "invalid_grant"})): + with patch(_HTTP_CLIENT, return_value=_client_raising_status(400, {"error": "invalid_grant"})): with pytest.raises(SubjectTokenRejected) as exc_info: await _post_exchange_endpoint("https://idp/token", {"grant_type": "x"}, {}) assert exc_info.value.claims is None @@ -148,6 +152,26 @@ async def test_post_gateway_fault_still_wins_when_claims_are_present(): # A gateway-fault code stays a 500-class TokenExchangeClientError even if the body carries # claims; the caller cannot fix invalid_client by stepping up. body = {"error": "invalid_client", "claims": '{"access_token":{}}'} - with patch(_HTTP_CLIENT, return_value=_client_raising_4xx(body)): + with patch(_HTTP_CLIENT, return_value=_client_raising_status(400, body)): with pytest.raises(TokenExchangeClientError): await _post_exchange_endpoint("https://idp/token", {"grant_type": "x"}, {}) + + +_CONFIG = TokenExchangeConfig( + token_exchange_endpoint="https://idp.example.com/token", + client_id="cid", + client_secret=SecretStr("csec"), + scopes=("s1",), +) +_SERVER = ServerSpec(server_id="srv", resource="https://up.example.com", config=_CONFIG) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("status", [408, 429]) +async def test_exchange_maps_throttled_or_timed_out_4xx_to_upstream_unavailable(status): + # 408/429 are the IdP shedding load, not the caller presenting a bad subject: the exchange must + # surface upstream_unavailable (503-class, retryable) and never tell the caller to sign in again. + with patch(_HTTP_CLIENT, return_value=_client_raising_status(status, {"error": "temporarily_unavailable"})): + result = await OboTokenExchanger(_post_exchange_endpoint).exchange("caller-jwt", _SERVER, _CONFIG) + assert isinstance(result, Error) + assert result.error.tag == "upstream_unavailable" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_caller_sign_in.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_caller_sign_in.py new file mode 100644 index 00000000000..aa853696462 --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_caller_sign_in.py @@ -0,0 +1,134 @@ +from collections.abc import Iterator, Mapping +from typing import Final + +import pytest + +import litellm +from litellm.integrations.custom_guardrail import CustomGuardrail +from litellm.proxy._experimental.mcp_server.caller_sign_in import ( + CallerSignIn, + CallerSignInProvider, + caller_sign_in_for, +) +from litellm.proxy._types import UserAPIKeyAuth +from litellm.types.mcp import MCPAuth, MCPTransport +from litellm.types.mcp_server.mcp_server_manager import MCPServer + + +class _SignInGuardrail(CustomGuardrail): + def __init__(self, issuer: str, scope: str, gated: bool = True) -> None: + super().__init__(guardrail_name=f"sign-in-{issuer}") + self.issuer: Final = issuer + self.scope: Final = scope + self.gated: Final = gated + self.seen: list[tuple[str, str | None]] = [] # mutable-ok: call recorder + + def caller_sign_in(self, server: MCPServer, user_api_key_auth: UserAPIKeyAuth | None) -> CallerSignIn | None: + self.seen.append((server.name, user_api_key_auth.user_id if user_api_key_auth else None)) + if not self.gated: + return None + return CallerSignIn(issuers=(self.issuer,), scopes=(self.scope,)) + + +class _PlainGuardrail(CustomGuardrail): + def __init__(self) -> None: + super().__init__(guardrail_name="plain") + + +def _server(auth_type: MCPAuth | None = None, scopes: list[str] | None = None) -> MCPServer: + return MCPServer( + server_id="tools-id", + name="tools", + server_name="tools", + transport=MCPTransport.http, + url="https://tools.test/mcp", + auth_type=auth_type, + scopes=scopes, + ) + + +@pytest.fixture +def registered() -> Iterator[tuple[_SignInGuardrail, _SignInGuardrail]]: + first: Final = _SignInGuardrail("https://idp-a.test", "scope-a") + second: Final = _SignInGuardrail("https://idp-b.test", "scope-b") + plain: Final = _PlainGuardrail() + for callback in (first, plain, second): + litellm.logging_callback_manager.add_litellm_callback(callback) + try: + yield first, second + finally: + for callback in (first, plain, second): + litellm.logging_callback_manager.remove_callback_from_list_by_object( + litellm.callbacks, callback, require_self=False + ) + + +def test_protocol_matches_only_guardrails_implementing_the_hook(): + assert isinstance(_SignInGuardrail("i", "s"), CallerSignInProvider) + assert not isinstance(_PlainGuardrail(), CallerSignInProvider) + + +def test_no_registered_provider_and_non_obo_advertises_nothing(): + assert caller_sign_in_for(_server(), None) is None + + +def test_registered_providers_merge_in_order_and_dedupe(registered): + first, second = registered + sign_in: Final = caller_sign_in_for(_server(), None) + + assert sign_in is not None + assert sign_in.issuers == ("https://idp-a.test", "https://idp-b.test") + assert sign_in.scopes == ("scope-a", "scope-b") + assert first.seen == [("tools", None)] + assert second.seen == [("tools", None)] + + +def test_provider_returning_none_contributes_nothing(registered): + ungated = _SignInGuardrail("https://idp-c.test", "scope-c", gated=False) + litellm.logging_callback_manager.add_litellm_callback(ungated) + try: + sign_in: Final = caller_sign_in_for(_server(), None) + assert sign_in is not None + assert "https://idp-c.test" not in sign_in.issuers + finally: + litellm.logging_callback_manager.remove_callback_from_list_by_object( + litellm.callbacks, ungated, require_self=False + ) + + +def test_obo_server_contributes_jwt_issuers_and_own_scopes(monkeypatch): + monkeypatch.setenv("JWT_ISSUER", "https://jwt-idp.test") + sign_in: Final = caller_sign_in_for( + _server(auth_type=MCPAuth.oauth2_token_exchange, scopes=["read"]), None + ) + assert sign_in == CallerSignIn(issuers=("https://jwt-idp.test",), scopes=("read",)) + + +def test_obo_server_and_provider_merge_and_dedupe(monkeypatch, registered): + monkeypatch.setenv("JWT_ISSUER", "https://idp-a.test") + sign_in: Final = caller_sign_in_for( + _server(auth_type=MCPAuth.oauth2_token_exchange, scopes=["read", "scope-a"]), None + ) + assert sign_in is not None + assert sign_in.issuers == ("https://idp-a.test", "https://idp-b.test") + assert sign_in.scopes == ("read", "scope-a", "scope-b") + + +def test_obo_server_without_jwt_issuer_still_signs_in_when_a_provider_gates(registered): + sign_in: Final = caller_sign_in_for(_server(auth_type=MCPAuth.oauth2_token_exchange), None) + assert sign_in is not None + assert sign_in.issuers == ("https://idp-a.test", "https://idp-b.test") + + +def test_oauth_utils_strips_the_route_relative_root_path(): + """Regression: Starlette sets ``app_root_path`` to ``""`` on an unmounted app, so the strip must + fall back to ``root_path`` (which is where ``/mcp`` lands when the MCP app is mounted).""" + from litellm.proxy._experimental.mcp_server.oauth_utils import get_route_relative_request_path + + scope: Final[Mapping[str, object]] = { + "type": "http", + "path": "/mcp/catalog", + "root_path": "/mcp", + "app_root_path": "", + } + assert get_route_relative_request_path(scope) == "/catalog" # pyright: ignore[reportArgumentType] diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 4e27ec134d4..7349db67d5d 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -7288,7 +7288,7 @@ async def test_token_exchange_persists_for_oauth2(): # ------------------------------------------------------------------- _OBO_RESOURCE = "https://litellm.example.com/mcp/obo_mcp" -_PATCH_ISSUERS = "litellm.proxy._experimental.mcp_server.discoverable_endpoints._jwt_auth_issuers" +_PATCH_ISSUERS = "litellm.proxy._experimental.mcp_server.caller_sign_in.jwt_auth_issuers" def _obo_server(scopes=None): @@ -7307,15 +7307,15 @@ def _obo_server(scopes=None): ) -def test_obo_protected_resource_response_names_jwt_issuers(): +def test_caller_sign_in_protected_resource_response_names_jwt_issuers(): """An OBO server's PRM points authorization_servers at the configured JWT issuers (the IdP that mints and validates the subject token), with the gateway resource echoed back.""" from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( - _obo_protected_resource_response, + _caller_sign_in_protected_resource_response, ) with patch(_PATCH_ISSUERS, return_value=["https://idp.example.com"]): - response = _obo_protected_resource_response(_obo_server(scopes=["read"]), _OBO_RESOURCE) + response = _caller_sign_in_protected_resource_response(_obo_server(scopes=["read"]), _OBO_RESOURCE) assert response == { "authorization_servers": ["https://idp.example.com"], "resource": _OBO_RESOURCE, @@ -7323,32 +7323,32 @@ def test_obo_protected_resource_response_names_jwt_issuers(): } -def test_obo_protected_resource_response_scopes_default_empty(): +def test_caller_sign_in_protected_resource_response_scopes_default_empty(): """A scopeless OBO server reports scopes_supported as [] rather than None.""" from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( - _obo_protected_resource_response, + _caller_sign_in_protected_resource_response, ) with patch(_PATCH_ISSUERS, return_value=["https://idp.example.com"]): - response = _obo_protected_resource_response(_obo_server(scopes=None), _OBO_RESOURCE) + response = _caller_sign_in_protected_resource_response(_obo_server(scopes=None), _OBO_RESOURCE) assert response["scopes_supported"] == [] -def test_obo_protected_resource_response_falls_back_when_no_issuer(): +def test_caller_sign_in_protected_resource_response_falls_back_when_no_issuer(): """With no JWT issuer configured, the OBO branch returns None so the caller falls back to the gateway-default PRM (discovery still works, it just can't name the IdP).""" from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( - _obo_protected_resource_response, + _caller_sign_in_protected_resource_response, ) with patch(_PATCH_ISSUERS, return_value=[]): - assert _obo_protected_resource_response(_obo_server(), _OBO_RESOURCE) is None + assert _caller_sign_in_protected_resource_response(_obo_server(), _OBO_RESOURCE) is None -def test_obo_protected_resource_response_ignores_non_obo_server(): +def test_caller_sign_in_protected_resource_response_ignores_non_obo_server(): """Non-OBO servers are not handled by this branch (returns None -> gateway default).""" from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( - _obo_protected_resource_response, + _caller_sign_in_protected_resource_response, ) from litellm.proxy._types import MCPTransport from litellm.types.mcp import MCPAuth @@ -7360,7 +7360,7 @@ def test_obo_protected_resource_response_ignores_non_obo_server(): transport=MCPTransport.http, auth_type=MCPAuth.oauth2, ) - assert _obo_protected_resource_response(oauth2_server, _OBO_RESOURCE) is None + assert _caller_sign_in_protected_resource_response(oauth2_server, _OBO_RESOURCE) is None @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 070d87eb012..8d2642a66f1 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -1,4 +1,3 @@ -from litellm.proxy._experimental.mcp_server import operations as mcp_operations import asyncio import contextlib import contextvars @@ -29,6 +28,9 @@ from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS, LATEST_HANDSHAKE_VERS from pydantic import TypeAdapter from starlette.types import Message, Receive, Scope, Send +import litellm +from litellm.integrations.custom_guardrail import CustomGuardrail +from litellm.proxy._experimental.mcp_server import operations as mcp_operations from litellm.proxy._experimental.mcp_server.mcp_context import active_mcp_request_ctx_var from litellm.proxy._experimental.mcp_server.mcp_server_manager import ListedToolsCaller from litellm.proxy._types import ( @@ -1159,6 +1161,7 @@ async def test_mcp_read_resource_success(): ) async def test_read_resource_preserves_content_metadata(_mcp_request_ctx, kind, metadata): from mcp.types import ReadResourceRequestParams + from litellm.proxy._experimental.mcp_server import operations, server uri: Final = "https://example.com/resource" @@ -9881,7 +9884,7 @@ class TestPreemptive401ModeAware: with ( patch.object( mcp_operations.global_mcp_server_manager, - "get_mcp_server_by_name", + "get_mcp_server_answering_to", return_value=server, ), patch.object( @@ -9980,7 +9983,7 @@ class TestPreemptive401ModeAware: patch.dict(os.environ, {"SERVER_ROOT_PATH": "/litellm"}), patch.object( mcp_operations.global_mcp_server_manager, - "get_mcp_server_by_name", + "get_mcp_server_answering_to", return_value=server, ), patch.object( @@ -10087,7 +10090,7 @@ class TestSingleServerPreflightReachesIdJag: with ( patch.object( # test-quality-ok: route wiring must use the manager's configured server mcp_operations.global_mcp_server_manager, - "get_mcp_server_by_name", + "get_mcp_server_answering_to", return_value=server, ), patch.object( # test-quality-ok: route wiring must invoke the manager preflight @@ -10147,7 +10150,7 @@ class TestSingleServerPreflightReachesIdJag: with ( patch.object( # test-quality-ok: route wiring must use the manager's configured server mcp_operations.global_mcp_server_manager, - "get_mcp_server_by_name", + "get_mcp_server_answering_to", return_value=token_exchange, ), patch.object( # test-quality-ok: route wiring must invoke the manager preflight @@ -10212,7 +10215,7 @@ class TestOboPreflightScopedToAllowedServers: preflight = AsyncMock() with ( patch.object( # test-quality-ok: route handler reads the module-level manager, no injection seam - mcp_operations.global_mcp_server_manager, "get_mcp_server_by_name", return_value=requested + mcp_operations.global_mcp_server_manager, "get_mcp_server_answering_to", return_value=requested ), patch.object( # test-quality-ok: the exchanger is the observable; a real one would call an IdP mcp_operations.global_mcp_server_manager, "preflight_token_exchange", preflight @@ -10228,6 +10231,10 @@ class TestOboPreflightScopedToAllowedServers: mcp_server_auth_headers=None, user_api_key_auth=user_api_key_auth, client_ip="10.0.0.7", + raw_headers={ + "x-litellm-api-key": user_api_key_auth.api_key if user_api_key_auth else "", + "authorization": self.SUBJECT_HEADERS["Authorization"], + }, ) return allowed_lookup, preflight @@ -10253,7 +10260,13 @@ class TestOboPreflightScopedToAllowedServers: _, preflight = await self._run(requested, allowed=[requested], user_api_key_auth=key) preflight.assert_awaited_once_with( - server=requested, oauth2_headers=self.SUBJECT_HEADERS, user_api_key_auth=key, raw_headers=None + server=requested, + oauth2_headers=self.SUBJECT_HEADERS, + user_api_key_auth=key, + raw_headers={ + "x-litellm-api-key": key.api_key, + "authorization": self.SUBJECT_HEADERS["Authorization"], + }, ) @@ -10774,8 +10787,8 @@ async def test_native_listing_preserves_empty_result_on_auth_failure(_mcp_reques @pytest.mark.asyncio @pytest.mark.parametrize("failure_hook_raises", [False, True]) async def test_tool_listing_preserves_permission_denial_when_failure_logging_fails(failure_hook_raises): - from litellm.proxy._experimental.mcp_server import operations from litellm.proxy import proxy_server + from litellm.proxy._experimental.mcp_server import operations auth = UserAPIKeyAuth(user_id="denied-caller") denial = HTTPException(status_code=403, detail="scope denied") @@ -10811,8 +10824,9 @@ async def test_legacy_sse_mount_emits_message_endpoint( ) -> None: from starlette.applications import Starlette from starlette.routing import Mount - from litellm.proxy._experimental.mcp_server.faults.list_outcomes import AggregateToolListing + from litellm.proxy._experimental.mcp_server import server as mcp_server + from litellm.proxy._experimental.mcp_server.faults.list_outcomes import AggregateToolListing app: Final = Starlette(routes=[Mount("/mcp", app=mcp_server.app)]) incoming: Final[asyncio.Queue[Message]] = asyncio.Queue() @@ -10970,6 +10984,7 @@ def test_protocol_header_respects_configured_advertisement(revision, rejected): @pytest.mark.asyncio async def test_discovery_adapter_preserves_authenticated_context(_mcp_request_ctx): from mcp.types import DiscoverResult, RequestParams, ServerCapabilities + from litellm.proxy._experimental.mcp_server import server expected = DiscoverResult(supported_versions=["2025-11-25"], capabilities=ServerCapabilities()) @@ -10988,3 +11003,151 @@ async def test_discovery_adapter_preserves_authenticated_context(_mcp_request_ct context = dispatched.await_args.args[1] assert context.user_api_key_auth.user_id == "discover-caller" assert context.mcp_servers == ("allowed",) + + +def _catalog_server() -> MCPServer: + return MCPServer( + server_id="catalog-server-id-001", + name="catalog", + alias="catalog", + server_name="catalog", + url="https://catalog.test/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + mcp_info={"server_name": "catalog"}, + ) + + +class _CallerSignInGuardrail(CustomGuardrail): + def caller_sign_in(self, server, user_api_key_auth): + from litellm.proxy._experimental.mcp_server.caller_sign_in import CallerSignIn + + return CallerSignIn(issuers=("https://idp.test",), scopes=("scope-a",)) + + +class TestConnectChallengeResolver: + """The connect-time sign-in challenge must resolve the server the same way the router resolves + ``/mcp/{name}``: server_id, case-insensitive name, and the short prefix all reach it.""" + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "route_name", + [ + "catalog-server-id-001", + "CATALOG", + ], + ids=["server_id", "uppercase_name"], + ) + async def test_provider_gated_server_challenged_on_every_route_spelling(self, route_name): + from litellm.proxy._experimental.mcp_server import server as server_module + + server = _catalog_server() + guardrail = _CallerSignInGuardrail(guardrail_name="sign-in-stub") + litellm.logging_callback_manager.add_litellm_callback(guardrail) + preflight = AsyncMock() + try: + with ( + patch.object( + mcp_operations.global_mcp_server_manager, + "get_filtered_registry", + return_value={server.server_id: server}, + ), + patch.object( + mcp_operations.global_mcp_server_manager, + "preflight_token_exchange", + preflight, + ), + patch.object( + mcp_operations, + "_get_allowed_mcp_servers", + AsyncMock(return_value=[server]), + ), + pytest.raises(HTTPException) as exc, + ): + await server_module._raise_preemptive_401_for_unauthenticated_servers( + scope={"type": "http", "method": "POST", "path": f"/mcp/{route_name}", "headers": []}, + mcp_servers=[route_name], + oauth2_headers=None, + mcp_server_auth_headers=None, + user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"), + client_ip=None, + ) + finally: + litellm.logging_callback_manager.remove_callback_from_list_by_object( + litellm.callbacks, guardrail, require_self=False + ) + + assert exc.value.status_code == 401 + headers = exc.value.headers or {} + assert "resource_metadata" in (headers.get("WWW-Authenticate") or "") + preflight.assert_not_awaited() + + @pytest.mark.asyncio + async def test_provider_gated_server_challenged_on_short_prefix_route(self): + from litellm.proxy._experimental.mcp_server import server as server_module + from litellm.proxy._experimental.mcp_server.utils import compute_short_server_prefix + + server = _catalog_server() + route_name = compute_short_server_prefix(server.server_id) + guardrail = _CallerSignInGuardrail(guardrail_name="sign-in-stub") + litellm.logging_callback_manager.add_litellm_callback(guardrail) + try: + with ( + patch.object( + mcp_operations.global_mcp_server_manager, + "get_filtered_registry", + return_value={server.server_id: server}, + ), + patch.object( + mcp_operations, + "_get_allowed_mcp_servers", + AsyncMock(return_value=[server]), + ), + pytest.raises(HTTPException) as exc, + ): + await server_module._raise_preemptive_401_for_unauthenticated_servers( + scope={"type": "http", "method": "POST", "path": f"/mcp/{route_name}", "headers": []}, + mcp_servers=[route_name], + oauth2_headers=None, + mcp_server_auth_headers=None, + user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"), + client_ip=None, + ) + finally: + litellm.logging_callback_manager.remove_callback_from_list_by_object( + litellm.callbacks, guardrail, require_self=False + ) + + assert exc.value.status_code == 401 + + @pytest.mark.asyncio + async def test_obo_challenge_www_authenticate_matches_main_byte_for_byte(self, monkeypatch): + """The provider redesign must not change what an OBO server challenges with: the relative + RFC 9728 resource_metadata path plus the RFC 6750 invalid_token triple, exactly as main.""" + from litellm.proxy._experimental.mcp_server import server as server_module + + monkeypatch.delenv("SERVER_ROOT_PATH", raising=False) + obo = _make_obo_server("obo") + with ( + patch.object( + mcp_operations.global_mcp_server_manager, + "get_mcp_server_answering_to", + return_value=obo, + ), + pytest.raises(HTTPException) as exc, + ): + await server_module._raise_preemptive_401_for_unauthenticated_servers( + scope={"type": "http", "method": "POST", "path": "/mcp/obo", "headers": []}, + mcp_servers=["obo"], + oauth2_headers=None, + mcp_server_auth_headers=None, + user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"), + client_ip=None, + ) + + assert exc.value.status_code == 401 + assert (exc.value.headers or {}).get("WWW-Authenticate") == ( + 'Bearer resource_metadata="/.well-known/oauth-protected-resource/mcp/obo", ' + 'error="invalid_token", ' + 'error_description="Missing or invalid subject token; authenticate with the IdP and retry"' + ) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index f361b8bfd09..97b455bfd39 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -1464,6 +1464,19 @@ class TestMCPServerManager: assert server.uses_per_server_oauth_relay is True assert server.advertises_gateway_authorization_server is False + @pytest.mark.asyncio + async def test_load_servers_from_config_does_not_advertise_gateway_as_for_token_exchange(self): + # keeps_caller_authorization includes oauth2_token_exchange so a sign-in provider may gate it, + # but named discovery must still fall to the server's own PRM rather than the aggregate AS. + manager = MCPServerManager() + + with patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=None)): + await manager.load_servers_from_config(self._client_forwarded_config(MCPAuth.oauth2_token_exchange)) + + server = next(iter(manager.config_mcp_servers.values())) + assert server.keeps_caller_authorization is True + assert server.advertises_gateway_authorization_server is False + @pytest.mark.asyncio @pytest.mark.parametrize( "config", diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py index e72b716665c..cfd6745308d 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py @@ -14,6 +14,14 @@ from litellm.exceptions import Timeout as LitellmTimeout from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.litellm_core_utils.secret_redaction import redact_string +from litellm.proxy._experimental.mcp_server.caller_sign_in import CallerSignIn, caller_sign_in_for +from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import OAuthToken +from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Error, Ok, Result +from litellm.proxy._experimental.mcp_server.outbound_credentials.types import ( + CredError, + ServerSpec, + TokenExchangeConfig, +) from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.guardrails.guardrail_hooks.agent_365 import ( Agent365Guardrail, @@ -27,9 +35,12 @@ from litellm.types.guardrails import ( LitellmParams, SupportedGuardrailIntegrations, ) +from litellm.types.mcp import MCPAuth, MCPTransport +from litellm.types.mcp_server.mcp_server_manager import MCPServer from litellm.types.proxy.guardrails.guardrail_hooks.agent_365 import ( AGENT_365_PROD_API_BASE, AGENT_365_PROD_RESOURCE_APP_ID, + AGENT_365_SCOPE_NAME, Agent365GuardrailConfigModel, ) @@ -45,8 +56,46 @@ def _response(status_code: int, payload: Any = None, text: str | None = None) -> return httpx.Response(status_code=status_code, text=text or "", request=request) -def _token_response(access_token: str = "obo-access-token", expires_in: int = 3599) -> httpx.Response: - return _response(200, {"access_token": access_token, "expires_in": expires_in}) +class StubTokenExchanger: + """The TokenExchanger the guardrail is injected with in tests: programmed Result queue plus a + per-subject cache honoring ``expires_at``, so cache and evaluate-401-invalidate behavior is + exercised the way the real OboTokenExchanger drives it.""" + + def __init__(self, results: list[Result[OAuthToken, CredError] | BaseException] | None = None): + self._results = list(results or []) + self._cache: dict[str, OAuthToken] = {} + self.calls: list[tuple[str, ServerSpec, TokenExchangeConfig]] = [] + self.invalidations: list[str] = [] + + async def exchange( + self, subject_token: str, server: ServerSpec, config: TokenExchangeConfig, *, tenant_id: str = "" + ) -> Result[OAuthToken, CredError]: + cached: Final = self._cache.get(subject_token) + if cached is not None and (cached.expires_at is None or cached.expires_at > time.time()): + return Ok(cached) + self.calls.append((subject_token, server, config)) + if not self._results: + raise AssertionError("StubTokenExchanger ran out of programmed results") + result = self._results.pop(0) + if isinstance(result, BaseException): + raise result + if isinstance(result, Ok): + self._cache[subject_token] = result.ok + return result + + async def invalidate( + self, subject_token: str, server: ServerSpec, config: TokenExchangeConfig, *, tenant_id: str = "" + ) -> None: + self.invalidations.append(subject_token) + self._cache.pop(subject_token, None) + + +def _ok_exchange(access_token: str = "obo-access-token", expires_in: int = 3599) -> Ok[OAuthToken, CredError]: + return Ok(OAuthToken(access_token=access_token, expires_at=time.time() + expires_in)) + + +def _obo_ok(access_token: str = "obo-access-token") -> list[Result[OAuthToken, CredError]]: + return [_ok_exchange(access_token)] def _allow_response(correlation_id: str = "corr-1") -> httpx.Response: @@ -121,7 +170,9 @@ class FakeHandler: def _make_guardrail( handler: FakeHandler, *, + exchanger: StubTokenExchanger | None = None, unreachable_fallback: str = "fail_closed", + default_on: bool = True, ) -> Agent365Guardrail: return Agent365Guardrail( guardrail_name="agent-365-guard", @@ -130,21 +181,28 @@ def _make_guardrail( client_secret="secret-123", unreachable_fallback=unreachable_fallback, async_handler=handler, + token_exchanger=exchanger if exchanger is not None else StubTokenExchanger(_obo_ok()), event_hook="pre_mcp_call", - default_on=True, + default_on=default_on, ) -def _default_fallback_guardrail(handler: FakeHandler) -> Agent365Guardrail: - return Agent365Guardrail( - guardrail_name="agent-365-guard", - tenant_id="tenant-abc", - client_id="client-xyz", - client_secret="secret-123", - async_handler=handler, - event_hook="pre_mcp_call", - default_on=True, - ) +def _default_fallback_guardrail( + handler: FakeHandler, exchanger: StubTokenExchanger | None = None +) -> Agent365Guardrail: + return _make_guardrail(handler, exchanger=exchanger) + + +def _server(**overrides: Any) -> MCPServer: + kwargs: Final[dict] = { + "server_id": "outlook-id", + "name": "outlook_mcp", + "server_name": "outlook_mcp", + "transport": MCPTransport.http, + "url": "https://outlook.test/mcp", + } + kwargs.update(overrides) + return MCPServer(**kwargs) def _mcp_data(**overrides: Any) -> dict: @@ -257,14 +315,13 @@ class TestInitializeGuardrail: resource_app_id="00000000-0000-0000-0000-000000000000", agent_id="yaml-agent", ) - handler: Final = FakeHandler([_token_response(), _allow_response()]) + handler: Final = FakeHandler([_allow_response()]) with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): - guardrail: Final = initialize_guardrail(params, {"guardrail_name": "a365-stale"}, async_handler=handler) + initialize_guardrail(params, {"guardrail_name": "a365-stale"}, async_handler=handler) assert "ignoring api_base, resource_app_id, agent_id" in caplog.text + guardrail: Final = _make_guardrail(handler) await _run(guardrail, _mcp_data()) - token_call, evaluate_call = handler.calls - assert token_call.url == TOKEN_URL - assert token_call.data["scope"] == f"{AGENT_365_PROD_RESOURCE_APP_ID}/ThreatProtection.Evaluate.All" + evaluate_call: Final = handler.calls[0] assert evaluate_call.url == EVALUATE_URL assert evaluate_call.json["agentId"] == "my-agent-key" @@ -305,7 +362,7 @@ def _guardrail_info(data: dict) -> dict: class TestAllowFlow: @pytest.mark.asyncio async def test_allowed_call_passes_through(self): - handler: Final = FakeHandler([_token_response(), _allow_response()]) + handler: Final = FakeHandler([_allow_response()]) guardrail: Final = _make_guardrail(handler) data: Final = _mcp_data() result: Final = await _run(guardrail, data) @@ -319,25 +376,27 @@ class TestAllowFlow: assert info["guardrail_response"]["latency_ms"] >= 0 @pytest.mark.asyncio - async def test_obo_exchange_form(self): - handler: Final = FakeHandler([_token_response(), _allow_response()]) - guardrail: Final = _make_guardrail(handler) + async def test_obo_exchange_uses_the_built_entra_obo_config(self): + exchanger: Final = StubTokenExchanger(_obo_ok()) + handler: Final = FakeHandler([_allow_response()]) + guardrail: Final = _make_guardrail(handler, exchanger=exchanger) await _run(guardrail, _mcp_data()) - token_call: Final = handler.calls[0] - assert token_call.url == TOKEN_URL - assert token_call.data["grant_type"] == "urn:ietf:params:oauth:grant-type:jwt-bearer" - assert token_call.data["requested_token_use"] == "on_behalf_of" - assert token_call.data["assertion"] == FAKE_ASSERTION - assert token_call.data["client_id"] == "client-xyz" - assert token_call.data["client_secret"] == "secret-123" - assert token_call.data["scope"] == f"{AGENT_365_PROD_RESOURCE_APP_ID}/ThreatProtection.Evaluate.All" + subject_token, server, config = exchanger.calls[0] + assert subject_token == FAKE_ASSERTION + assert server.server_id == "agent-365:tenant-abc" + assert server.resource == AGENT_365_PROD_API_BASE + assert config.profile == "entra_obo" + assert config.token_exchange_endpoint == TOKEN_URL + assert config.client_id == "client-xyz" + assert config.client_secret is not None and config.client_secret.get_secret_value() == "secret-123" + assert config.scopes == (f"{AGENT_365_PROD_RESOURCE_APP_ID}/{AGENT_365_SCOPE_NAME}",) @pytest.mark.asyncio async def test_evaluate_payload(self): - handler: Final = FakeHandler([_token_response(), _allow_response()]) + handler: Final = FakeHandler([_allow_response()]) guardrail: Final = _make_guardrail(handler) await _run(guardrail, _mcp_data()) - evaluate_call: Final = handler.calls[1] + evaluate_call: Final = handler.calls[0] assert evaluate_call.url == EVALUATE_URL assert evaluate_call.headers["Authorization"] == "Bearer obo-access-token" assert evaluate_call.json["tool"] == {"name": "send_email"} @@ -348,11 +407,11 @@ class TestAllowFlow: @pytest.mark.asyncio async def test_evaluate_payload_includes_listed_tool_metadata(self): - handler: Final = FakeHandler([_token_response(), _allow_response()]) + handler: Final = FakeHandler([_allow_response()]) guardrail: Final = _make_guardrail(handler) schema: Final = {"type": "object", "properties": {"to": {"type": "string"}}, "required": ["to"]} - await _run(guardrail, _mcp_data(mcp_tool_description="Send an email", mcp_input_schema=schema)) - assert handler.calls[1].json["tool"] == { + await _run(guardrail, _mcp_data(mcp_tool_description="Send an email", mcp_tool_input_schema=schema)) + assert handler.calls[0].json["tool"] == { "name": "send_email", "description": "Send an email", "inputSchema": schema, @@ -364,10 +423,17 @@ class TestAllowFlow: [(None, None), ("", None), (None, ["not", "a", "schema"]), (42, "type: object")], ) async def test_evaluate_payload_omits_missing_or_malformed_tool_metadata(self, description, schema): - handler: Final = FakeHandler([_token_response(), _allow_response()]) + handler: Final = FakeHandler([_allow_response()]) guardrail: Final = _make_guardrail(handler) - await _run(guardrail, _mcp_data(mcp_tool_description=description, mcp_input_schema=schema)) - assert handler.calls[1].json["tool"] == {"name": "send_email"} + await _run(guardrail, _mcp_data(mcp_tool_description=description, mcp_tool_input_schema=schema)) + assert handler.calls[0].json["tool"] == {"name": "send_email"} + + @pytest.mark.asyncio + async def test_agent_id_falls_back_to_key_alias(self): + handler: Final = FakeHandler([_allow_response()]) + guardrail: Final = _make_guardrail(handler) + await _run(guardrail, _mcp_data()) + assert handler.calls[0].json["agentId"] == "my-agent-key" @pytest.mark.asyncio async def test_non_mcp_call_type_skipped(self): @@ -385,26 +451,26 @@ class TestConversationId: @pytest.mark.asyncio async def test_two_calls_in_one_session_share_the_conversation_id(self): - handler: Final = FakeHandler([_token_response(), _allow_response(), _allow_response()]) + handler: Final = FakeHandler([_allow_response(), _allow_response()]) guardrail: Final = _make_guardrail(handler) for call_id in ("call-1", "call-2"): await _run( guardrail, _mcp_data(litellm_call_id=call_id, litellm_logging_obj=_logging_obj(call_id, mcp_session_id="sess-A")), ) - assert [call.json["conversationId"] for call in handler.calls[1:]] == ["sess-A", "sess-A"] + assert [call.json["conversationId"] for call in handler.calls[:]] == ["sess-A", "sess-A"] @pytest.mark.asyncio async def test_server_recorded_session_beats_the_client_header(self): - handler: Final = FakeHandler([_token_response(), _allow_response()]) + handler: Final = FakeHandler([_allow_response()]) guardrail: Final = _make_guardrail(handler) data: Final = _mcp_data(litellm_logging_obj=_logging_obj("call-id-1", mcp_session_id="sess-from-logging")) await _run(guardrail, data) - assert handler.calls[1].json["conversationId"] == "sess-from-logging" + assert handler.calls[0].json["conversationId"] == "sess-from-logging" @pytest.mark.asyncio async def test_sessionless_call_falls_back_to_the_request_call_id(self): - handler: Final = FakeHandler([_token_response(), _allow_response()]) + handler: Final = FakeHandler([_allow_response()]) guardrail: Final = _make_guardrail(handler) data: Final = _mcp_data( metadata={"headers": {}}, @@ -412,37 +478,37 @@ class TestConversationId: litellm_logging_obj=_logging_obj("call-id-from-logging"), ) await _run(guardrail, data) - assert handler.calls[1].json["conversationId"] == "call-id-from-data" + assert handler.calls[0].json["conversationId"] == "call-id-from-data" @pytest.mark.asyncio async def test_sessionless_call_without_request_call_id_uses_the_logging_call_id(self): - handler: Final = FakeHandler([_token_response(), _allow_response()]) + handler: Final = FakeHandler([_allow_response()]) guardrail: Final = _make_guardrail(handler) data: Final = _mcp_data(metadata={"headers": {}}, litellm_logging_obj=_logging_obj("call-id-from-logging")) await _run(guardrail, data) - assert handler.calls[1].json["conversationId"] == "call-id-from-logging" + assert handler.calls[0].json["conversationId"] == "call-id-from-logging" @pytest.mark.asyncio async def test_session_id_header_case_insensitive(self): - handler: Final = FakeHandler([_token_response(), _allow_response()]) + handler: Final = FakeHandler([_allow_response()]) guardrail: Final = _make_guardrail(handler) data: Final = _mcp_data(metadata={"headers": {"Mcp-Session-Id": "sess-CASED"}}) await _run(guardrail, data) - assert handler.calls[1].json["conversationId"] == "sess-CASED" + assert handler.calls[0].json["conversationId"] == "sess-CASED" @pytest.mark.asyncio async def test_generates_uuid_when_no_identifier_available(self): - handler: Final = FakeHandler([_token_response(), _allow_response()]) + handler: Final = FakeHandler([_allow_response()]) guardrail: Final = _make_guardrail(handler) await _run(guardrail, _mcp_data(metadata={"headers": {}}, litellm_logging_obj=_logging_obj(""))) - conversation_id: Final = handler.calls[1].json["conversationId"] + conversation_id: Final = handler.calls[0].json["conversationId"] assert uuid.UUID(conversation_id).version == 4 class TestBlockFlow: @pytest.mark.asyncio async def test_blocked_call_raises_400(self): - handler: Final = FakeHandler([_token_response(), _block_response(message="Injection detected")]) + handler: Final = FakeHandler([_block_response(message="Injection detected")]) guardrail: Final = _make_guardrail(handler) data: Final = _mcp_data() with pytest.raises(HTTPException) as exc_info: @@ -458,7 +524,7 @@ class TestBlockFlow: @pytest.mark.asyncio async def test_blocked_even_with_fail_open(self): - handler: Final = FakeHandler([_token_response(), _block_response()]) + handler: Final = FakeHandler([_block_response()]) guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open") with pytest.raises(HTTPException) as exc_info: await _run(guardrail, _mcp_data()) @@ -467,7 +533,7 @@ class TestBlockFlow: @pytest.mark.asyncio @pytest.mark.parametrize("status", ["Skipped", "FailedOpen"]) async def test_explicit_block_wins_over_non_evaluated_status(self, status): - handler: Final = FakeHandler([_token_response(), _block_response(status=status)]) + handler: Final = FakeHandler([_block_response(status=status)]) guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open") data: Final = _mcp_data() with pytest.raises(HTTPException) as exc_info: @@ -483,7 +549,7 @@ class TestDefenderNotEvaluated: @pytest.mark.asyncio @pytest.mark.parametrize("status", ["Skipped", "FailedOpen"]) async def test_fail_closed_blocks_allowed_but_unevaluated_call(self, status): - handler: Final = FakeHandler([_token_response(), _not_evaluated_response(status)]) + handler: Final = FakeHandler([_not_evaluated_response(status)]) guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_closed") data: Final = _mcp_data() with pytest.raises(HTTPException) as exc_info: @@ -500,7 +566,7 @@ class TestDefenderNotEvaluated: @pytest.mark.asyncio @pytest.mark.parametrize("status", ["Skipped", "FailedOpen"]) async def test_fail_open_allows_unevaluated_call_as_unscanned(self, status): - handler: Final = FakeHandler([_token_response(), _not_evaluated_response(status)]) + handler: Final = FakeHandler([_not_evaluated_response(status)]) guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open") data: Final = _mcp_data() result: Final = await _run(guardrail, data) @@ -514,7 +580,7 @@ class TestDefenderNotEvaluated: @pytest.mark.asyncio @pytest.mark.parametrize("payload", [{"allowed": True}, {"allowed": True, "defender": {"verdict": "Allow"}}]) async def test_allowed_without_defender_status_is_not_an_evaluated_allow(self, payload): - handler: Final = FakeHandler([_token_response(), _response(200, payload)]) + handler: Final = FakeHandler([_response(200, payload)]) guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_closed") data: Final = _mcp_data() with pytest.raises(HTTPException) as exc_info: @@ -525,7 +591,7 @@ class TestDefenderNotEvaluated: @pytest.mark.asyncio async def test_http_400_always_blocks_even_fail_open(self): - handler: Final = FakeHandler([_token_response(), _response(400, text="Bad request: serverName missing")]) + handler: Final = FakeHandler([_response(400, text="Bad request: serverName missing")]) guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open") with pytest.raises(HTTPException) as exc_info: await _run(guardrail, _mcp_data()) @@ -534,18 +600,20 @@ class TestDefenderNotEvaluated: AVAILABILITY_FAILURES: Final = ( - pytest.param([_token_response(), httpx.ReadTimeout("timed out")], id="evaluate-timeout"), - pytest.param([_token_response(), _response(502, text="bad gateway")], id="evaluate-5xx"), - pytest.param([_token_response(), _not_evaluated_response("Skipped")], id="evaluate-skipped"), - pytest.param([_response(503, text="entra down")], id="entra-5xx"), + pytest.param([httpx.ReadTimeout("timed out")], _obo_ok(), id="evaluate-timeout"), + pytest.param([_response(502, text="bad gateway")], _obo_ok(), id="evaluate-5xx"), + pytest.param([_not_evaluated_response("Skipped")], _obo_ok(), id="evaluate-skipped"), + pytest.param([], [Error(CredError.of_upstream_unavailable("entra down"))], id="entra-5xx"), ) class TestFailOpenOptIn: @pytest.mark.asyncio - @pytest.mark.parametrize("responses", AVAILABILITY_FAILURES) - async def test_constructor_default_blocks_each_availability_failure_with_503(self, responses): - guardrail: Final = _default_fallback_guardrail(FakeHandler(responses)) + @pytest.mark.parametrize(("responses", "exchange_results"), AVAILABILITY_FAILURES) + async def test_constructor_default_blocks_each_availability_failure_with_503(self, responses, exchange_results): + guardrail: Final = _default_fallback_guardrail( + FakeHandler(responses), StubTokenExchanger(exchange_results) + ) assert guardrail.unreachable_fallback == "fail_closed" with pytest.raises(HTTPException) as exc_info: await _run(guardrail, _mcp_data()) @@ -553,9 +621,15 @@ class TestFailOpenOptIn: assert "fail_closed" in exc_info.value.detail["message"] @pytest.mark.asyncio - @pytest.mark.parametrize("responses", AVAILABILITY_FAILURES) - async def test_opted_in_fail_open_lets_each_availability_failure_through_as_failed_to_respond(self, responses): - guardrail: Final = _make_guardrail(FakeHandler(responses), unreachable_fallback="fail_open") + @pytest.mark.parametrize(("responses", "exchange_results"), AVAILABILITY_FAILURES) + async def test_opted_in_fail_open_lets_each_availability_failure_through_as_failed_to_respond( + self, responses, exchange_results + ): + guardrail: Final = _make_guardrail( + FakeHandler(responses), + exchanger=StubTokenExchanger(exchange_results), + unreachable_fallback="fail_open", + ) data: Final = _mcp_data() assert await _run(guardrail, data) is data info: Final = _guardrail_info(data) @@ -564,7 +638,7 @@ class TestFailOpenOptIn: @pytest.mark.asyncio async def test_opted_in_fail_open_logs_the_unscanned_call_at_error_level(self, caplog): - handler: Final = FakeHandler([_token_response(), httpx.ReadTimeout("timed out")]) + handler: Final = FakeHandler([httpx.ReadTimeout("timed out")]) guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open") with caplog.at_level(logging.ERROR, logger="LiteLLM Proxy"): await _run(guardrail, _mcp_data()) @@ -573,7 +647,7 @@ class TestFailOpenOptIn: @pytest.mark.asyncio async def test_opted_in_fail_open_still_blocks_a_policy_block(self): - handler: Final = FakeHandler([_token_response(), _block_response()]) + handler: Final = FakeHandler([_block_response()]) guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open") with pytest.raises(HTTPException) as exc_info: await _run(guardrail, _mcp_data()) @@ -584,10 +658,7 @@ class TestUnreachableFallback: @pytest.mark.asyncio async def test_evaluate_litellm_timeout_fail_closed(self): handler: Final = FakeHandler( - [ - _token_response(), - LitellmTimeout(message="Connection timed out", model="default-model-name", llm_provider="httpx"), - ] + [LitellmTimeout(message="Connection timed out", model="default-model-name", llm_provider="httpx")] ) guardrail: Final = _make_guardrail(handler) with pytest.raises(HTTPException) as exc_info: @@ -596,7 +667,7 @@ class TestUnreachableFallback: @pytest.mark.asyncio async def test_evaluate_timeout_fail_closed(self): - handler: Final = FakeHandler([_token_response(), httpx.ReadTimeout("timed out")]) + handler: Final = FakeHandler([httpx.ReadTimeout("timed out")]) guardrail: Final = _make_guardrail(handler) with pytest.raises(HTTPException) as exc_info: await _run(guardrail, _mcp_data()) @@ -605,7 +676,7 @@ class TestUnreachableFallback: @pytest.mark.asyncio async def test_evaluate_timeout_fail_open(self): - handler: Final = FakeHandler([_token_response(), httpx.ReadTimeout("timed out")]) + handler: Final = FakeHandler([httpx.ReadTimeout("timed out")]) guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open") data: Final = _mcp_data() result: Final = await _run(guardrail, data) @@ -616,7 +687,7 @@ class TestUnreachableFallback: @pytest.mark.asyncio async def test_evaluate_5xx_fail_closed(self): - handler: Final = FakeHandler([_token_response(), _response(502, text="bad gateway")]) + handler: Final = FakeHandler([_response(502, text="bad gateway")]) guardrail: Final = _make_guardrail(handler) with pytest.raises(HTTPException) as exc_info: await _run(guardrail, _mcp_data()) @@ -634,11 +705,13 @@ class TestUnreachableFallback: @pytest.mark.asyncio async def test_non_jwt_bearer_token_fail_closed(self): + exchanger: Final = StubTokenExchanger() handler: Final = FakeHandler([]) - guardrail: Final = _make_guardrail(handler) + guardrail: Final = _make_guardrail(handler, exchanger=exchanger) with pytest.raises(HTTPException) as exc_info: await _run(guardrail, _mcp_data(incoming_bearer_token="sk-litellm-virtual-key")) assert exc_info.value.status_code == 401 + assert exchanger.calls == [] @pytest.mark.asyncio async def test_missing_bearer_token_blocks_even_fail_open(self): @@ -655,18 +728,19 @@ class TestUnreachableFallback: @pytest.mark.asyncio async def test_obo_rejected_blocks_even_fail_open(self): - handler: Final = FakeHandler( - [_response(400, {"error": "invalid_grant", "error_description": "AADSTS50013: bad assertion"})] + exchanger: Final = StubTokenExchanger( + [Error(CredError.of_unauthorized("IdP rejected the subject token (HTTP 400)", claims="invalid_grant"))] ) - guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open") + handler: Final = FakeHandler([]) + guardrail: Final = _make_guardrail(handler, exchanger=exchanger, unreachable_fallback="fail_open") with pytest.raises(HTTPException) as exc_info: await _run(guardrail, _mcp_data()) assert exc_info.value.status_code == 401 - assert "invalid_grant" in exc_info.value.detail["message"] + assert "On-Behalf-Of token exchange was rejected" in exc_info.value.detail["message"] @pytest.mark.asyncio async def test_evaluate_4xx_blocks_even_fail_open(self): - handler: Final = FakeHandler([_token_response(), _response(403, text="obo token lacks the scope")]) + handler: Final = FakeHandler([_response(403, text="obo token lacks the scope")]) guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open") data: Final = _mcp_data() with pytest.raises(HTTPException) as exc_info: @@ -680,24 +754,24 @@ class TestUnreachableFallback: @pytest.mark.asyncio async def test_obo_rejected_fail_closed(self): - handler: Final = FakeHandler( - [_response(400, {"error": "invalid_grant", "error_description": "AADSTS50013: bad assertion"})] + exchanger: Final = StubTokenExchanger( + [Error(CredError.of_unauthorized("IdP rejected the subject token (HTTP 400)"))] ) - guardrail: Final = _make_guardrail(handler) + handler: Final = FakeHandler([]) + guardrail: Final = _make_guardrail(handler, exchanger=exchanger) with pytest.raises(HTTPException) as exc_info: await _run(guardrail, _mcp_data()) assert exc_info.value.status_code == 401 - assert "invalid_grant" in exc_info.value.detail["message"] + assert "On-Behalf-Of token exchange was rejected" in exc_info.value.detail["message"] @pytest.mark.asyncio @pytest.mark.parametrize( "error_code", ["invalid_client", "unauthorized_client", "invalid_scope", "invalid_resource"] ) async def test_gateway_credential_rejection_is_unavailable_not_a_caller_401(self, error_code: str): - handler: Final = FakeHandler( - [_response(401, {"error": error_code, "error_description": "AADSTS7000215: invalid client secret"})] - ) - guardrail: Final = _make_guardrail(handler) + exchanger: Final = StubTokenExchanger([Error(CredError.of_misconfigured(error_code))]) + handler: Final = FakeHandler([]) + guardrail: Final = _make_guardrail(handler, exchanger=exchanger) data: Final = _mcp_data() with pytest.raises(HTTPException) as exc_info: await _run(guardrail, data) @@ -710,23 +784,12 @@ class TestUnreachableFallback: assert "client_secret" in info["guardrail_response"]["reason"] @pytest.mark.asyncio - @pytest.mark.parametrize("aadsts_code", [5002710, 5002723], ids=["malformed-header", "no-kid"]) - async def test_malformed_assertion_reported_as_invalid_client_is_a_caller_401(self, aadsts_code: int): - """Entra answers ``invalid_client`` for a forged or garbled assertion (AADSTS50027xx) exactly as for a - bad gateway secret; the sub-code is what says the caller, not the gateway, has to fix it.""" - handler: Final = FakeHandler( - [ - _response( - 401, - { - "error": "invalid_client", - "error_description": f"AADSTS{aadsts_code}: Invalid JWT token.", - "error_codes": [aadsts_code], - }, - ) - ] + async def test_caller_rejection_reason_does_not_blame_the_gateway_credentials(self): + exchanger: Final = StubTokenExchanger( + [Error(CredError.of_unauthorized("IdP rejected the subject token (HTTP 401)"))] ) - guardrail: Final = _make_guardrail(handler) + handler: Final = FakeHandler([]) + guardrail: Final = _make_guardrail(handler, exchanger=exchanger) data: Final = _mcp_data() with pytest.raises(HTTPException) as exc_info: await _run(guardrail, data) @@ -735,10 +798,9 @@ class TestUnreachableFallback: @pytest.mark.asyncio async def test_gateway_credential_rejection_follows_fail_open(self): - handler: Final = FakeHandler( - [_response(401, {"error": "invalid_client", "error_description": "AADSTS7000215: invalid client secret"})] - ) - guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open") + exchanger: Final = StubTokenExchanger([Error(CredError.of_misconfigured("invalid_client"))]) + handler: Final = FakeHandler([]) + guardrail: Final = _make_guardrail(handler, exchanger=exchanger, unreachable_fallback="fail_open") data: Final = _mcp_data() result: Final = await _run(guardrail, data) assert result is data @@ -747,10 +809,43 @@ class TestUnreachableFallback: assert info["guardrail_response"]["verdict"] == "Unscanned" assert "invalid_client" in info["guardrail_response"]["reason"] + @pytest.mark.asyncio + async def test_exchange_upstream_unavailable_is_unavailable_with_the_summary_not_a_caller_401(self): + exchanger: Final = StubTokenExchanger( + [Error(CredError.of_upstream_unavailable("token exchange did not return a usable access token"))] + ) + handler: Final = FakeHandler([]) + guardrail: Final = _make_guardrail(handler, exchanger=exchanger) + data: Final = _mcp_data() + with pytest.raises(HTTPException) as exc_info: + await _run(guardrail, data) + assert exc_info.value.status_code == 503 + info: Final = _guardrail_info(data) + assert info["guardrail_status"] == "guardrail_failed_to_respond" + assert info["guardrail_response"]["verdict"] == "Unavailable" + assert ( + info["guardrail_response"]["reason"] + == "the Entra token exchange failed (upstream unavailable: token exchange did not return a usable access token)" + ) + + @pytest.mark.asyncio + async def test_exchange_upstream_unavailable_follows_fail_open(self): + exchanger: Final = StubTokenExchanger([Error(CredError.of_upstream_unavailable("Entra throttled the exchange"))]) + handler: Final = FakeHandler([]) + guardrail: Final = _make_guardrail(handler, exchanger=exchanger, unreachable_fallback="fail_open") + data: Final = _mcp_data() + result: Final = await _run(guardrail, data) + assert result is data + info: Final = _guardrail_info(data) + assert info["guardrail_status"] == "guardrail_failed_to_respond" + assert info["guardrail_response"]["verdict"] == "Unscanned" + assert "Entra throttled the exchange" in info["guardrail_response"]["reason"] + @pytest.mark.asyncio async def test_obo_endpoint_5xx_fail_open(self): - handler: Final = FakeHandler([_response(503, text="entra down")]) - guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open") + exchanger: Final = StubTokenExchanger([Error(CredError.of_upstream_unavailable("entra down"))]) + handler: Final = FakeHandler([]) + guardrail: Final = _make_guardrail(handler, exchanger=exchanger, unreachable_fallback="fail_open") data: Final = _mcp_data() result: Final = await _run(guardrail, data) assert result is data @@ -762,47 +857,33 @@ class TestUnreachableFallback: class TestOboTokenCache: @pytest.mark.asyncio async def test_same_assertion_reuses_token(self): - handler: Final = FakeHandler([_token_response(), _allow_response(), _allow_response()]) - guardrail: Final = _make_guardrail(handler) + exchanger: Final = StubTokenExchanger(_obo_ok()) + handler: Final = FakeHandler([_allow_response(), _allow_response()]) + guardrail: Final = _make_guardrail(handler, exchanger=exchanger) await _run(guardrail, _mcp_data()) await _run(guardrail, _mcp_data()) - token_calls: Final = [c for c in handler.calls if c.url == TOKEN_URL] - assert len(token_calls) == 1 + assert len(exchanger.calls) == 1 @pytest.mark.asyncio async def test_different_assertions_get_distinct_tokens(self): other_assertion: Final = "eyJhbGciOi.eyJvdGhlciI.b3RoZXJzaWc" - handler: Final = FakeHandler( - [ - _token_response(access_token="token-a"), - _allow_response(), - _token_response(access_token="token-b"), - _allow_response(), - ] - ) - guardrail: Final = _make_guardrail(handler) + exchanger: Final = StubTokenExchanger([_ok_exchange("token-a"), _ok_exchange("token-b")]) + handler: Final = FakeHandler([_allow_response(), _allow_response()]) + guardrail: Final = _make_guardrail(handler, exchanger=exchanger) await _run(guardrail, _mcp_data()) await _run(guardrail, _mcp_data(incoming_bearer_token=other_assertion)) - token_calls: Final = [c for c in handler.calls if c.url == TOKEN_URL] - assert len(token_calls) == 2 - assert handler.calls[3].headers["Authorization"] == "Bearer token-b" + assert len(exchanger.calls) == 2 + assert handler.calls[1].headers["Authorization"] == "Bearer token-b" @pytest.mark.asyncio async def test_expired_token_refreshed(self): - handler: Final = FakeHandler( - [ - _token_response(access_token="short-lived", expires_in=1), - _allow_response(), - _token_response(access_token="fresh"), - _allow_response(), - ] - ) - guardrail: Final = _make_guardrail(handler) + exchanger: Final = StubTokenExchanger([_ok_exchange("short-lived", -1), _ok_exchange("fresh")]) + handler: Final = FakeHandler([_allow_response(), _allow_response()]) + guardrail: Final = _make_guardrail(handler, exchanger=exchanger) await _run(guardrail, _mcp_data()) await _run(guardrail, _mcp_data()) - token_calls: Final = [c for c in handler.calls if c.url == TOKEN_URL] - assert len(token_calls) == 2 - assert handler.calls[3].headers["Authorization"] == "Bearer fresh" + assert len(exchanger.calls) == 2 + assert handler.calls[1].headers["Authorization"] == "Bearer fresh" class TestEarlyPhasePassthrough: @@ -835,34 +916,27 @@ class TestRegistryDiscovery: class TestMalformedResponses: @pytest.mark.asyncio - async def test_obo_html_body_fail_open(self): - handler: Final = FakeHandler([_response(200, text="blocked by egress proxy")]) - guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open") + async def test_obo_upstream_unavailable_fail_open(self): + exchanger: Final = StubTokenExchanger([Error(CredError.of_upstream_unavailable("entra returned no token"))]) + handler: Final = FakeHandler([]) + guardrail: Final = _make_guardrail(handler, exchanger=exchanger, unreachable_fallback="fail_open") data: Final = _mcp_data() result: Final = await _run(guardrail, data) assert result is data assert _guardrail_info(data)["guardrail_response"]["verdict"] == "Unscanned" @pytest.mark.asyncio - async def test_obo_html_body_fail_closed(self): - handler: Final = FakeHandler([_response(200, text="outage")]) - guardrail: Final = _make_guardrail(handler) - with pytest.raises(HTTPException) as exc_info: - await _run(guardrail, _mcp_data()) - assert exc_info.value.status_code == 503 - assert "non-JSON" in exc_info.value.detail["message"] - - @pytest.mark.asyncio - async def test_obo_non_object_json_fail_closed(self): - handler: Final = FakeHandler([_response(200, ["not", "a", "dict"])]) - guardrail: Final = _make_guardrail(handler) + async def test_obo_upstream_unavailable_fail_closed(self): + exchanger: Final = StubTokenExchanger([Error(CredError.of_upstream_unavailable("entra returned no token"))]) + handler: Final = FakeHandler([]) + guardrail: Final = _make_guardrail(handler, exchanger=exchanger) with pytest.raises(HTTPException) as exc_info: await _run(guardrail, _mcp_data()) assert exc_info.value.status_code == 503 @pytest.mark.asyncio async def test_evaluate_html_body_fail_open(self): - handler: Final = FakeHandler([_token_response(), _response(200, text="waf page")]) + handler: Final = FakeHandler([_response(200, text="waf page")]) guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open") data: Final = _mcp_data() result: Final = await _run(guardrail, data) @@ -871,7 +945,7 @@ class TestMalformedResponses: @pytest.mark.asyncio async def test_evaluate_html_body_fail_closed(self): - handler: Final = FakeHandler([_token_response(), _response(200, text="waf page")]) + handler: Final = FakeHandler([_response(200, text="waf page")]) guardrail: Final = _make_guardrail(handler) with pytest.raises(HTTPException) as exc_info: await _run(guardrail, _mcp_data()) @@ -879,7 +953,7 @@ class TestMalformedResponses: @pytest.mark.asyncio async def test_evaluate_non_object_json_fail_closed(self): - handler: Final = FakeHandler([_token_response(), _response(200, "allowed")]) + handler: Final = FakeHandler([_response(200, "allowed")]) guardrail: Final = _make_guardrail(handler) with pytest.raises(HTTPException) as exc_info: await _run(guardrail, _mcp_data()) @@ -892,7 +966,7 @@ class TestMalformedResponses: ids=["missing", "null", "string-true", "int-one", "string-false"], ) async def test_evaluate_non_boolean_allowed_fail_closed(self, verdict: dict): - handler: Final = FakeHandler([_token_response(), _response(200, verdict)]) + handler: Final = FakeHandler([_response(200, verdict)]) guardrail: Final = _make_guardrail(handler) data: Final = _mcp_data() with pytest.raises(HTTPException) as exc_info: @@ -908,7 +982,7 @@ class TestMalformedResponses: ids=["missing", "null", "string-true", "int-one", "string-false"], ) async def test_evaluate_non_boolean_allowed_fail_open(self, verdict: dict): - handler: Final = FakeHandler([_token_response(), _response(200, verdict)]) + handler: Final = FakeHandler([_response(200, verdict)]) guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open") data: Final = _mcp_data() result: Final = await _run(guardrail, data) @@ -917,21 +991,21 @@ class TestMalformedResponses: assert _guardrail_info(data)["guardrail_status"] == "guardrail_failed_to_respond" @pytest.mark.asyncio - async def test_bad_expires_in_still_allows(self): - handler: Final = FakeHandler( - [_response(200, {"access_token": "tok-1", "expires_in": "soon"}), _allow_response()] - ) - guardrail: Final = _make_guardrail(handler) + async def test_obo_token_without_expiry_still_allows(self): + exchanger: Final = StubTokenExchanger([Ok(OAuthToken(access_token="tok-1", expires_at=None))]) + handler: Final = FakeHandler([_allow_response()]) + guardrail: Final = _make_guardrail(handler, exchanger=exchanger) data: Final = _mcp_data() result: Final = await _run(guardrail, data) assert result is data @pytest.mark.asyncio async def test_obo_litellm_timeout_fail_open(self): - handler: Final = FakeHandler( + exchanger: Final = StubTokenExchanger( [LitellmTimeout(message="Connection timed out", model="default-model-name", llm_provider="httpx")] ) - guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open") + handler: Final = FakeHandler([]) + guardrail: Final = _make_guardrail(handler, exchanger=exchanger, unreachable_fallback="fail_open") data: Final = _mcp_data() result: Final = await _run(guardrail, data) assert result is data @@ -940,28 +1014,18 @@ class TestMalformedResponses: class TestDeltaHardening: @pytest.mark.asyncio - async def test_non_string_access_token_fail_closed(self): - handler: Final = FakeHandler([_response(200, {"access_token": None, "expires_in": 3599})]) - guardrail: Final = _make_guardrail(handler) - with pytest.raises(HTTPException) as exc_info: - await _run(guardrail, _mcp_data()) - assert exc_info.value.status_code == 503 - assert "access_token" in exc_info.value.detail["message"] - - @pytest.mark.asyncio - async def test_numeric_string_expires_in_honored(self): - handler: Final = FakeHandler( - [_response(200, {"access_token": "tok-9", "expires_in": "120"}), _allow_response()] - ) - guardrail: Final = _make_guardrail(handler) + async def test_unexpired_exchange_result_is_reused(self): + exchanger: Final = StubTokenExchanger([_ok_exchange("tok-9", 120)]) + handler: Final = FakeHandler([_allow_response(), _allow_response()]) + guardrail: Final = _make_guardrail(handler, exchanger=exchanger) await _run(guardrail, _mcp_data()) - entries: Final = list(guardrail._obo_token_cache.values()) - assert len(entries) == 1 - assert entries[0][1] - time.time() < 200 + await _run(guardrail, _mcp_data()) + assert len(exchanger.calls) == 1 + assert handler.calls[1].headers["Authorization"] == "Bearer tok-9" @pytest.mark.asyncio async def test_evaluate_400_records_intervention(self): - handler: Final = FakeHandler([_token_response(), _response(400, text="bad request shape")]) + handler: Final = FakeHandler([_response(400, text="bad request shape")]) guardrail: Final = _make_guardrail(handler) data: Final = _mcp_data() with pytest.raises(HTTPException) as exc_info: @@ -975,25 +1039,19 @@ class TestDeltaHardening: class TestVeriaHardening: @pytest.mark.asyncio async def test_evaluate_401_evicts_cached_obo_token(self): - handler: Final = FakeHandler( - [ - _token_response(), - _response(401, text="token expired"), - _token_response(access_token="tok-2"), - _allow_response(), - ] - ) - guardrail: Final = _make_guardrail(handler) + exchanger: Final = StubTokenExchanger(_obo_ok() + _obo_ok("tok-2")) + handler: Final = FakeHandler([_response(401, text="token expired"), _allow_response()]) + guardrail: Final = _make_guardrail(handler, exchanger=exchanger) with pytest.raises(HTTPException): await _run(guardrail, _mcp_data()) result: Final = await _run(guardrail, _mcp_data()) assert result is not None - token_calls: Final = [c for c in handler.calls if c.url == TOKEN_URL] - assert len(token_calls) == 2 + assert exchanger.invalidations == [FAKE_ASSERTION] + assert len(exchanger.calls) == 2 @pytest.mark.asyncio async def test_evaluate_429_blocks_even_fail_open_as_throttled(self): - handler: Final = FakeHandler([_token_response(), _response(429, text="slow down")]) + handler: Final = FakeHandler([_response(429, text="slow down")]) guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open") data: Final = _mcp_data() with pytest.raises(HTTPException) as exc_info: @@ -1006,7 +1064,7 @@ class TestVeriaHardening: @pytest.mark.asyncio async def test_evaluate_500_is_unavailable(self): - handler: Final = FakeHandler([_token_response(), _response(500, text="oops")]) + handler: Final = FakeHandler([_response(500, text="oops")]) guardrail: Final = _make_guardrail(handler) with pytest.raises(HTTPException) as exc_info: await _run(guardrail, _mcp_data()) @@ -1014,51 +1072,23 @@ class TestVeriaHardening: assert "500" in exc_info.value.detail["message"] @pytest.mark.asyncio - async def test_token_endpoint_429_blocks_even_fail_open_as_throttled(self): - handler: Final = FakeHandler( - [_response(429, {"error": "temporarily_throttled", "error_description": "AADSTS90056"})] + async def test_token_endpoint_unauthorized_is_a_caller_401_even_fail_open(self): + exchanger: Final = StubTokenExchanger( + [Error(CredError.of_unauthorized("IdP rejected the subject token (HTTP 429)"))] ) - guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open") + handler: Final = FakeHandler([]) + guardrail: Final = _make_guardrail(handler, exchanger=exchanger, unreachable_fallback="fail_open") data: Final = _mcp_data() with pytest.raises(HTTPException) as exc_info: await _run(guardrail, data) - assert exc_info.value.status_code == 503 - assert "429" in exc_info.value.detail["message"] + assert exc_info.value.status_code == 401 info: Final = _guardrail_info(data) - assert info["guardrail_status"] == "guardrail_failed_to_respond" - assert info["guardrail_response"]["verdict"] == "Throttled" - - @pytest.mark.asyncio - async def test_token_endpoint_408_non_json_blocks_as_throttled(self): - handler: Final = FakeHandler([_response(408, text="Request Timeout")]) - guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open") - data: Final = _mcp_data() - with pytest.raises(HTTPException) as exc_info: - await _run(guardrail, data) - assert exc_info.value.status_code == 503 - assert _guardrail_info(data)["guardrail_response"]["verdict"] == "Throttled" - - @pytest.mark.asyncio - async def test_token_endpoint_4xx_html_stays_infra_fail_open(self): - handler: Final = FakeHandler([_response(403, text="waf block page")]) - guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open") - data: Final = _mcp_data() - result: Final = await _run(guardrail, data) - assert result is data - assert _guardrail_info(data)["guardrail_response"]["verdict"] == "Unscanned" - - @pytest.mark.asyncio - async def test_entra_200_missing_access_token_is_malformed(self): - handler: Final = FakeHandler([_response(200, {"token_type": "Bearer"})]) - guardrail: Final = _make_guardrail(handler) - with pytest.raises(HTTPException) as exc_info: - await _run(guardrail, _mcp_data()) - assert exc_info.value.status_code == 503 - assert "access_token" in exc_info.value.detail["message"] + assert info["guardrail_status"] == "guardrail_intervened" + assert info["guardrail_response"]["verdict"] == "Rejected" @pytest.mark.asyncio async def test_evaluate_5xx_fail_open_allows_unscanned_once(self): - handler: Final = FakeHandler([_token_response(), _response(502, text='{"error": "bad gateway"}')]) + handler: Final = FakeHandler([_response(502, text='{"error": "bad gateway"}')]) guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open") data: Final = _mcp_data() result: Final = await _run(guardrail, data) @@ -1099,7 +1129,7 @@ class TestFinalArgumentsEvaluated: @pytest.mark.asyncio @pytest.mark.parametrize("agent_365_first", [True, False], ids=["agent_365_then_masker", "masker_then_agent_365"]) async def test_agent_365_receives_the_arguments_sent_upstream(self, agent_365_first: bool): - handler: Final = FakeHandler([_token_response(), _allow_response()]) + handler: Final = FakeHandler([_allow_response()]) guardrail: Final = _make_guardrail(handler) masker: Final = _ArgumentMasker("arg-rewrite") registered: Final = (guardrail, masker) if agent_365_first else (masker, guardrail) @@ -1116,4 +1146,61 @@ class TestFinalArgumentsEvaluated: litellm.callbacks, callback, require_self=False ) assert result["modified_arguments"] == {"turn": "please [REWRITE_ME_REDACTED] now"} - assert handler.calls[1].json["arguments"] == {"turn": "please [REWRITE_ME_REDACTED] now"} + assert handler.calls[0].json["arguments"] == {"turn": "please [REWRITE_ME_REDACTED] now"} + + +class TestCallerSignIn: + """The guardrail is a CallerSignInProvider: it tells the MCP connect path which issuer and scope + the caller must sign in for before the first tool call, and only for servers that keep the + caller's Authorization with the gateway.""" + + def test_gated_server_advertises_entra_issuer_and_gateway_scope(self): + guardrail: Final = _make_guardrail(FakeHandler([])) + sign_in: Final = guardrail.caller_sign_in(_server(), None) + assert sign_in == CallerSignIn( + issuers=("https://login.microsoftonline.com/tenant-abc/v2.0",), + scopes=("api://client-xyz/access_as_user",), + ) + + def test_default_off_guardrail_does_not_gate(self): + guardrail: Final = _make_guardrail(FakeHandler([]), default_on=False) + assert guardrail.caller_sign_in(_server(), None) is None + + def test_server_that_fills_authorization_itself_is_not_gated(self): + guardrail: Final = _make_guardrail(FakeHandler([])) + assert guardrail.caller_sign_in(_server(auth_type=MCPAuth.oauth2), None) is None + assert guardrail.caller_sign_in(_server(extra_headers=["authorization"]), None) is None + + def test_opted_out_key_does_not_gate(self): + class _OptedOut(Agent365Guardrail): + def should_run_guardrail(self, data, event_type) -> bool: + return False + + guardrail: Final = _OptedOut( + guardrail_name="a365-off", + tenant_id="tenant-abc", + client_id="client-xyz", + client_secret="secret-123", + token_exchanger=StubTokenExchanger(), + default_on=True, + ) + assert guardrail.caller_sign_in(_server(), UserAPIKeyAuth(api_key="k", user_id="u-1")) is None + assert guardrail.caller_sign_in(_server(), None) is not None + + def test_obo_server_with_provider_advertises_both_issuers_and_scopes(self, monkeypatch): + monkeypatch.setenv("JWT_ISSUER", "https://jwt-idp.test") + guardrail: Final = _make_guardrail(FakeHandler([])) + litellm.logging_callback_manager.add_litellm_callback(guardrail) + try: + server: Final = _server(auth_type=MCPAuth.oauth2_token_exchange, scopes=["read"]) + sign_in: Final = caller_sign_in_for(server, None) + finally: + litellm.logging_callback_manager.remove_callback_from_list_by_object( + litellm.callbacks, guardrail, require_self=False + ) + assert sign_in is not None + assert sign_in.issuers == ( + "https://jwt-idp.test", + "https://login.microsoftonline.com/tenant-abc/v2.0", + ) + assert sign_in.scopes == ("read", "api://client-xyz/access_as_user") From df076ce24cb4d82b602f97ac08a01242543e0554 Mon Sep 17 00:00:00 2001 From: yucheng Date: Sat, 26 Sep 2026 10:01:42 +0000 Subject: [PATCH 02/51] fix(mcp): keep the exact-name fallback in the challenge resolver and reformat get_mcp_server_answering_to now falls back to get_mcp_server_by_name when no published prefix form matches, preserving the exact-name lookup the preemptive path had before the router-equivalent resolver. Applies ruff format to caller_sign_in.py and agent_365.py per the lint gate. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/_experimental/mcp_server/caller_sign_in.py | 4 +--- .../proxy/_experimental/mcp_server/mcp_server_manager.py | 5 +++-- .../guardrails/guardrail_hooks/agent_365/agent_365.py | 9 ++++++--- 3 files changed, 10 insertions(+), 8 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/caller_sign_in.py b/litellm/proxy/_experimental/mcp_server/caller_sign_in.py index 9d9b8b7297e..95476bd7c20 100644 --- a/litellm/proxy/_experimental/mcp_server/caller_sign_in.py +++ b/litellm/proxy/_experimental/mcp_server/caller_sign_in.py @@ -107,9 +107,7 @@ def caller_sign_in_for(server: MCPServer, user_api_key_auth: UserAPIKeyAuth | No contribution for contribution in ( *( - ( - CallerSignIn(issuers=jwt_auth_issuers(), scopes=tuple(server.scopes or ())), - ) + (CallerSignIn(issuers=jwt_auth_issuers(), scopes=tuple(server.scopes or ())),) if server.auth_type == MCPAuth.oauth2_token_exchange else () ), diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index c38d156dd46..1bb3cedbd02 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -7193,7 +7193,8 @@ class MCPServerManager: def get_mcp_server_answering_to(self, name: str, client_ip: str | None = None) -> MCPServer | None: """The server a scoped ``/mcp/{name}`` connect resolves to, matched the way the router matches - it: case-insensitive over server_id, name and every published prefix form.""" + it: case-insensitive over server_id, name and every published prefix form, then the exact + name lookup as the fallback.""" return next( ( server @@ -7201,7 +7202,7 @@ class MCPServerManager: if server_answers_to_name(server, name) ), None, - ) + ) or self.get_mcp_server_by_name(name, client_ip=client_ip) def get_filtered_registry(self, client_ip: str | None = None) -> dict[str, MCPServer]: """ 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 e5397a7269d..cc2df181574 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py +++ b/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py @@ -445,9 +445,12 @@ class Agent365Guardrail(CustomGuardrail): "user_api_key_team_metadata": user_api_key_auth.team_metadata, # pyright: ignore[reportUnknownMemberType] # UserAPIKeyAuth.team_metadata is a raw dict } } - if self.should_run_guardrail( # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # should_run_guardrail takes an untyped data dict - data=probe, event_type=GuardrailEventHooks.pre_mcp_call - ) is not True: + 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),), From 4be107e779e064713399c06ab2f2ee522f5aa433 Mon Sep 17 00:00:00 2001 From: yucheng Date: Sat, 26 Sep 2026 10:23:13 +0000 Subject: [PATCH 03/51] fix(mcp): flatten the sign-in merge without stacked comprehension clauses LIT014 budgets one for clause per comprehension; chain.from_iterable keeps the dedupe immutable Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/_experimental/mcp_server/caller_sign_in.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/caller_sign_in.py b/litellm/proxy/_experimental/mcp_server/caller_sign_in.py index 95476bd7c20..7b1d538df5e 100644 --- a/litellm/proxy/_experimental/mcp_server/caller_sign_in.py +++ b/litellm/proxy/_experimental/mcp_server/caller_sign_in.py @@ -14,6 +14,7 @@ never imports a concrete provider. from __future__ import annotations +import itertools from collections.abc import Mapping from dataclasses import dataclass from typing import TYPE_CHECKING, Final, Protocol, cast, runtime_checkable @@ -117,6 +118,6 @@ def caller_sign_in_for(server: MCPServer, user_api_key_auth: UserAPIKeyAuth | No ] if not contributions: return None - issuers: Final = tuple(dict.fromkeys(issuer for contribution in contributions for issuer in contribution.issuers)) - scopes: Final = tuple(dict.fromkeys(scope for contribution in contributions for scope in contribution.scopes)) + 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) From 692086a0ae46883e07aaeabd24f0ab75db4b8cac Mon Sep 17 00:00:00 2001 From: yucheng Date: Sat, 26 Sep 2026 23:21:02 +0000 Subject: [PATCH 04/51] fix(mcp): restore the raw bearer hook kwarg, exact-name-first resolution, and the base OBO challenge gate Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/mcp_server_manager.py | 20 ++++-- .../proxy/_experimental/mcp_server/server.py | 6 +- .../guardrail_hooks/agent_365/agent_365.py | 2 +- .../mcp_server/test_mcp_server.py | 44 ++++++++++++ .../mcp_server/test_mcp_server_manager.py | 71 +++++++++++++++++++ .../guardrail_hooks/test_agent_365.py | 20 ++++-- 6 files changed, 147 insertions(+), 16 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 1bb3cedbd02..30e9d7f833e 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -5971,8 +5971,14 @@ class MCPServerManager: if proxy_logging_obj is None: return hook_result - # Admission credentials are never handed to guardrails as the caller's assertion. - incoming_bearer_token: Final = self._extract_subject_token(None, raw_headers, user_api_key_auth) + inbound_authorization: Final = next( + (v for k, v in (raw_headers or {}).items() if isinstance(k, str) and k.lower() == "authorization"), + "", + ) + incoming_bearer_token: Final = ( + inbound_authorization[len("bearer ") :] if inbound_authorization.lower().startswith("bearer ") else None + ) + incoming_subject_token: Final = self._extract_subject_token(None, raw_headers, user_api_key_auth) pre_hook_kwargs: Final = { "guardrail_context": guardrail_context, @@ -5988,6 +5994,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, @@ -7192,17 +7199,16 @@ class MCPServerManager: return None def get_mcp_server_answering_to(self, name: str, client_ip: str | None = None) -> MCPServer | None: - """The server a scoped ``/mcp/{name}`` connect resolves to, matched the way the router matches - it: case-insensitive over server_id, name and every published prefix form, then the exact - name lookup as the fallback.""" - return next( + """The server a scoped ``/mcp/{name}`` connect resolves to: the alias-first exact lookup, then + the router's case-insensitive prefix match.""" + return self.get_mcp_server_by_name(name, client_ip=client_ip) or next( ( server for server in self.get_filtered_registry(client_ip).values() if server_answers_to_name(server, name) ), None, - ) or self.get_mcp_server_by_name(name, client_ip=client_ip) + ) def get_filtered_registry(self, client_ip: str | None = None) -> dict[str, MCPServer]: """ diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index b73d183db86..b9dcb79dafe 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -1715,8 +1715,8 @@ if MCP_AVAILABLE: continue # Caller sign-in: challenge at connect because a tool-call-time 401 is wrapped into a - # JSON-RPC error and the WWW-Authenticate header is lost. Non-OBO gates fire only on a - # single-server connect the key's grant admits. + # 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. sign_in = caller_sign_in_for(server, user_api_key_auth) if server is not None else None if ( server @@ -1726,7 +1726,7 @@ if MCP_AVAILABLE: ) is None and ( - server.auth_type == MCPAuth.oauth2_token_exchange + (server.auth_type == MCPAuth.oauth2_token_exchange and not oauth2_headers) or await _key_granted_single_server(server, mcp_servers, user_api_key_auth, client_ip) ) ): 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 cc2df181574..b33c61951d0 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py +++ b/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py @@ -195,7 +195,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, diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 8d2642a66f1..48d43519b6a 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -10270,6 +10270,50 @@ class TestOboPreflightScopedToAllowedServers: ) +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"}, + ) + + @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.""" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 97b455bfd39..8194d7ab5c5 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -6186,6 +6186,24 @@ 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-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 + def test_remove_server_drops_only_its_own_tool_mapping_rows(self): manager = self._manager_with_deepwiki_and_huggingface() @@ -6813,6 +6831,59 @@ 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 eyJ.x.y"}, + "eyJ.x.y", + "eyJ.x.y", + None, + id="idp-token-as-admission-stays-raw-bearer", + ), + 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", + ), + ], + ) + 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 async def test_check_tool_permission_for_key_team_allows_permitted_tool(self): """ diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py index cfd6745308d..47f1d9f5f0c 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py @@ -210,7 +210,7 @@ def _mcp_data(**overrides: Any) -> 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) @@ -699,17 +699,27 @@ 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_raw_bearer_without_subject_token_is_no_bearer(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=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_bearer_token="sk-litellm-virtual-key")) + await _run(guardrail, _mcp_data(incoming_subject_token="sk-litellm-virtual-key")) assert exc_info.value.status_code == 401 assert exchanger.calls == [] @@ -717,7 +727,7 @@ class TestUnreachableFallback: 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 @@ -871,7 +881,7 @@ class TestOboTokenCache: 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)) + await _run(guardrail, _mcp_data(incoming_subject_token=other_assertion)) assert len(exchanger.calls) == 2 assert handler.calls[1].headers["Authorization"] == "Bearer token-b" From 6a71a3b1a0fb7b9697b8c290490452c1f7fa8efd Mon Sep 17 00:00:00 2001 From: yucheng Date: Sun, 27 Sep 2026 01:12:38 +0000 Subject: [PATCH 05/51] test(mcp): add integration coverage for the caller sign-in gates Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp/test_mcp_caller_sign_in.py | 218 ++++++++++++++++++ 1 file changed, 218 insertions(+) create mode 100644 tests/integration/mcp/test_mcp_caller_sign_in.py 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..f80a0c086af --- /dev/null +++ b/tests/integration/mcp/test_mcp_caller_sign_in.py @@ -0,0 +1,218 @@ +import json +import uuid +from collections.abc import Mapping +from pathlib import Path +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]) -> httpx.Response: + return gateway.client.post( + path, + json={"jsonrpc": "2.0", "id": 1, "method": "initialize", "params": INITIALIZE}, + headers={"x-litellm-api-key": key, **ACCEPT, **headers}, + ) + + +def _sign_in_config(guardrail_params: dict[str, object], path: Path) -> 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}] + 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_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() == () + + +def test_agent_365_gated_server_challenges_at_connect_and_advertises_entra(gateway: Gateway, tmp_path: Path) -> None: + def nothing(request: Request) -> Reply: + return Reply(status=500) + + with wire_server(nothing) as api: + config: Final = _sign_in_config( + { + "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", + "api_base": api.url, + }, + 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="/.well-known/oauth-protected-resource/mcp/{alias}"' in authenticate + assert 'error="invalid_token"' in 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"] == [ + "https://login.microsoftonline.com/00000000-0000-0000-0000-000000000000/v2.0" + ] + assert document["scopes_supported"] == ["api://22222222-2222-2222-2222-222222222222/access_as_user"] + + refused: Final = _rpc(candidate, f"/mcp/{alias}", denied, {}) + assert refused.status_code == 403, refused.text + assert "www-authenticate" not in refused.headers + assert api.drain() == () From fcfd7b4f1409ae70bc0a4b8d1e76b79cbddfd2cc Mon Sep 17 00:00:00 2001 From: yucheng Date: Sun, 27 Sep 2026 02:17:03 +0000 Subject: [PATCH 06/51] fix(mcp): name the connected segment in sign-in challenge resource metadata Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/discoverable_endpoints.py | 5 +- .../outbound_credentials/adapter.py | 7 ++- .../proxy/_experimental/mcp_server/server.py | 4 +- .../mcp/test_mcp_caller_sign_in.py | 43 +++++++++++++++ .../mcp_server/test_discoverable_endpoints.py | 50 +++++++++++++++++ .../mcp_server/test_mcp_server.py | 55 +++++++++++++++++++ 6 files changed, 159 insertions(+), 5 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 6940e8d4f02..849c1feb524 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -531,7 +531,10 @@ def _resolve_mcp_server_by_name_or_id(lookup: str, client_ip: str | None) -> MCP 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) + by_id: Final = global_mcp_server_manager.get_mcp_server_by_id(lookup, client_ip=client_ip) + if by_id is not None: + return by_id + return global_mcp_server_manager.get_mcp_server_answering_to(lookup, client_ip=client_ip) def _resolve_oauth2_server_for_root_endpoints( diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/adapter.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/adapter.py index 42947e39530..a4b723b5d67 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/adapter.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/adapter.py @@ -298,7 +298,7 @@ def raise_public(error: CredError) -> NoReturn: assert_never(error.tag) -def oauth_protected_resource_path(root_path: str, server: MCPServer) -> str: +def oauth_protected_resource_path(root_path: str, server: MCPServer, *, connected_as: str | None = None) -> str: """The server's RFC 9728 Protected Resource Metadata path, the shared anchor of both challenges. ``root_path`` is the prefix the request was routed under, resolved by the caller (the imperative @@ -320,7 +320,7 @@ def oauth_protected_resource_path(root_path: str, server: MCPServer) -> str: challenge would then disagree on where the resource metadata lives. """ prefix: Final = "" if root_path == "/" else root_path - name: Final = server.alias or server.server_name or server.name or server.server_id + name: Final = connected_as or server.alias or server.server_name or server.name or server.server_id scalar_env: Final = os.getenv("SERVER_ROOT_PATH", "").rstrip("/") if not prefix or (scalar_env and prefix == scalar_env): return f"/.well-known/oauth-protected-resource{prefix}/mcp/{name}" @@ -348,6 +348,7 @@ def raise_token_exchange_challenge( *, root_path: str, claims: str | None = None, + connected_as: 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. @@ -366,7 +367,7 @@ def raise_token_exchange_challenge( two literals) and the base64 claims draw from a fixed alphabet, so nothing from the IdP body reaches the header unescaped. """ - resource_metadata: Final = oauth_protected_resource_path(root_path, server) + resource_metadata: Final = oauth_protected_resource_path(root_path, server, connected_as=connected_as) 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 = ( diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index b9dcb79dafe..7570ba36905 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -1737,7 +1737,9 @@ 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(), connected_as=server_name + ) # 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 diff --git a/tests/integration/mcp/test_mcp_caller_sign_in.py b/tests/integration/mcp/test_mcp_caller_sign_in.py index f80a0c086af..1d41be0d78c 100644 --- a/tests/integration/mcp/test_mcp_caller_sign_in.py +++ b/tests/integration/mcp/test_mcp_caller_sign_in.py @@ -216,3 +216,46 @@ def test_agent_365_gated_server_challenges_at_connect_and_advertises_entra(gatew assert refused.status_code == 403, refused.text assert "www-authenticate" not in refused.headers assert api.drain() == () + + +def test_challenge_and_prm_resolve_the_connected_case_variant(gateway: Gateway, tmp_path: Path) -> None: + def nothing(request: Request) -> Reply: + return Reply(status=500) + + with wire_server(nothing) as api: + config: Final = _sign_in_config( + { + "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", + "api_base": api.url, + }, + 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="/.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"] == [ + "https://login.microsoftonline.com/00000000-0000-0000-0000-000000000000/v2.0" + ] + assert document["scopes_supported"] == ["api://22222222-2222-2222-2222-222222222222/access_as_user"] + assert api.drain() == () diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 7349db67d5d..cb24b3a32b7 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -3743,6 +3743,56 @@ async def test_protected_resource_metadata_resolves_server_by_id_when_name_looku by_id.assert_called_once_with(server.server_id, client_ip=None) +@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(): from fastapi import Request diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 48d43519b6a..0c4c6ce1ab8 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -11063,11 +11063,22 @@ def _catalog_server() -> MCPServer: 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 @@ -11195,3 +11206,47 @@ class TestConnectChallengeResolver: 'error="invalid_token", ' 'error_description="Missing or invalid subject token; authenticate with the IdP and retry"' ) + + @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, + ) + 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}"') From 0099b96ddc62a40513455406d46420c9e7cb9ec4 Mon Sep 17 00:00:00 2001 From: yucheng Date: Sun, 27 Sep 2026 02:17:42 +0000 Subject: [PATCH 07/51] feat(mcp): challenge rejected caller sign-in subjects at connect Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/caller_sign_in.py | 65 +++++++++++++- .../proxy/_experimental/mcp_server/server.py | 34 +++++-- .../guardrail_hooks/agent_365/agent_365.py | 56 +++++++++++- .../mcp_server/test_caller_sign_in.py | 7 ++ .../mcp_server/test_mcp_server.py | 89 +++++++++++++++++++ .../guardrail_hooks/test_agent_365.py | 67 +++++++++++++- 6 files changed, 308 insertions(+), 10 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/caller_sign_in.py b/litellm/proxy/_experimental/mcp_server/caller_sign_in.py index 7b1d538df5e..13ef8edfb21 100644 --- a/litellm/proxy/_experimental/mcp_server/caller_sign_in.py +++ b/litellm/proxy/_experimental/mcp_server/caller_sign_in.py @@ -17,7 +17,7 @@ from __future__ import annotations import itertools from collections.abc import Mapping from dataclasses import dataclass -from typing import TYPE_CHECKING, Final, Protocol, cast, runtime_checkable +from typing import TYPE_CHECKING, Final, Protocol, assert_never, cast, runtime_checkable from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError @@ -38,6 +38,30 @@ class CallerSignIn: 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: @@ -46,6 +70,13 @@ class CallerSignInProvider(Protocol): 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") @@ -121,3 +152,35 @@ def caller_sign_in_for(server: MCPServer, user_api_key_auth: UserAPIKeyAuth | No 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, + connected_as: str | None, +) -> None: + """Run every provider's connect-time check against the subject token, so a bearer the IdP will + reject surfaces as a challenge here rather than a JSON-RPC error at the first tool call.""" + 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, + ) + + for provider in _providers(): + if provider.caller_sign_in(server, user_api_key_auth) is None: + continue + 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, connected_as=connected_as) + case Unavailable(detail=detail, 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/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 7570ba36905..a4ba9e4d4bb 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -1718,13 +1718,17 @@ if MCP_AVAILABLE: # 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. sign_in = caller_sign_in_for(server, user_api_key_auth) if server is not None else None + subject_token: Final = ( + operations.global_mcp_server_manager._extract_subject_token( # pyright: ignore[reportPrivateUsage] # the manager owns the subject/admission filter shared with the preflight + oauth2_headers, raw_headers, user_api_key_auth + ) + if server is not None + else None + ) if ( server and sign_in is not None - and operations.global_mcp_server_manager._extract_subject_token( # pyright: ignore[reportPrivateUsage] # the manager owns the subject/admission filter shared with the preflight - oauth2_headers, raw_headers, user_api_key_auth - ) - is None + and subject_token is None and ( (server.auth_type == MCPAuth.oauth2_token_exchange and not oauth2_headers) or await _key_granted_single_server(server, mcp_servers, user_api_key_auth, client_ip) @@ -1737,8 +1741,26 @@ if MCP_AVAILABLE: get_request_root_path, ) - raise_token_exchange_challenge( - server, root_path=get_request_root_path(), connected_as=server_name + raise_token_exchange_challenge(server, root_path=get_request_root_path(), connected_as=server_name) + if ( + server + and sign_in is not None + and subject_token is not None + and await _key_granted_single_server(server, mcp_servers, user_api_key_auth, client_ip) + ): + 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(), + connected_as=server_name, ) # Exchange-backed modes (token_exchange's OBO mint, id_jag's stored-assertion mint): run 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 b33c61951d0..8bbb4c54b03 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py +++ b/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py @@ -31,13 +31,21 @@ from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, ) -from litellm.proxy._experimental.mcp_server.caller_sign_in import CallerSignIn -from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Error, Ok +from litellm.proxy._experimental.mcp_server.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, ) from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchanger import TokenExchanger from litellm.proxy._experimental.mcp_server.outbound_credentials.types import ( + CredError, ServerSpec, TokenExchangeConfig, ) @@ -457,6 +465,50 @@ class Agent365Guardrail(CustomGuardrail): scopes=(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 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 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 SignedIn() + try: + exchange_result: Final = await self._exchange_caller_assertion(assertion) + except (httpx.HTTPError, LitellmTimeout, TimeoutError) as exc: + return Unavailable( + detail=f"the Entra token endpoint could not be reached ({type(exc).__name__})", + fail_open=self.unreachable_fallback == "fail_open", + ) + match exchange_result: + case Ok(_): + return SignedIn() + case Error(error): + match error.tag: + case "unauthorized": + return Rejected(detail=error.unauthorized.detail, claims=error.unauthorized.claims) + case "misconfigured": + return Unavailable( + detail=( + f"Entra rejected the gateway's own Agent 365 credentials ({error.misconfigured}); " + "check the guardrail's client_id, client_secret and resource_app_id" + ), + fail_open=self.unreachable_fallback == "fail_open", + ) + case _: + return Unavailable( + detail=f"the Entra token exchange failed ({error.summary})", + fail_open=self.unreachable_fallback == "fail_open", + ) + async def _post_allowing_error_status( self, url: str, diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_caller_sign_in.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_caller_sign_in.py index aa853696462..0e66500fad0 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_caller_sign_in.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_caller_sign_in.py @@ -7,7 +7,9 @@ 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, ) from litellm.proxy._types import UserAPIKeyAuth @@ -29,6 +31,11 @@ class _SignInGuardrail(CustomGuardrail): 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: diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 0c4c6ce1ab8..bfaa1767456 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -11250,3 +11250,92 @@ class TestConnectChallengeResolver: 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): + 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=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"}, + ) + + @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 + 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/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py index 47f1d9f5f0c..d64ec31574a 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py @@ -14,7 +14,13 @@ from litellm.exceptions import Timeout as LitellmTimeout from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.litellm_core_utils.secret_redaction import redact_string -from litellm.proxy._experimental.mcp_server.caller_sign_in import CallerSignIn, caller_sign_in_for +from litellm.proxy._experimental.mcp_server.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 ( @@ -1214,3 +1220,62 @@ class TestCallerSignIn: "https://login.microsoftonline.com/tenant-abc/v2.0", ) assert sign_in.scopes == ("read", "api://client-xyz/access_as_user") + + +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_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 + + @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_signs_in_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 verdict == SignedIn() + assert exchanger.calls == [] From e6a6e3ec359fb0e81a4fd223b593c29386883f8c Mon Sep 17 00:00:00 2001 From: yucheng Date: Sun, 27 Sep 2026 02:18:11 +0000 Subject: [PATCH 08/51] style(mcp): format the new connect sign-in preflight tests Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../proxy/_experimental/mcp_server/test_mcp_server.py | 10 +++++++--- 1 file changed, 7 insertions(+), 3 deletions(-) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index bfaa1767456..d642c5de2c8 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -11249,7 +11249,9 @@ class TestConnectChallengeResolver: ) 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}"') + assert authenticate.startswith( + f'Bearer resource_metadata="/.well-known/oauth-protected-resource/mcp/{route_name}"' + ) class TestConnectSignInPreflight: @@ -11289,7 +11291,8 @@ class TestConnectSignInPreflight: server = _catalog_server() guardrail = _CallerSignInGuardrail( - guardrail_name="sign-in-stub", preflight_result=Rejected("the Entra OBO exchange was rejected (AADSTS70002)") + guardrail_name="sign-in-stub", + preflight_result=Rejected("the Entra OBO exchange was rejected (AADSTS70002)"), ) litellm.logging_callback_manager.add_litellm_callback(guardrail) try: @@ -11312,7 +11315,8 @@ class TestConnectSignInPreflight: server = _catalog_server() guardrail = _CallerSignInGuardrail( - guardrail_name="sign-in-stub", preflight_result=Unavailable("the Entra token endpoint could not be reached", fail_open=False) + 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: From c2f85ca2e78d833f9175dcc5a0881d90a0930227 Mon Sep 17 00:00:00 2001 From: yucheng Date: Sun, 27 Sep 2026 02:52:46 +0000 Subject: [PATCH 09/51] fix(mcp): share the allowed lookup between the sign-in and exchange preflights Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/caller_sign_in.py | 3 ++- .../proxy/_experimental/mcp_server/server.py | 19 ++++++++++--------- 2 files changed, 12 insertions(+), 10 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/caller_sign_in.py b/litellm/proxy/_experimental/mcp_server/caller_sign_in.py index 13ef8edfb21..b6b5b728924 100644 --- a/litellm/proxy/_experimental/mcp_server/caller_sign_in.py +++ b/litellm/proxy/_experimental/mcp_server/caller_sign_in.py @@ -17,9 +17,10 @@ from __future__ import annotations import itertools from collections.abc import Mapping from dataclasses import dataclass -from typing import TYPE_CHECKING, Final, Protocol, assert_never, cast, runtime_checkable +from typing import TYPE_CHECKING, Final, Protocol, cast, runtime_checkable from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError +from typing_extensions import assert_never import litellm from litellm.integrations.custom_guardrail import CustomGuardrail diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index a4ba9e4d4bb..a42c2b54cec 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -1718,7 +1718,7 @@ if MCP_AVAILABLE: # 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. sign_in = caller_sign_in_for(server, user_api_key_auth) if server is not None else None - subject_token: Final = ( + subject_token = ( operations.global_mcp_server_manager._extract_subject_token( # pyright: ignore[reportPrivateUsage] # the manager owns the subject/admission filter shared with the preflight oauth2_headers, raw_headers, user_api_key_auth ) @@ -1742,11 +1742,18 @@ if MCP_AVAILABLE: ) raise_token_exchange_challenge(server, root_path=get_request_root_path(), connected_as=server_name) + 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 server and len(mcp_servers or []) == 1 + else [] + ) if ( server and sign_in is not None and subject_token is not None - and await _key_granted_single_server(server, mcp_servers, user_api_key_auth, client_ip) + and any(allowed.server_id == server.server_id for allowed in allowed_single) ): from litellm.proxy._experimental.mcp_server.caller_sign_in import ( # noqa: PLC0415 # lazy: provider discovery pulls the guardrail registry preflight_caller_sign_in, @@ -1773,13 +1780,7 @@ if MCP_AVAILABLE: 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 - ) - ) + and server.server_id in frozenset(allowed.server_id for allowed in allowed_single) ): await operations.global_mcp_server_manager.preflight_token_exchange( server=server, From 072898ab1958c8f88ebaf8750ba46c24fb3cf165 Mon Sep 17 00:00:00 2001 From: yucheng Date: Wed, 30 Sep 2026 02:10:08 +0000 Subject: [PATCH 10/51] fix(mcp): satisfy type discipline gate and the merged input-schema key Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../_experimental/mcp_server/caller_sign_in.py | 4 ++-- .../mcp_server/discoverable_endpoints.py | 4 ++-- litellm/proxy/_experimental/mcp_server/server.py | 6 +++--- .../guardrail_hooks/agent_365/agent_365.py | 6 ++++-- .../guardrails/guardrail_hooks/test_agent_365.py | 16 +++++++--------- 5 files changed, 18 insertions(+), 18 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/caller_sign_in.py b/litellm/proxy/_experimental/mcp_server/caller_sign_in.py index b6b5b728924..e11c313f94a 100644 --- a/litellm/proxy/_experimental/mcp_server/caller_sign_in.py +++ b/litellm/proxy/_experimental/mcp_server/caller_sign_in.py @@ -136,7 +136,7 @@ def caller_sign_in_for(server: MCPServer, user_api_key_auth: UserAPIKeyAuth | No """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 = [ + contributions: Final = tuple( contribution for contribution in ( *( @@ -147,7 +147,7 @@ def caller_sign_in_for(server: MCPServer, user_api_key_auth: UserAPIKeyAuth | No *(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))) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 849c1feb524..5c4894c3bd6 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -2625,9 +2625,9 @@ def _caller_sign_in_protected_resource_response( if sign_in is None or not sign_in.issuers: return None return { - "authorization_servers": list(sign_in.issuers), + "authorization_servers": sign_in.issuers, "resource": resource_url, - "scopes_supported": list(sign_in.scopes), + "scopes_supported": sign_in.scopes, } diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index a42c2b54cec..9763c20cd4d 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -1589,7 +1589,7 @@ if MCP_AVAILABLE: ) -> bool: """Sign-in challenges are issued only on a single-server connect the key's grant admits, so a key without access gets the grant's 403 instead of a sign-in it could not use.""" - if len(mcp_servers or []) != 1: + if len(mcp_servers or ()) != 1: return False allowed: Final = await operations._get_allowed_mcp_servers( user_api_key_auth=user_api_key_auth, mcp_servers=mcp_servers, client_ip=client_ip @@ -1746,8 +1746,8 @@ if MCP_AVAILABLE: 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 len(mcp_servers or []) == 1 - else [] + if server and len(mcp_servers or ()) == 1 + else () ) if ( server 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 8bbb4c54b03..7208041f4da 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py +++ b/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py @@ -447,8 +447,10 @@ class Agent365Guardrail(CustomGuardrail): if not (self.default_on and server.keeps_caller_authorization): return None if user_api_key_auth is not None: - probe: Final[dict[str, Mapping[str, object]]] = { # pyright: ignore[reportUnknownVariableType] # UserAPIKeyAuth metadata dicts are untyped - "metadata": { + probe: Final[ + dict[str, Mapping[str, object]] + ] = { # mutable-ok: should_run_guardrail takes a mutable data dict # pyright: ignore[reportUnknownVariableType] # UserAPIKeyAuth metadata dicts are untyped + "metadata": { # mutable-ok: should_run_guardrail takes a mutable data dict "user_api_key_metadata": user_api_key_auth.metadata, # pyright: ignore[reportUnknownMemberType] # UserAPIKeyAuth.metadata is a raw dict "user_api_key_team_metadata": user_api_key_auth.team_metadata, # pyright: ignore[reportUnknownMemberType] # UserAPIKeyAuth.team_metadata is a raw dict } diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py index d64ec31574a..3d4fab45322 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py @@ -193,9 +193,7 @@ def _make_guardrail( ) -def _default_fallback_guardrail( - handler: FakeHandler, exchanger: StubTokenExchanger | None = None -) -> Agent365Guardrail: +def _default_fallback_guardrail(handler: FakeHandler, exchanger: StubTokenExchanger | None = None) -> Agent365Guardrail: return _make_guardrail(handler, exchanger=exchanger) @@ -416,7 +414,7 @@ class TestAllowFlow: 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_tool_input_schema=schema)) + await _run(guardrail, _mcp_data(mcp_tool_description="Send an email", mcp_input_schema=schema)) assert handler.calls[0].json["tool"] == { "name": "send_email", "description": "Send an email", @@ -431,7 +429,7 @@ class TestAllowFlow: async def test_evaluate_payload_omits_missing_or_malformed_tool_metadata(self, description, schema): handler: Final = FakeHandler([_allow_response()]) guardrail: Final = _make_guardrail(handler) - await _run(guardrail, _mcp_data(mcp_tool_description=description, mcp_tool_input_schema=schema)) + await _run(guardrail, _mcp_data(mcp_tool_description=description, mcp_input_schema=schema)) assert handler.calls[0].json["tool"] == {"name": "send_email"} @pytest.mark.asyncio @@ -617,9 +615,7 @@ class TestFailOpenOptIn: @pytest.mark.asyncio @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) - ) + 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()) @@ -846,7 +842,9 @@ class TestUnreachableFallback: @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"))]) + 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() From 4236a43aa7b27531cf7d6789673a568b9c3f7533 Mon Sep 17 00:00:00 2001 From: yucheng Date: Wed, 30 Sep 2026 01:57:39 +0000 Subject: [PATCH 11/51] feat(mcp): challenge opaque caller bearers at connect and move Agent 365 sign-in onto the fixed production constants Rework the Agent 365 sign-in provider for the guardrail shape 43189 landed on main: the OBO scope and resource come from the fixed AGENT_365_PROD_* constants instead of the removed resource_app_id/api_base fields, and the exchange runs through the shared TokenExchanger so the connect preflight and the tool call reuse one cached token per caller assertion. A present but non-JWS bearer is now rejected in preflight_caller_sign_in, so the connect answers 401 with the RFC 9728 challenge instead of letting the call reach tools/call and lose WWW-Authenticate in the JSON-RPC error. The OBO-only tool-call challenge in operations.py stays narrowed to token_exchange servers. Immutable rewrites (tuple, MappingProxyType, explicit None checks) keep the LIT002 total within the budget without a mutable-ok Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/mcp_server_manager.py | 5 +- .../_experimental/mcp_server/operations.py | 10 +- .../proxy/_experimental/mcp_server/server.py | 4 +- .../guardrail_hooks/agent_365/agent_365.py | 23 ++- litellm/proxy/utils.py | 1 + .../mcp/test_mcp_agent_365_guardrail.py | 10 +- .../mcp/test_mcp_caller_sign_in.py | 140 ++++++++---------- .../mcp_server/test_discoverable_endpoints.py | 16 +- .../guardrail_hooks/test_agent_365.py | 56 +++++-- .../utils/proxy_logging/test_mcp_bridging.py | 5 + 10 files changed, 146 insertions(+), 124 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 30e9d7f833e..287913d8eb7 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -5971,10 +5971,7 @@ class MCPServerManager: if proxy_logging_obj is None: return hook_result - inbound_authorization: Final = next( - (v for k, v in (raw_headers or {}).items() if isinstance(k, str) and k.lower() == "authorization"), - "", - ) + 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 ) diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index abc7c52f5a7..50101975dca 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -1766,13 +1766,11 @@ 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. """ - from litellm.proxy._experimental.mcp_server.caller_sign_in import ( - caller_sign_in_for, # noqa: PLC0415 # lazy: caller_sign_in pulls the proxy graph - ) - - if server is None or caller_sign_in_for(server, user_api_key_auth) is None: + if server is None or server.auth_type != MCPAuth.oauth2_token_exchange: return if requested_server is not None and requested_server.server_id != server.server_id: return diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 9763c20cd4d..cd1b0f9ddb6 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -1589,7 +1589,7 @@ if MCP_AVAILABLE: ) -> bool: """Sign-in challenges are issued only on a single-server connect the key's grant admits, so a key without access gets the grant's 403 instead of a sign-in it could not use.""" - if len(mcp_servers or ()) != 1: + if mcp_servers is None or len(mcp_servers) != 1: return False allowed: Final = await operations._get_allowed_mcp_servers( user_api_key_auth=user_api_key_auth, mcp_servers=mcp_servers, client_ip=client_ip @@ -1746,7 +1746,7 @@ if MCP_AVAILABLE: 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 len(mcp_servers or ()) == 1 + if server and mcp_servers is not None and len(mcp_servers) == 1 else () ) if ( 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 7208041f4da..9ffd1111d62 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py +++ b/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py @@ -100,6 +100,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) @@ -447,10 +456,8 @@ class Agent365Guardrail(CustomGuardrail): if not (self.default_on and server.keeps_caller_authorization): return None if user_api_key_auth is not None: - probe: Final[ - dict[str, Mapping[str, object]] - ] = { # mutable-ok: should_run_guardrail takes a mutable data dict # pyright: ignore[reportUnknownVariableType] # UserAPIKeyAuth metadata dicts are untyped - "metadata": { # mutable-ok: should_run_guardrail takes a mutable data dict + probe: Final[_AdmissionProbe] = { + "metadata": { "user_api_key_metadata": user_api_key_auth.metadata, # pyright: ignore[reportUnknownMemberType] # UserAPIKeyAuth.metadata is a raw dict "user_api_key_team_metadata": user_api_key_auth.team_metadata, # pyright: ignore[reportUnknownMemberType] # UserAPIKeyAuth.team_metadata is a raw dict } @@ -477,12 +484,12 @@ class Agent365Guardrail(CustomGuardrail): 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 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.""" + """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 SignedIn() + return Rejected(detail="the caller's bearer is not an Entra token; sign in with Entra and retry") try: exchange_result: Final = await self._exchange_caller_assertion(assertion) except (httpx.HTTPError, LitellmTimeout, TimeoutError) as exc: diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 8a325eba699..8406d021ad3 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -1519,6 +1519,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/tests/integration/mcp/test_mcp_agent_365_guardrail.py b/tests/integration/mcp/test_mcp_agent_365_guardrail.py index 5842d3ce4a7..e636cd45c44 100644 --- a/tests/integration/mcp/test_mcp_agent_365_guardrail.py +++ b/tests/integration/mcp/test_mcp_agent_365_guardrail.py @@ -136,13 +136,17 @@ def test_a_missing_or_malformed_caller_bearer_blocks_on_every_entry_point_whatev assert f"{rig.alias}-add" in rig.caller().list_tools().tools, "the catalog needs only the virtual key" 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, outcome in (("without a bearer", missing), ("opaque bearer", malformed)): + assert outcome.error is not None, f"{entry} {label}: {outcome.raw}" + if entry == "server_mcp": + assert outcome.status == 401, f"{entry} {label} skips the connect sign-in challenge: {outcome.raw}" + else: + assert 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 index 1d41be0d78c..9e940470031 100644 --- a/tests/integration/mcp/test_mcp_caller_sign_in.py +++ b/tests/integration/mcp/test_mcp_caller_sign_in.py @@ -171,91 +171,75 @@ def test_jwt_signer_verifies_the_bearer_that_admitted_the_call(gateway: Gateway, 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: - def nothing(request: Request) -> Reply: - return Reply(status=500) + 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"]}) - with wire_server(nothing) as api: - config: Final = _sign_in_config( - { - "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", - "api_base": api.url, - }, - tmp_path / "agent365.yaml", + 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="/.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="/.well-known/oauth-protected-resource/mcp/{alias}"' in opaque.headers.get( + "www-authenticate", "" ) - 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="/.well-known/oauth-protected-resource/mcp/{alias}"' in authenticate - assert 'error="invalid_token"' in 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] - 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"] == [ - "https://login.microsoftonline.com/00000000-0000-0000-0000-000000000000/v2.0" - ] - assert document["scopes_supported"] == ["api://22222222-2222-2222-2222-222222222222/access_as_user"] - - refused: Final = _rpc(candidate, f"/mcp/{alias}", denied, {}) - assert refused.status_code == 403, refused.text - assert "www-authenticate" not in refused.headers - assert api.drain() == () + 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_challenge_and_prm_resolve_the_connected_case_variant(gateway: Gateway, tmp_path: Path) -> None: - def nothing(request: Request) -> Reply: - return Reply(status=500) + 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() - with wire_server(nothing) as api: - config: Final = _sign_in_config( - { - "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", - "api_base": api.url, - }, - 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="/.well-known/oauth-protected-resource/mcp/{connected_as}"' in authenticate + assert 'error="invalid_token"' in authenticate - 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="/.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"] == [ - "https://login.microsoftonline.com/00000000-0000-0000-0000-000000000000/v2.0" - ] - assert document["scopes_supported"] == ["api://22222222-2222-2222-2222-222222222222/access_as_user"] - assert api.drain() == () + 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/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index cb24b3a32b7..361d78059ad 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -3787,9 +3787,9 @@ async def test_protected_resource_metadata_resolves_the_connected_case_variant() use_standard_pattern=True, ) - assert result["authorization_servers"] == [ - "https://login.microsoftonline.com/00000000-0000-0000-0000-000000000000/v2.0" - ] + assert result["authorization_servers"] == ( + "https://login.microsoftonline.com/00000000-0000-0000-0000-000000000000/v2.0", + ) assert result["resource"] == "https://llm.example.com/mcp/CATALOG" @@ -7367,21 +7367,21 @@ def test_caller_sign_in_protected_resource_response_names_jwt_issuers(): with patch(_PATCH_ISSUERS, return_value=["https://idp.example.com"]): 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_caller_sign_in_protected_resource_response_scopes_default_empty(): - """A scopeless OBO server reports scopes_supported as [] rather than None.""" + """A scopeless OBO server reports scopes_supported as an empty array rather than None.""" from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( _caller_sign_in_protected_resource_response, ) with patch(_PATCH_ISSUERS, return_value=["https://idp.example.com"]): response = _caller_sign_in_protected_resource_response(_obo_server(scopes=None), _OBO_RESOURCE) - assert response["scopes_supported"] == [] + assert response["scopes_supported"] == () def test_caller_sign_in_protected_resource_response_falls_back_when_no_issuer(): @@ -7440,7 +7440,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/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py index 3d4fab45322..a43ffd34145 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py @@ -1135,6 +1135,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 @@ -1185,20 +1213,17 @@ class TestCallerSignIn: assert guardrail.caller_sign_in(_server(auth_type=MCPAuth.oauth2), None) is None assert guardrail.caller_sign_in(_server(extra_headers=["authorization"]), None) is None - def test_opted_out_key_does_not_gate(self): - class _OptedOut(Agent365Guardrail): - def should_run_guardrail(self, data, event_type) -> bool: - return False - - guardrail: Final = _OptedOut( - guardrail_name="a365-off", - tenant_id="tenant-abc", - client_id="client-xyz", - client_secret="secret-123", - token_exchanger=StubTokenExchanger(), - default_on=True, + 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]} ) - assert guardrail.caller_sign_in(_server(), UserAPIKeyAuth(api_key="k", user_id="u-1")) is None + 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_obo_server_with_provider_advertises_both_issuers_and_scopes(self, monkeypatch): @@ -1269,11 +1294,12 @@ class TestPreflightCallerSignIn: assert verdict.fail_open is True @pytest.mark.asyncio - async def test_non_assertion_subject_signs_in_without_exchanging(self): + 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 verdict == SignedIn() + assert isinstance(verdict, Rejected) + assert verdict.claims is None assert exchanger.calls == [] diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_mcp_bridging.py b/tests/test_litellm/proxy/utils/proxy_logging/test_mcp_bridging.py index 25be3b5de6b..c7fb0848f76 100644 --- a/tests/test_litellm/proxy/utils/proxy_logging/test_mcp_bridging.py +++ b/tests/test_litellm/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, } From 1bddacfb68baaff8888675333717c2c627e2c3c4 Mon Sep 17 00:00:00 2001 From: yucheng Date: Wed, 30 Sep 2026 02:25:58 +0000 Subject: [PATCH 12/51] fix(guardrails): stop pointing admins at the removed resource_app_id in the connect-time 503 Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../proxy/guardrails/guardrail_hooks/agent_365/agent_365.py | 2 +- .../proxy/guardrails/guardrail_hooks/test_agent_365.py | 3 +++ 2 files changed, 4 insertions(+), 1 deletion(-) 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 9ffd1111d62..6e5849a92a4 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py +++ b/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py @@ -508,7 +508,7 @@ class Agent365Guardrail(CustomGuardrail): return Unavailable( detail=( f"Entra rejected the gateway's own Agent 365 credentials ({error.misconfigured}); " - "check the guardrail's client_id, client_secret and resource_app_id" + "check the guardrail's client_id and client_secret" ), fail_open=self.unreachable_fallback == "fail_open", ) diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py index a43ffd34145..fd2f14471ce 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py @@ -1280,6 +1280,9 @@ class TestPreflightCallerSignIn: 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_endpoint_unreachable_fail_open_is_unavailable(self): From f00a2c1c185d3fddd115f5385fd30323b630c2d4 Mon Sep 17 00:00:00 2001 From: yucheng Date: Wed, 30 Sep 2026 02:36:04 +0000 Subject: [PATCH 13/51] fix(mcp): run the single-server admission lookup once for the challenge, sign-in preflight and exchange Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../proxy/_experimental/mcp_server/server.py | 60 ++++++------------- 1 file changed, 18 insertions(+), 42 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index cd1b0f9ddb6..470fc6c2308 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -1581,21 +1581,6 @@ if MCP_AVAILABLE: ) return user_api_key_auth.model_copy(update={"object_permission": updated_op, "mcp_toolset_id": toolset_id}) - async def _key_granted_single_server( - server: MCPServer, - mcp_servers: Sequence[str] | None, - user_api_key_auth: UserAPIKeyAuth | None, - client_ip: str | None, - ) -> bool: - """Sign-in challenges are issued only on a single-server connect the key's grant admits, so a key - without access gets the grant's 403 instead of a sign-in it could not use.""" - if mcp_servers is None or len(mcp_servers) != 1: - return False - allowed: Final = await operations._get_allowed_mcp_servers( - user_api_key_auth=user_api_key_auth, mcp_servers=mcp_servers, client_ip=client_ip - ) - return any(granted.server_id == server.server_id for granted in allowed) - async def _raise_preemptive_401_for_unauthenticated_servers( scope: Scope, mcp_servers: list[str] | None, @@ -1716,7 +1701,9 @@ if MCP_AVAILABLE: # 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. + # 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 below 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 subject_token = ( operations.global_mcp_server_manager._extract_subject_token( # pyright: ignore[reportPrivateUsage] # the manager owns the subject/admission filter shared with the preflight @@ -1725,15 +1712,20 @@ if MCP_AVAILABLE: if server is not None else None ) - if ( - server - and sign_in is not None - and subject_token is None - and ( - (server.auth_type == MCPAuth.oauth2_token_exchange and not oauth2_headers) - or await _key_granted_single_server(server, mcp_servers, user_api_key_auth, client_ip) + obo_without_subject = ( + server is not None and server.auth_type == MCPAuth.oauth2_token_exchange and not oauth2_headers + ) + 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 server and not obo_without_subject and mcp_servers is not None and len(mcp_servers) == 1 + else () + ) + granted_single = server is not None and any( + allowed.server_id == server.server_id for allowed in allowed_single + ) + 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, ) @@ -1742,19 +1734,7 @@ if MCP_AVAILABLE: ) raise_token_exchange_challenge(server, root_path=get_request_root_path(), connected_as=server_name) - 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 server and mcp_servers is not None and len(mcp_servers) == 1 - else () - ) - if ( - server - and sign_in is not None - and subject_token is not None - and any(allowed.server_id == server.server_id for allowed in allowed_single) - ): + 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, ) @@ -1777,11 +1757,7 @@ 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 allowed_single) - ): + if server and granted_single: await operations.global_mcp_server_manager.preflight_token_exchange( server=server, oauth2_headers=oauth2_headers, From be5a5ccf81d1ec6a11ebed813f3dc242c27c2ba8 Mon Sep 17 00:00:00 2001 From: yucheng Date: Wed, 30 Sep 2026 08:39:05 +0000 Subject: [PATCH 14/51] fix(guardrails): return explicitly from every Agent 365 preflight exchange outcome Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../guardrail_hooks/agent_365/agent_365.py | 39 +++++++++---------- 1 file changed, 19 insertions(+), 20 deletions(-) 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 6e5849a92a4..ff06f09114d 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py +++ b/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py @@ -497,26 +497,25 @@ class Agent365Guardrail(CustomGuardrail): detail=f"the Entra token endpoint could not be reached ({type(exc).__name__})", fail_open=self.unreachable_fallback == "fail_open", ) - match exchange_result: - case Ok(_): - return SignedIn() - case Error(error): - match error.tag: - case "unauthorized": - return Rejected(detail=error.unauthorized.detail, claims=error.unauthorized.claims) - case "misconfigured": - return Unavailable( - detail=( - f"Entra rejected the gateway's own Agent 365 credentials ({error.misconfigured}); " - "check the guardrail's client_id and client_secret" - ), - fail_open=self.unreachable_fallback == "fail_open", - ) - case _: - return Unavailable( - detail=f"the Entra token exchange failed ({error.summary})", - fail_open=self.unreachable_fallback == "fail_open", - ) + if isinstance(exchange_result, Ok): + return SignedIn() + error: Final = exchange_result.error + match error.tag: + case "unauthorized": + return Rejected(detail=error.unauthorized.detail, claims=error.unauthorized.claims) + case "misconfigured": + return Unavailable( + detail=( + f"Entra rejected the gateway's own Agent 365 credentials ({error.misconfigured}); " + "check the guardrail's client_id and client_secret" + ), + fail_open=self.unreachable_fallback == "fail_open", + ) + case _: + return Unavailable( + detail=f"the Entra token exchange failed ({error.summary})", + fail_open=self.unreachable_fallback == "fail_open", + ) async def _post_allowing_error_status( self, From 06e5d9f7c5da4b8b1338eb3532223addead298b5 Mon Sep 17 00:00:00 2001 From: yucheng Date: Wed, 30 Sep 2026 09:16:03 +0000 Subject: [PATCH 15/51] fix(guardrails): return the Entra exchange fallback verdict from one path so CodeQL sees no implicit None Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../guardrail_hooks/agent_365/agent_365.py | 25 +++++++------------ .../guardrail_hooks/test_agent_365.py | 11 ++++++++ 2 files changed, 20 insertions(+), 16 deletions(-) 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 ff06f09114d..56cc5763e30 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py +++ b/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py @@ -500,22 +500,15 @@ class Agent365Guardrail(CustomGuardrail): if isinstance(exchange_result, Ok): return SignedIn() error: Final = exchange_result.error - match error.tag: - case "unauthorized": - return Rejected(detail=error.unauthorized.detail, claims=error.unauthorized.claims) - case "misconfigured": - return Unavailable( - detail=( - f"Entra rejected the gateway's own Agent 365 credentials ({error.misconfigured}); " - "check the guardrail's client_id and client_secret" - ), - fail_open=self.unreachable_fallback == "fail_open", - ) - case _: - return Unavailable( - detail=f"the Entra token exchange failed ({error.summary})", - fail_open=self.unreachable_fallback == "fail_open", - ) + if error.tag == "unauthorized": + return Rejected(detail=error.unauthorized.detail, claims=error.unauthorized.claims) + detail: Final = ( + f"Entra rejected the gateway's own Agent 365 credentials ({error.misconfigured}); " + "check the guardrail's client_id and client_secret" + if error.tag == "misconfigured" + else f"the Entra token exchange failed ({error.summary})" + ) + return Unavailable(detail=detail, fail_open=self.unreachable_fallback == "fail_open") async def _post_allowing_error_status( self, diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py index fd2f14471ce..ba90460d573 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py @@ -1284,6 +1284,17 @@ class TestPreflightCallerSignIn: "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( From eeadf7bdc1f61d036a491a1d1aa2471287dc646e Mon Sep 17 00:00:00 2001 From: yucheng Date: Wed, 30 Sep 2026 09:48:28 +0000 Subject: [PATCH 16/51] fix(mcp): name the connected route in OBO rejection challenges and test the initialized Agent 365 guardrail Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/mcp_server_manager.py | 8 ++++-- .../proxy/_experimental/mcp_server/server.py | 1 + .../guardrail_hooks/agent_365/__init__.py | 3 +++ .../mcp_server/test_mcp_server_manager.py | 25 +++++++++++++++++++ .../guardrail_hooks/test_agent_365.py | 20 ++++++++++++--- 5 files changed, 52 insertions(+), 5 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 287913d8eb7..bdd3c983717 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -4211,6 +4211,7 @@ class MCPServerManager: oauth2_headers: dict[str, str] | None, user_api_key_auth: UserAPIKeyAuth | None, raw_headers: Mapping[str, str] | None = None, + connected_as: str | None = None, ) -> None: """Mint an exchange-backed server's upstream credential at the transport edge. @@ -4247,14 +4248,16 @@ class MCPServerManager: ) if subject_token is None and caller_sign_in_for(server, user_api_key_auth) is not None: - raise_token_exchange_challenge(server, root_path=get_request_root_path()) + raise_token_exchange_challenge(server, root_path=get_request_root_path(), connected_as=connected_as) 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(), connected_as=connected_as + ) match await self._cred_provider.resolve_credentials(to_subject(user_api_key_auth, subject_token), spec): case Ok(_): return @@ -4264,6 +4267,7 @@ class MCPServerManager: resolved_server, root_path=get_request_root_path(), claims=err.unauthorized.claims, + connected_as=connected_as, ) raise_public(err) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 470fc6c2308..988f77ea98b 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -1763,6 +1763,7 @@ if MCP_AVAILABLE: oauth2_headers=oauth2_headers, user_api_key_auth=user_api_key_auth, raw_headers=raw_headers, + connected_as=server_name, ) # Pass-through OAuth: when the admin has opted a server into diff --git a/litellm/proxy/guardrails/guardrail_hooks/agent_365/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/agent_365/__init__.py index 2a8c6479ae6..dc0257f2017 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/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 8194d7ab5c5..14d71f98900 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -2930,6 +2930,31 @@ 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, + connected_as=server.server_id, + ) + headers = exc_info.value.headers or {} + www_authenticate = headers.get("WWW-Authenticate") or headers.get("www-authenticate") or "" + assert f"/.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) diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py index ba90460d573..c57a3b42975 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py @@ -194,7 +194,16 @@ def _make_guardrail( def _default_fallback_guardrail(handler: FakeHandler, exchanger: StubTokenExchanger | None = None) -> Agent365Guardrail: - return _make_guardrail(handler, exchanger=exchanger) + 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 _server(**overrides: Any) -> MCPServer: @@ -320,11 +329,16 @@ class TestInitializeGuardrail: agent_id="yaml-agent", ) handler: Final = FakeHandler([_allow_response()]) + exchanger: Final = StubTokenExchanger(_obo_ok()) with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): - 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 - guardrail: Final = _make_guardrail(handler) await _run(guardrail, _mcp_data()) + _, 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" From bc07570af83e3b81e39d12424ca8cfa88731a87a Mon Sep 17 00:00:00 2001 From: yucheng Date: Wed, 30 Sep 2026 10:02:00 +0000 Subject: [PATCH 17/51] test(mcp): expect the connected route name in the connect-time OBO preflight Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../proxy/_experimental/mcp_server/test_mcp_server.py | 1 + 1 file changed, 1 insertion(+) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index d642c5de2c8..c7dff4fa82e 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -10267,6 +10267,7 @@ class TestOboPreflightScopedToAllowedServers: "x-litellm-api-key": key.api_key, "authorization": self.SUBJECT_HEADERS["Authorization"], }, + connected_as=requested.alias, ) From 674356d3d841816b36784d09a88e8c68fa4c7b52 Mon Sep 17 00:00:00 2001 From: yucheng Date: Wed, 30 Sep 2026 10:18:19 +0000 Subject: [PATCH 18/51] fix(mcp): resolve case-variant scoped connects with the same alias-first priority as the exact name Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/mcp_server_manager.py | 16 +++++++++++++--- .../mcp_server/test_mcp_server_manager.py | 2 ++ 2 files changed, 15 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index bdd3c983717..0611b153de5 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -7200,9 +7200,19 @@ class MCPServerManager: return None def get_mcp_server_answering_to(self, name: str, client_ip: str | None = None) -> MCPServer | None: - """The server a scoped ``/mcp/{name}`` connect resolves to: the alias-first exact lookup, then - the router's case-insensitive prefix match.""" - return self.get_mcp_server_by_name(name, client_ip=client_ip) or next( + """The server a scoped ``/mcp/{name}`` connect resolves to: alias, then server_name, then name, each + case-insensitive so ``/mcp/GH`` and ``/mcp/gh`` agree, then the router's prefix match.""" + 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() diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 14d71f98900..1a4dd2cb93f 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -6225,6 +6225,8 @@ class TestMCPServerManager: ) 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 From 23eb6bd8a58481192c62c9d6e41d87885b43e8ec Mon Sep 17 00:00:00 2001 From: yucheng Date: Wed, 30 Sep 2026 10:32:13 +0000 Subject: [PATCH 19/51] test(mcp): patch the scoped-connect resolver the preemptive challenge reads Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/test_mcp_stale_session.py | 32 +++++++++++++++++++ 1 file changed, 32 insertions(+) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py index ec6fdef69ee..9a35227daf2 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py @@ -647,6 +647,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", @@ -735,6 +739,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", @@ -1030,6 +1038,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", @@ -1133,6 +1145,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", @@ -1221,6 +1237,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", @@ -1320,6 +1340,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", @@ -1577,6 +1601,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", @@ -1645,6 +1673,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", From f89b92763bcdbc7f1d42895fd480a55caf781c89 Mon Sep 17 00:00:00 2001 From: yucheng Date: Wed, 30 Sep 2026 10:47:23 +0000 Subject: [PATCH 20/51] fix(mcp): resolve /mcp/{name} routes through one exact-first lookup for connect, discovery and the scoped router Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/mcp_server_manager.py | 8 +++-- .../_experimental/mcp_server/operations.py | 22 +++++++----- .../mcp_server/test_mcp_server.py | 36 +++++++++++++++++++ .../mcp_server/test_mcp_server_manager.py | 22 ++++++++++++ .../mcp_server/test_mcp_server.py | 2 ++ 5 files changed, 80 insertions(+), 10 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 0611b153de5..707fa7625d7 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -7200,8 +7200,12 @@ class MCPServerManager: return None def get_mcp_server_answering_to(self, name: str, client_ip: str | None = None) -> MCPServer | None: - """The server a scoped ``/mcp/{name}`` connect resolves to: alias, then server_name, then name, each - case-insensitive so ``/mcp/GH`` and ``/mcp/gh`` agree, then the router's prefix match.""" + """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 same priority case-insensitively, then any prefix form routing accepts.""" + exact: Final = self.get_mcp_server_by_name(name, client_ip=client_ip) + if exact is not None: + return exact requested: Final = name.lower() servers: Final = tuple(self.get_registry().values()) identifiers: Final[tuple[Callable[[MCPServer], str | None], ...]] = ( diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index 50101975dca..9473f139def 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -455,15 +455,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 + if (scoped := _scoped_server(server_or_group, allowed_mcp_servers)) is not None: + filtered_server[scoped.server_id] = scoped - 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: + if scoped is None: try: access_group_server_ids = await MCPRequestHandler._get_mcp_servers_from_access_groups( [server_or_group] @@ -500,6 +495,17 @@ def _server_answers_to(server: MCPServer, name: str) -> bool: return server_answers_to_name(server, name) +def _scoped_server(name: str, allowed_mcp_servers: Sequence[MCPServer]) -> MCPServer | None: + """The granted server a scoped ``name`` selects: the registry's ``get_mcp_server_answering_to`` pick when + the caller holds it, so the router agrees with the connect preflight and discovery, and none when the + registry names a server the caller does not hold. Names the registry cannot place fall back to the first + granted server answering to them.""" + registry_pick: Final = global_mcp_server_manager.get_mcp_server_answering_to(name) + if registry_pick is not None: + return next((s for s in allowed_mcp_servers if s.server_id == registry_pick.server_id), None) + return next((s for s in allowed_mcp_servers if s and _server_answers_to(s, name)), None) + + async def raise_denied_scoped_mcp_access( requested_names: Sequence[str], user_api_key_auth: UserAPIKeyAuth | None, diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index c7dff4fa82e..03d4414ec7a 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -1279,6 +1279,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 = lambda name, client_ip=None: None # 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, @@ -6764,6 +6765,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 = MagicMock(return_value=None) 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) @@ -8588,6 +8590,39 @@ 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 only_b == [], "a name the registry gives to an ungranted server must not fall through to another" + + @pytest.mark.asyncio async def test_get_allowed_mcp_servers_from_mcp_server_names_mixed_known_and_unknown(): """ @@ -9729,6 +9764,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 = lambda name, client_ip=None: None 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): diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 1a4dd2cb93f..ccf9d7ce516 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -6231,6 +6231,28 @@ class TestMCPServerManager: 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 + def test_remove_server_drops_only_its_own_tool_mapping_rows(self): manager = self._manager_with_deepwiki_and_huggingface() 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 f8bf72428aa..cb4c690e18c 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 = 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( @@ -1002,6 +1003,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( From 2172b8b60e7d9a7dd7baac99ac78818b82810ee5 Mon Sep 17 00:00:00 2001 From: yucheng Date: Wed, 30 Sep 2026 11:26:45 +0000 Subject: [PATCH 21/51] fix(mcp): resolve scoped, connect and discovery routes through one exact-first, ip-aware lookup and stop denied names widening to access groups Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/discoverable_endpoints.py | 6 -- .../mcp_server/mcp_server_manager.py | 6 +- .../_experimental/mcp_server/operations.py | 29 +++++--- .../mcp/test_mcp_caller_sign_in.py | 71 +++++++++++++++++- .../mcp_server/test_mcp_server.py | 72 ++++++++++++++++++- .../mcp_server/test_mcp_server_manager.py | 28 ++++++++ 6 files changed, 190 insertions(+), 22 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 5c4894c3bd6..9e262432ade 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -528,12 +528,6 @@ 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 - by_id: Final = global_mcp_server_manager.get_mcp_server_by_id(lookup, client_ip=client_ip) - if by_id is not None: - return by_id return global_mcp_server_manager.get_mcp_server_answering_to(lookup, client_ip=client_ip) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 707fa7625d7..8fbb075c9e2 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -7202,10 +7202,14 @@ class MCPServerManager: def get_mcp_server_answering_to(self, name: str, client_ip: str | 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 same priority case-insensitively, then any prefix form routing accepts.""" + priority first, then the exact ``server_id``, then the name priority case-insensitively, then any prefix + form routing accepts.""" exact: Final = self.get_mcp_server_by_name(name, client_ip=client_ip) if exact is not None: return exact + by_id: Final = self.get_mcp_server_by_id(name, client_ip=client_ip) + if by_id is not None: + return by_id requested: Final = name.lower() servers: Final = tuple(self.get_registry().values()) identifiers: Final[tuple[Callable[[MCPServer], str | None], ...]] = ( diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index 9473f139def..c04c32a6f29 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -6,7 +6,7 @@ import types import uuid from collections.abc import Mapping, Sequence from datetime import datetime -from typing import Any, Final, NoReturn, TypeAlias, overload +from typing import Any, Final, Literal, NoReturn, TypeAlias, overload from fastapi import HTTPException from mcp import ReadResourceResult, Resource @@ -439,6 +439,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. @@ -455,10 +456,12 @@ 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: - if (scoped := _scoped_server(server_or_group, allowed_mcp_servers)) is not None: + scoped = _scoped_server(server_or_group, allowed_mcp_servers, client_ip) + if isinstance(scoped, str): + verbose_logger.debug("MCP scope name %s names a server the caller does not hold", server_or_group) + elif scoped is not None: filtered_server[scoped.server_id] = scoped - - if scoped is None: + else: try: access_group_server_ids = await MCPRequestHandler._get_mcp_servers_from_access_groups( [server_or_group] @@ -495,14 +498,16 @@ def _server_answers_to(server: MCPServer, name: str) -> bool: return server_answers_to_name(server, name) -def _scoped_server(name: str, allowed_mcp_servers: Sequence[MCPServer]) -> MCPServer | None: - """The granted server a scoped ``name`` selects: the registry's ``get_mcp_server_answering_to`` pick when - the caller holds it, so the router agrees with the connect preflight and discovery, and none when the - registry names a server the caller does not hold. Names the registry cannot place fall back to the first - granted server answering to them.""" - registry_pick: Final = global_mcp_server_manager.get_mcp_server_answering_to(name) +def _scoped_server( + name: str, allowed_mcp_servers: Sequence[MCPServer], client_ip: str | None +) -> MCPServer | Literal["denied"] | None: + """The granted server a scoped ``name`` selects: the registry's ``get_mcp_server_answering_to`` pick, made + with the same ``client_ip`` the connect preflight and discovery use, when the caller holds it. ``"denied"`` + when the registry names a server the caller does not hold, so the name is not retried as an access group. + ``None`` when the registry cannot place the name, after trying the granted servers answering to it.""" + registry_pick: Final = global_mcp_server_manager.get_mcp_server_answering_to(name, client_ip=client_ip) if registry_pick is not None: - return next((s for s in allowed_mcp_servers if s.server_id == registry_pick.server_id), None) + return next((s for s in allowed_mcp_servers if s.server_id == registry_pick.server_id), "denied") return next((s for s in allowed_mcp_servers if s and _server_answers_to(s, name)), None) @@ -684,6 +689,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 @@ -2431,6 +2437,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/tests/integration/mcp/test_mcp_caller_sign_in.py b/tests/integration/mcp/test_mcp_caller_sign_in.py index 9e940470031..bf5d7de3e32 100644 --- a/tests/integration/mcp/test_mcp_caller_sign_in.py +++ b/tests/integration/mcp/test_mcp_caller_sign_in.py @@ -24,14 +24,34 @@ 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]) -> httpx.Response: +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": "initialize", "params": INITIALIZE}, + 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 _sign_in_config(guardrail_params: dict[str, object], path: Path) -> 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}] @@ -115,6 +135,53 @@ def test_alias_first_lookup_wins_over_a_server_whose_name_matches_the_alias(gate 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) + + challenged: Final = _rpc(candidate, f"/mcp/{stem}", key, {}) + assert challenged.status_code == 401, challenged.text + assert f'resource_metadata="/.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_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())) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 03d4414ec7a..7da59cae4d0 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -6,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 @@ -8623,6 +8623,74 @@ async def test_scoped_router_selects_the_server_the_connect_preflight_resolves(a assert only_b == [], "a name the registry gives to an ungranted server must not fall through to another" +@pytest.mark.asyncio +async def test_scoped_name_of_an_ungranted_server_is_not_retried_as_an_access_group(): + 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="shared", 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({"p-id": private, "m-id": member}) + 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=["m-id"], + ) as groups: + denied = await _get_allowed_mcp_servers_from_mcp_server_names( + mcp_servers=["shared"], allowed_mcp_servers=[member] + ) + unknown = await _get_allowed_mcp_servers_from_mcp_server_names( + mcp_servers=["team"], allowed_mcp_servers=[member] + ) + finally: + global_mcp_server_manager.registry.clear() + + assert denied == [], "a denied server name must not widen to an access group of the same name" + assert [s.server_id for s in unknown] == ["m-id"] + assert groups.await_args_list == [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 + ) + 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 [s.server_id for s in external] == ["u-id"], "the router must apply the connect preflight's IP filter" + assert internal == [] + + @pytest.mark.asyncio async def test_get_allowed_mcp_servers_from_mcp_server_names_mixed_known_and_unknown(): """ @@ -9470,7 +9538,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"]) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index ccf9d7ce516..87d87a3e74c 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -6253,6 +6253,34 @@ class TestMCPServerManager: 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("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() From f2a876f23f6c3a875c7c3b4444c2213fd9ccfbe4 Mon Sep 17 00:00:00 2001 From: yucheng Date: Wed, 30 Sep 2026 12:24:16 +0000 Subject: [PATCH 22/51] fix(mcp): keep a route hidden from a client ip from rerouting to a case variant of its name get_mcp_server_answering_to stops at the pass that finds an exact name or id and hides it from client_ip instead of falling through to the case-insensitive and prefix passes, and _scoped_server treats a name the registry knows for some caller but not this one as denied, so connect, discovery and scoped routing all refuse the hidden route instead of serving a public server whose alias only differs by case Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/mcp_server_manager.py | 11 +++-- .../_experimental/mcp_server/operations.py | 7 ++- .../mcp/test_mcp_caller_sign_in.py | 48 +++++++++++++++++++ .../mcp_server/test_discoverable_endpoints.py | 20 ++++---- .../mcp_server/test_mcp_server.py | 6 ++- .../mcp_server/test_mcp_server_manager.py | 18 +++++++ 6 files changed, 92 insertions(+), 18 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 8fbb075c9e2..5282b228f92 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -7203,13 +7203,14 @@ class MCPServerManager: """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.""" - exact: Final = self.get_mcp_server_by_name(name, client_ip=client_ip) + 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.""" + exact: Final = self.get_mcp_server_by_name(name) if exact is not None: - return exact - by_id: Final = self.get_mcp_server_by_id(name, client_ip=client_ip) + 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 + 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], ...]] = ( diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index c04c32a6f29..15a9183ef7a 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -503,11 +503,14 @@ def _scoped_server( ) -> MCPServer | Literal["denied"] | None: """The granted server a scoped ``name`` selects: the registry's ``get_mcp_server_answering_to`` pick, made with the same ``client_ip`` the connect preflight and discovery use, when the caller holds it. ``"denied"`` - when the registry names a server the caller does not hold, so the name is not retried as an access group. - ``None`` when the registry cannot place the name, after trying the granted servers answering to it.""" + when the registry names a server the caller does not hold, or one hidden from ``client_ip``, so the name + is neither rerouted to another granted server nor retried as an access group. ``None`` when the registry + cannot place the name for any caller, after trying the granted servers answering to it.""" registry_pick: Final = global_mcp_server_manager.get_mcp_server_answering_to(name, client_ip=client_ip) if registry_pick is not None: return next((s for s in allowed_mcp_servers if s.server_id == registry_pick.server_id), "denied") + if global_mcp_server_manager.get_mcp_server_answering_to(name) is not None: + return "denied" return next((s for s in allowed_mcp_servers if s and _server_answers_to(s, name)), None) diff --git a/tests/integration/mcp/test_mcp_caller_sign_in.py b/tests/integration/mcp/test_mcp_caller_sign_in.py index bf5d7de3e32..9755031fe3c 100644 --- a/tests/integration/mcp/test_mcp_caller_sign_in.py +++ b/tests/integration/mcp/test_mcp_caller_sign_in.py @@ -182,6 +182,54 @@ def test_exact_name_wins_over_a_case_folded_config_alias_for_connect_discovery_a 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 == 403 + + 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())) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 361d78059ad..ccf968e8179 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -3642,8 +3642,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 +3679,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 +3711,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 +3739,8 @@ 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 @@ -3815,8 +3815,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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 7da59cae4d0..5cf35fb597c 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -8682,13 +8682,17 @@ async def test_scoped_router_hides_a_private_server_from_an_external_ip_like_the 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 [s.server_id for s in external] == ["u-id"], "the router must apply the connect preflight's IP filter" + assert external == [], "a name the preflight hides from this IP must not reroute to a case variant" assert internal == [] + assert [s.server_id for s in by_own_name] == ["u-id"] @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 87d87a3e74c..d7a1b3090a0 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -6253,6 +6253,24 @@ class TestMCPServerManager: 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 6adff6ffacb503136cff3553efe049845344e797 Mon Sep 17 00:00:00 2001 From: yucheng Date: Wed, 30 Sep 2026 12:45:46 +0000 Subject: [PATCH 23/51] test(mcp): cover throttled token exchange surfacing as an outage rather than a sign-in challenge Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- tests/integration/mcp/test_mcp_oauth_flows.py | 30 +++++++++++++++++++ 1 file changed, 30 insertions(+) diff --git a/tests/integration/mcp/test_mcp_oauth_flows.py b/tests/integration/mcp/test_mcp_oauth_flows.py index 7a60c8ede30..8b5fc14f213 100644 --- a/tests/integration/mcp/test_mcp_oauth_flows.py +++ b/tests/integration/mcp/test_mcp_oauth_flows.py @@ -1,5 +1,6 @@ import base64 import hashlib +import json import secrets import uuid from dataclasses import dataclass @@ -21,6 +22,7 @@ from integration._support.mcp import ( tool_calls, ) 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" @@ -197,6 +199,34 @@ 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 _assert_subject_token_challenge(response: httpx.Response, alias: str) -> None: assert response.status_code == 401, response.text challenge: Final = response.headers["www-authenticate"] From 265919f9a92bac6cc2f6b0a92503e3781f660114 Mon Sep 17 00:00:00 2001 From: yucheng Date: Wed, 30 Sep 2026 18:50:04 +0000 Subject: [PATCH 24/51] fix(mcp): annotate the general_settings cast for the type-discipline gate Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/_experimental/mcp_server/caller_sign_in.py | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/_experimental/mcp_server/caller_sign_in.py b/litellm/proxy/_experimental/mcp_server/caller_sign_in.py index e11c313f94a..89a920ea7e4 100644 --- a/litellm/proxy/_experimental/mcp_server/caller_sign_in.py +++ b/litellm/proxy/_experimental/mcp_server/caller_sign_in.py @@ -125,7 +125,12 @@ def jwt_auth_issuers() -> tuple[str, ...]: env_issuer: Final = os.getenv("JWT_ISSUER") env: Final[tuple[str, ...]] = (env_issuer,) if env_issuer else () - settings: Final[Mapping[str, object]] = cast(Mapping[str, object], general_settings) + settings: Final = ( + cast( # cast-ok: general_settings is a raw dict; the value is validated by _jwt_auth_issuer_entries + Mapping[str, object], + general_settings, + ) + ) configured: Final = tuple( entry.issuer for entry in _jwt_auth_issuer_entries(settings.get("litellm_jwtauth")) if entry.issuer ) From 86c3b910a6abc38baa9a3a5a8887bba025578c58 Mon Sep 17 00:00:00 2001 From: Yucheng He Date: Thu, 1 Oct 2026 18:00:25 -0700 Subject: [PATCH 25/51] fix(mcp): route a scoped name to the caller's granted server before the registry's pick A key granted only the server named `docs` was refused with 403 on /mcp/docs when an ungranted server held `docs` as its alias, and a key granted only `github` was refused on /mcp/GITHUB when an ungranted `GitHub` existed: the scoped router took the registry-wide pick and answered "denied" whenever that pick was not among the caller's servers. The router now runs the registry's own pass order (exact alias, server_name, name, server_id; the same case-insensitively; prefix forms) over the caller's granted servers first, through get_mcp_server_answering_to(among=...), with IP hiding applied at the pass that found the server. A name hidden from the client IP stays denied before any grant lookup, and a name the registry places only on an ungranted server stays denied rather than being retried as an access group. The connect preflight reuses the router's selection for a single scoped name as the server it challenges, signs in and exchanges for, so a 401 names the granted server; an ungranted caller keeps the registry pick and the downstream 403, and the no-key path is unchanged. Unauthenticated RFC 9728 discovery has no grant list and stays on the registry pick. Tests: the `only_b` assertion in test_scoped_router_selects_the_server_the_connect_preflight_resolves and the no-IP `internal` assertion in test_scoped_router_hides_a_private_server_from_an_external_ip_like_the_connect_preflight now expect the granted server, which is what the base branch returned for both shapes; the unentitled-key exchange test's fixture becomes the empty selection the router returns for such a key; listing tests that stub the manager now point its lookup at a real empty manager so the `among` pass runs. New tests cover the grant-first router, the manager `among` pass order and IP hiding, and the connect preflight under an alias collision. --- .../mcp_server/mcp_server_manager.py | 37 ++- .../_experimental/mcp_server/operations.py | 23 +- .../proxy/_experimental/mcp_server/server.py | 34 +-- .../mcp_server/test_mcp_server.py | 2 +- .../test_mcp_server_tool_calls_and_headers.py | 210 +++++++++++++++++- 5 files changed, 270 insertions(+), 36 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 5282b228f92..4ffa2ba5534 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -7199,12 +7199,18 @@ class MCPServerManager: return server return None - def get_mcp_server_answering_to(self, name: str, client_ip: str | None = None) -> MCPServer | 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.""" + 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 @@ -7230,6 +7236,33 @@ class MCPServerManager: 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 15a9183ef7a..8636afbebef 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -501,17 +501,22 @@ def _server_answers_to(server: MCPServer, name: str) -> bool: def _scoped_server( name: str, allowed_mcp_servers: Sequence[MCPServer], client_ip: str | None ) -> MCPServer | Literal["denied"] | None: - """The granted server a scoped ``name`` selects: the registry's ``get_mcp_server_answering_to`` pick, made - with the same ``client_ip`` the connect preflight and discovery use, when the caller holds it. ``"denied"`` - when the registry names a server the caller does not hold, or one hidden from ``client_ip``, so the name - is neither rerouted to another granted server nor retried as an access group. ``None`` when the registry - cannot place the name for any caller, after trying the granted servers answering to it.""" + """The server a scoped ``name`` selects for the caller, in this order. ``"denied"`` when the registry's + ``get_mcp_server_answering_to`` pick, made with the same ``client_ip`` the connect preflight and discovery + use, is a server hidden from that IP. Otherwise the granted server answering to ``name``: 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. With no granted server answering: ``"denied"`` when the registry + places the name on a server the caller does not hold, so it is not retried as an access group; ``None`` + when the registry cannot place the name for any caller.""" registry_pick: Final = global_mcp_server_manager.get_mcp_server_answering_to(name, client_ip=client_ip) - if registry_pick is not None: - return next((s for s in allowed_mcp_servers if s.server_id == registry_pick.server_id), "denied") - if global_mcp_server_manager.get_mcp_server_answering_to(name) is not None: + if registry_pick is None and global_mcp_server_manager.get_mcp_server_answering_to(name) is not None: return "denied" - return next((s for s in allowed_mcp_servers if s and _server_answers_to(s, name)), None) + granted: Final = global_mcp_server_manager.get_mcp_server_answering_to( + name, client_ip=client_ip, among=allowed_mcp_servers + ) + if granted is not None: + return granted + return "denied" if registry_pick is not None else None async def raise_denied_scoped_mcp_access( diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 988f77ea98b..49fa2f898ed 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -1603,7 +1603,24 @@ 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_answering_to(server_name, client_ip=client_ip) + registry_pick = operations.global_mcp_server_manager.get_mcp_server_answering_to( + server_name, client_ip=client_ip + ) + obo_without_subject = ( + registry_pick is not None + and registry_pick.auth_type == MCPAuth.oauth2_token_exchange + and not oauth2_headers + ) + 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 not obo_without_subject and mcp_servers is not None and len(mcp_servers) == 1 + else () + ) + granted = next(iter(allowed_single), None) + server = granted if granted is not None else registry_pick + granted_single = granted is not None 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 @@ -1703,7 +1720,7 @@ if MCP_AVAILABLE: # 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 below serves the challenge, the sign-in preflight and the exchange. + # 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 subject_token = ( operations.global_mcp_server_manager._extract_subject_token( # pyright: ignore[reportPrivateUsage] # the manager owns the subject/admission filter shared with the preflight @@ -1712,19 +1729,6 @@ if MCP_AVAILABLE: if server is not None else None ) - obo_without_subject = ( - server is not None and server.auth_type == MCPAuth.oauth2_token_exchange and not oauth2_headers - ) - 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 server and not obo_without_subject and mcp_servers is not None and len(mcp_servers) == 1 - else () - ) - granted_single = server is not None and any( - allowed.server_id == server.server_id for allowed in allowed_single - ) 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, 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 cb4c690e18c..d7ba2a3cc93 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py @@ -924,7 +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 = MagicMock(return_value=None) + 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( 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 14f284389e0..c733ba58b1c 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 @@ -32,7 +32,7 @@ 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, @@ -1279,7 +1279,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 = lambda name, client_ip=None: None + 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, @@ -6765,7 +6765,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 = MagicMock(return_value=None) + 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) @@ -8620,7 +8620,9 @@ async def test_scoped_router_selects_the_server_the_connect_preflight_resolves(a finally: global_mcp_server_manager.registry.clear() - assert only_b == [], "a name the registry gives to an ungranted server must not fall through to another" + 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 @@ -8691,10 +8693,202 @@ async def test_scoped_router_hides_a_private_server_from_an_external_ip_like_the global_mcp_server_manager.registry.clear() assert external == [], "a name the preflight hides from this IP must not reroute to a case variant" - assert internal == [] + 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"), + pytest.param(("p-id", "m-id"), "shared", ("m-id",), [], id="ungranted-only-name-stays-denied"), + ], +) +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), + "p-id": MCPServer(server_id="p-id", name="p", server_name="p", alias="shared", transport=MCPTransport.http), + "m-id": 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({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=["m-id"], + ) 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"}, + ) + + @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} as {kwargs['connected_as']}" + ) + + 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 as 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"] + + @pytest.mark.asyncio async def test_get_allowed_mcp_servers_from_mcp_server_names_mixed_known_and_unknown(): """ @@ -9836,7 +10030,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 = lambda name, client_ip=None: None + 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): @@ -10351,9 +10545,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( From 0a1134d06ab18eece9a74142477a2d922527e13e Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 2 Oct 2026 08:28:39 +0000 Subject: [PATCH 26/51] fix(mcp): satisfy the strict-gate typing import ban in caller sign-in Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/_experimental/mcp_server/caller_sign_in.py | 9 ++------- 1 file changed, 2 insertions(+), 7 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/caller_sign_in.py b/litellm/proxy/_experimental/mcp_server/caller_sign_in.py index 89a920ea7e4..ac054d7f737 100644 --- a/litellm/proxy/_experimental/mcp_server/caller_sign_in.py +++ b/litellm/proxy/_experimental/mcp_server/caller_sign_in.py @@ -17,7 +17,7 @@ from __future__ import annotations import itertools from collections.abc import Mapping from dataclasses import dataclass -from typing import TYPE_CHECKING, Final, Protocol, cast, runtime_checkable +from typing import TYPE_CHECKING, Final, Protocol, runtime_checkable from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError from typing_extensions import assert_never @@ -125,12 +125,7 @@ def jwt_auth_issuers() -> tuple[str, ...]: env_issuer: Final = os.getenv("JWT_ISSUER") env: Final[tuple[str, ...]] = (env_issuer,) if env_issuer else () - settings: Final = ( - cast( # cast-ok: general_settings is a raw dict; the value is validated by _jwt_auth_issuer_entries - Mapping[str, object], - general_settings, - ) - ) + 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 ) From 0b262d6594bf0603d01cad884ef6fa3bd7e677da Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 2 Oct 2026 17:57:24 +0000 Subject: [PATCH 27/51] fix(mcp): keep the plain OBO challenge naming the configured alias on every route The connect preflight passes connected_as only for caller sign-in servers, so an oauth2_token_exchange server challenges with the merge-base resource_metadata path (its alias) on /mcp/, the aggregate route and the legacy route as well as on /mcp/. The case-variant integration assertion that expected 403 is corrected to the 200 both base and head return Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/_experimental/mcp_server/server.py | 9 ++++++--- tests/integration/mcp/test_mcp_caller_sign_in.py | 2 +- .../test_mcp_server_tool_calls_and_headers.py | 14 ++++++++------ 3 files changed, 15 insertions(+), 10 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index cc02a88bfb0..9e579500f69 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -1741,6 +1741,9 @@ if MCP_AVAILABLE: # 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 + challenge_route: str | None = ( + None if server is not None and server.auth_type == MCPAuth.oauth2_token_exchange else server_name + ) subject_token = ( operations.global_mcp_server_manager._extract_subject_token( # pyright: ignore[reportPrivateUsage] # the manager owns the subject/admission filter shared with the preflight oauth2_headers, raw_headers, user_api_key_auth @@ -1756,7 +1759,7 @@ if MCP_AVAILABLE: get_request_root_path, ) - raise_token_exchange_challenge(server, root_path=get_request_root_path(), connected_as=server_name) + raise_token_exchange_challenge(server, root_path=get_request_root_path(), connected_as=challenge_route) 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, @@ -1770,7 +1773,7 @@ if MCP_AVAILABLE: user_api_key_auth, subject_token, root_path=get_request_root_path(), - connected_as=server_name, + connected_as=challenge_route, ) # Exchange-backed modes (token_exchange's OBO mint, id_jag's stored-assertion mint): run @@ -1786,7 +1789,7 @@ if MCP_AVAILABLE: oauth2_headers=oauth2_headers, user_api_key_auth=user_api_key_auth, raw_headers=raw_headers, - connected_as=server_name, + connected_as=challenge_route, ) # Pass-through OAuth: when the admin has opted a server into diff --git a/tests/integration/mcp/test_mcp_caller_sign_in.py b/tests/integration/mcp/test_mcp_caller_sign_in.py index 9755031fe3c..ddf0491f50e 100644 --- a/tests/integration/mcp/test_mcp_caller_sign_in.py +++ b/tests/integration/mcp/test_mcp_caller_sign_in.py @@ -211,7 +211,7 @@ def test_name_of_a_server_hidden_from_an_external_ip_does_not_reroute_to_a_case_ 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 == 403 + 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 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 01809d77b37..e73168a4448 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 @@ -10664,7 +10664,7 @@ class TestOboPreflightScopedToAllowedServers: "x-litellm-api-key": key.api_key, "authorization": self.SUBJECT_HEADERS["Authorization"], }, - connected_as=requested.alias, + connected_as=None, ) @@ -11574,13 +11574,15 @@ class TestConnectChallengeResolver: assert exc.value.status_code == 401 @pytest.mark.asyncio - async def test_obo_challenge_www_authenticate_matches_main_byte_for_byte(self, monkeypatch): + @pytest.mark.parametrize("route_name", ["obo", "obo_server"], ids=["alias_route", "server_name_route"]) + async def test_obo_challenge_www_authenticate_matches_main_byte_for_byte(self, monkeypatch, route_name): """The provider redesign must not change what an OBO server challenges with: the relative - RFC 9728 resource_metadata path plus the RFC 6750 invalid_token triple, exactly as main.""" + RFC 9728 resource_metadata path naming the configured alias whichever route the client used, + plus the RFC 6750 invalid_token triple, exactly as main.""" from litellm.proxy._experimental.mcp_server import server as server_module monkeypatch.delenv("SERVER_ROOT_PATH", raising=False) - obo = _make_obo_server("obo") + obo = _make_obo_server("obo").model_copy(update={"name": "obo_server", "server_name": "obo_server"}) with ( patch.object( mcp_operations.global_mcp_server_manager, @@ -11590,8 +11592,8 @@ class TestConnectChallengeResolver: pytest.raises(HTTPException) as exc, ): await server_module._raise_preemptive_401_for_unauthenticated_servers( - scope={"type": "http", "method": "POST", "path": "/mcp/obo", "headers": []}, - mcp_servers=["obo"], + 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"), From ee7c8a20d058547ee7b4d66b7c1bb5b3313d577d Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 2 Oct 2026 18:05:39 +0000 Subject: [PATCH 28/51] fix(mcp): resolve the caller's granted server before the registry pick at connect The connect preflight runs the scoped router's grant-first lookup for every single-server connect, so a granted plain server named like an ungranted OBO server's alias is served instead of intercepted by that server's 401, and obo_without_subject is read off the server the lookup picks Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../proxy/_experimental/mcp_server/server.py | 10 +- .../test_mcp_server_tool_calls_and_headers.py | 94 +++++++++++++++++++ 2 files changed, 98 insertions(+), 6 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 9e579500f69..5fb84c64258 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -1625,21 +1625,19 @@ if MCP_AVAILABLE: registry_pick = operations.global_mcp_server_manager.get_mcp_server_answering_to( server_name, client_ip=client_ip ) - obo_without_subject = ( - registry_pick is not None - and registry_pick.auth_type == MCPAuth.oauth2_token_exchange - and not oauth2_headers - ) 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 not obo_without_subject and mcp_servers is not None and len(mcp_servers) == 1 + if registry_pick and mcp_servers is not None and len(mcp_servers) == 1 else () ) granted = next(iter(allowed_single), 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 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 e73168a4448..7db98279e3a 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 @@ -8985,6 +8985,100 @@ class TestConnectPreflightRoutesLikeTheScopedRouter: ) assert guardrail.asked_about == ["d-id"] + @staticmethod + def _register_obo_alias_collision() -> tuple[MCPServer, MCPServer]: + 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}) + return obo, 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"}, + ) + + @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): + _, plain = 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=["p-id"]) + ), + 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 + ), + ): + await self._connect_to(route_name) + + assert preflight.await_args is not None, "the granted plain server must reach the preflight, not a 401" + assert preflight.await_args.kwargs["server"] is plain + + @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( + 'Bearer resource_metadata="/.well-known/oauth-protected-resource/mcp/obo"' + ) + + @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_get_allowed_mcp_servers_from_mcp_server_names_mixed_known_and_unknown(): From d2e0d4e9e128923bf036eef210bf9435f23cd81e Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 2 Oct 2026 18:13:32 +0000 Subject: [PATCH 29/51] test(mcp): assert the granted plain server's exchange through what it reports, not through the mock Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../test_mcp_server_tool_calls_and_headers.py | 21 ++++++++++++------- 1 file changed, 13 insertions(+), 8 deletions(-) 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 7db98279e3a..bd0e18d766b 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 @@ -8986,7 +8986,7 @@ class TestConnectPreflightRoutesLikeTheScopedRouter: assert guardrail.asked_about == ["d-id"] @staticmethod - def _register_obo_alias_collision() -> tuple[MCPServer, MCPServer]: + def _register_obo_alias_collision() -> None: obo: Final = MCPServer( server_id="o-id", name="obo_server", @@ -9008,7 +9008,6 @@ class TestConnectPreflightRoutesLikeTheScopedRouter: auth_type=MCPAuth.none, ) mcp_operations.global_mcp_server_manager.registry.update({"o-id": obo, "p-id": plain}) - return obo, plain @staticmethod async def _connect_to(route_name: str) -> None: @@ -9027,20 +9026,26 @@ class TestConnectPreflightRoutesLikeTheScopedRouter: @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): - _, plain = self._register_obo_alias_collision() - preflight: Final = AsyncMock() + self._register_obo_alias_collision() + + async def report_exchange(server, **kwargs): + raise HTTPException( + status_code=401, detail=f"exchange ran for {server.server_id} as {kwargs['connected_as']}" + ) + 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 records which server ran - mcp_operations.global_mcp_server_manager, "preflight_token_exchange", preflight + 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 preflight.await_args is not None, "the granted plain server must reach the preflight, not a 401" - assert preflight.await_args.kwargs["server"] is plain + assert exc.value.status_code == 401 + assert exc.value.detail == f"exchange ran for p-id as {route_name}" @pytest.mark.asyncio @pytest.mark.parametrize("route_name", ["obo", "obo_server"], ids=["alias", "server_name"]) From ba0dcf0f63ad523878f3b3e4b63cc770fb77d7a7 Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 2 Oct 2026 18:22:52 +0000 Subject: [PATCH 30/51] test(mcp): expect the granted exact-name server on the colliding alias route, keep the challenge for an unscoped key Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- tests/integration/mcp/test_mcp_caller_sign_in.py | 12 +++++++++++- 1 file changed, 11 insertions(+), 1 deletion(-) diff --git a/tests/integration/mcp/test_mcp_caller_sign_in.py b/tests/integration/mcp/test_mcp_caller_sign_in.py index ddf0491f50e..26fb8485ace 100644 --- a/tests/integration/mcp/test_mcp_caller_sign_in.py +++ b/tests/integration/mcp/test_mcp_caller_sign_in.py @@ -173,7 +173,17 @@ def test_exact_name_wins_over_a_case_folded_config_alias_for_connect_discovery_a assert tool_calls(obo_peer.drain()) == () assert _advertised(candidate, cased) == _advertised(candidate, by_name) - challenged: Final = _rpc(candidate, f"/mcp/{stem}", key, {}) + 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="/.well-known/oauth-protected-resource/mcp/{stem}"' in challenged.headers.get( "www-authenticate", "" From b232b8cfa725a06a1606d56334a19f3a1fb3ed27 Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 2 Oct 2026 20:40:49 +0000 Subject: [PATCH 31/51] fix(mcp): caller-fault AADSTS assertions, server scopes and route-exact absolute challenge metadata A malformed caller assertion that Entra reports as invalid_client with an AADSTS50027xx error_codes entry is now the caller's 401 sign-in challenge in the shared token exchange provider, never a 503 gateway-credential fault or a fail-open pass, so both the Agent 365 guardrail and plain OBO servers classify it the same way The Agent 365 sign-in provider advertises a server's configured scopes for every server and falls back to api:///access_as_user only when none are configured Connect-time challenges from both the Agent 365 provider and plain OBO carry resource_metadata as an absolute URL built from the request origin (trusted forwarded headers honored) and naming the route the client actually used, so //mcp and /mcp/ each point at their own protected-resource document Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/caller_sign_in.py | 6 +- .../mcp_server/mcp_server_manager.py | 10 +- .../outbound_credentials/adapter.py | 14 +- .../token_exchange_provider.py | 50 +++++-- .../proxy/_experimental/mcp_server/server.py | 13 +- .../guardrail_hooks/agent_365/agent_365.py | 4 +- .../mcp/test_mcp_caller_sign_in.py | 102 +++++++++++++- tests/integration/mcp/test_mcp_oauth_flows.py | 35 +++++ .../outbound_credentials/test_adapter.py | 15 ++ .../test_token_exchange_provider.py | 28 ++++ .../mcp_server/test_mcp_server_manager.py | 7 +- .../test_mcp_server_tool_calls_and_headers.py | 132 ++++++++++++++++-- .../guardrail_hooks/test_agent_365.py | 86 +++++++++++- 13 files changed, 447 insertions(+), 55 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/caller_sign_in.py b/litellm/proxy/_experimental/mcp_server/caller_sign_in.py index ac054d7f737..1d417c01dbd 100644 --- a/litellm/proxy/_experimental/mcp_server/caller_sign_in.py +++ b/litellm/proxy/_experimental/mcp_server/caller_sign_in.py @@ -161,7 +161,7 @@ async def preflight_caller_sign_in( subject_token: str, *, root_path: str, - connected_as: str | None, + resource_metadata: str | None, ) -> None: """Run every provider's connect-time check against the subject token, so a bearer the IdP will reject surfaces as a challenge here rather than a JSON-RPC error at the first tool call.""" @@ -178,7 +178,9 @@ async def preflight_caller_sign_in( case SignedIn(): continue case Rejected(detail=_, claims=claims): - raise_token_exchange_challenge(server, root_path=root_path, claims=claims, connected_as=connected_as) + raise_token_exchange_challenge( + server, root_path=root_path, claims=claims, resource_metadata=resource_metadata + ) case Unavailable(detail=detail, fail_open=True): continue case Unavailable(detail=detail, fail_open=False): diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 5f74f39002e..9e3a9c2194e 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -4226,7 +4226,7 @@ class MCPServerManager: oauth2_headers: dict[str, str] | None, user_api_key_auth: UserAPIKeyAuth | None, raw_headers: Mapping[str, str] | None = None, - connected_as: str | None = None, + resource_metadata: str | None = None, ) -> None: """Mint an exchange-backed server's upstream credential at the transport edge. @@ -4263,7 +4263,9 @@ class MCPServerManager: ) if subject_token is None and caller_sign_in_for(server, user_api_key_auth) is not None: - raise_token_exchange_challenge(server, root_path=get_request_root_path(), connected_as=connected_as) + 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) @@ -4271,7 +4273,7 @@ class MCPServerManager: return if subject_token is None and isinstance(spec.config, TokenExchangeConfig): raise_token_exchange_challenge( - resolved_server, root_path=get_request_root_path(), connected_as=connected_as + 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(_): @@ -4282,7 +4284,7 @@ class MCPServerManager: resolved_server, root_path=get_request_root_path(), claims=err.unauthorized.claims, - connected_as=connected_as, + resource_metadata=resource_metadata, ) raise_public(err) diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/adapter.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/adapter.py index a4b723b5d67..abbde8f16f6 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/adapter.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/adapter.py @@ -298,7 +298,7 @@ def raise_public(error: CredError) -> NoReturn: assert_never(error.tag) -def oauth_protected_resource_path(root_path: str, server: MCPServer, *, connected_as: str | None = None) -> str: +def oauth_protected_resource_path(root_path: str, server: MCPServer) -> str: """The server's RFC 9728 Protected Resource Metadata path, the shared anchor of both challenges. ``root_path`` is the prefix the request was routed under, resolved by the caller (the imperative @@ -320,7 +320,7 @@ def oauth_protected_resource_path(root_path: str, server: MCPServer, *, connecte challenge would then disagree on where the resource metadata lives. """ prefix: Final = "" if root_path == "/" else root_path - name: Final = connected_as or server.alias or server.server_name or server.name or server.server_id + name: Final = server.alias or server.server_name or server.name or server.server_id scalar_env: Final = os.getenv("SERVER_ROOT_PATH", "").rstrip("/") if not prefix or (scalar_env and prefix == scalar_env): return f"/.well-known/oauth-protected-resource{prefix}/mcp/{name}" @@ -348,7 +348,7 @@ def raise_token_exchange_challenge( *, root_path: str, claims: str | None = None, - connected_as: 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. @@ -366,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, connected_as=connected_as) + 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 = ( @@ -377,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 7f5c0c99145..ad503996155 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 @@ -34,11 +35,29 @@ 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"} ) +# Entra reports a forged or garbled assertion as ``invalid_client`` with an AADSTS50027xx sub-code, +# the same top-level code as a bad gateway secret; the sub-code is what says the caller has to fix it. +_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 +67,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 (), ) @@ -85,18 +108,19 @@ async def _post_exchange_endpoint( 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 diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 5fb84c64258..d0efbf7c0fb 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -63,6 +63,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, @@ -1739,9 +1740,7 @@ if MCP_AVAILABLE: # 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 - challenge_route: str | None = ( - None if server is not None and server.auth_type == MCPAuth.oauth2_token_exchange else server_name - ) + resource_metadata = get_passthrough_resource_metadata_url(scope, server_name) subject_token = ( operations.global_mcp_server_manager._extract_subject_token( # pyright: ignore[reportPrivateUsage] # the manager owns the subject/admission filter shared with the preflight oauth2_headers, raw_headers, user_api_key_auth @@ -1757,7 +1756,9 @@ if MCP_AVAILABLE: get_request_root_path, ) - raise_token_exchange_challenge(server, root_path=get_request_root_path(), connected_as=challenge_route) + 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, @@ -1771,7 +1772,7 @@ if MCP_AVAILABLE: user_api_key_auth, subject_token, root_path=get_request_root_path(), - connected_as=challenge_route, + resource_metadata=resource_metadata, ) # Exchange-backed modes (token_exchange's OBO mint, id_jag's stored-assertion mint): run @@ -1787,7 +1788,7 @@ if MCP_AVAILABLE: oauth2_headers=oauth2_headers, user_api_key_auth=user_api_key_auth, raw_headers=raw_headers, - connected_as=challenge_route, + resource_metadata=resource_metadata, ) # Pass-through OAuth: when the admin has opted a server into 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 d5e0473abc7..ac09821feb3 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py +++ b/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py @@ -471,7 +471,9 @@ class Agent365Guardrail(CustomGuardrail): return None return CallerSignIn( issuers=(ENTRA_ISSUER_TEMPLATE.format(tenant_id=self.tenant_id),), - scopes=(GATEWAY_SCOPE_TEMPLATE.format(client_id=self.client_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]: diff --git a/tests/integration/mcp/test_mcp_caller_sign_in.py b/tests/integration/mcp/test_mcp_caller_sign_in.py index 26fb8485ace..53cec8a61e3 100644 --- a/tests/integration/mcp/test_mcp_caller_sign_in.py +++ b/tests/integration/mcp/test_mcp_caller_sign_in.py @@ -2,6 +2,7 @@ import json import uuid from collections.abc import Mapping from pathlib import Path +from types import MappingProxyType from typing import Final import httpx @@ -52,9 +53,16 @@ def _advertised(gateway: Gateway, segment: str) -> tuple[int, tuple[str, ...], o return response.status_code, issuers, document.get("scopes_supported") -def _sign_in_config(guardrail_params: dict[str, object], path: Path) -> Path: +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 @@ -185,8 +193,9 @@ def test_exact_name_wins_over_a_case_folded_config_alias_for_connect_discovery_a challenged: Final = _rpc(candidate, f"/mcp/{stem}", scenario.key(), {}) assert challenged.status_code == 401, challenged.text - assert f'resource_metadata="/.well-known/oauth-protected-resource/mcp/{stem}"' in challenged.headers.get( - "www-authenticate", "" + 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) @@ -323,13 +332,16 @@ def test_agent_365_gated_server_challenges_at_connect_and_advertises_entra(gatew 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="/.well-known/oauth-protected-resource/mcp/{alias}"' in 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="/.well-known/oauth-protected-resource/mcp/{alias}"' in opaque.headers.get( - "www-authenticate", "" + 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}") @@ -344,6 +356,79 @@ def test_agent_365_gated_server_challenges_at_connect_and_advertises_entra(gatew 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, 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 ( @@ -359,7 +444,10 @@ def test_challenge_and_prm_resolve_the_connected_case_variant(gateway: Gateway, 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="/.well-known/oauth-protected-resource/mcp/{connected_as}"' in 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}") diff --git a/tests/integration/mcp/test_mcp_oauth_flows.py b/tests/integration/mcp/test_mcp_oauth_flows.py index 252fb6228ea..0e65ee9d2b1 100644 --- a/tests/integration/mcp/test_mcp_oauth_flows.py +++ b/tests/integration/mcp/test_mcp_oauth_flows.py @@ -233,6 +233,41 @@ def test_a_throttled_token_exchange_is_an_outage_not_a_sign_in_challenge(gateway 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 _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 002e70da9d2..589c9ce16f5 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 @@ -89,6 +89,34 @@ async def test_post_maps_gateway_fault_4xx_to_client_error(code): 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"}, {}) + + @pytest.mark.asyncio @pytest.mark.parametrize( "body", 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 9ad7068cac0..4c3c10474bf 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 @@ -2952,11 +2952,14 @@ class TestMCPServerManager: server=server, oauth2_headers={"Authorization": "Bearer rejected-subject"}, user_api_key_auth=None, - connected_as=server.server_id, + 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"/.well-known/oauth-protected-resource/mcp/{server.server_id}" in www_authenticate, www_authenticate + 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): 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 bd0e18d766b..258d24ff619 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 @@ -8941,7 +8941,7 @@ class TestConnectPreflightRoutesLikeTheScopedRouter: async def refuse_exchange(server, **kwargs): raise HTTPException( - status_code=401, detail=f"exchange refused for {server.server_id} as {kwargs['connected_as']}" + status_code=401, detail=f"exchange refused for {server.server_id} at {kwargs['resource_metadata']}" ) with ( @@ -8956,7 +8956,7 @@ class TestConnectPreflightRoutesLikeTheScopedRouter: await self._connect_to_docs() assert exc.value.status_code == 401 - assert exc.value.detail == "exchange refused for d-id as docs" + 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): @@ -9030,7 +9030,7 @@ class TestConnectPreflightRoutesLikeTheScopedRouter: async def report_exchange(server, **kwargs): raise HTTPException( - status_code=401, detail=f"exchange ran for {server.server_id} as {kwargs['connected_as']}" + status_code=401, detail=f"exchange ran for {server.server_id} at {kwargs['resource_metadata']}" ) with ( @@ -9045,7 +9045,7 @@ class TestConnectPreflightRoutesLikeTheScopedRouter: await self._connect_to(route_name) assert exc.value.status_code == 401 - assert exc.value.detail == f"exchange ran for p-id as {route_name}" + 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"]) @@ -9062,7 +9062,7 @@ class TestConnectPreflightRoutesLikeTheScopedRouter: 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/obo"' + f'Bearer resource_metadata="/.well-known/oauth-protected-resource/mcp/{route_name}"' ) @pytest.mark.asyncio @@ -10763,7 +10763,7 @@ class TestOboPreflightScopedToAllowedServers: "x-litellm-api-key": key.api_key, "authorization": self.SUBJECT_HEADERS["Authorization"], }, - connected_as=None, + resource_metadata=f"/.well-known/oauth-protected-resource/mcp/{requested.alias}", ) @@ -11546,6 +11546,22 @@ async def test_discovery_adapter_preserves_authenticated_context(_mcp_request_ct 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", @@ -11673,14 +11689,33 @@ class TestConnectChallengeResolver: assert exc.value.status_code == 401 @pytest.mark.asyncio - @pytest.mark.parametrize("route_name", ["obo", "obo_server"], ids=["alias_route", "server_name_route"]) - async def test_obo_challenge_www_authenticate_matches_main_byte_for_byte(self, monkeypatch, route_name): - """The provider redesign must not change what an OBO server challenges with: the relative - RFC 9728 resource_metadata path naming the configured alias whichever route the client used, - plus the RFC 6750 invalid_token triple, exactly as main.""" + @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( @@ -11691,7 +11726,7 @@ class TestConnectChallengeResolver: 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": []}, + scope=_connect_scope(path), mcp_servers=[route_name], oauth2_headers=None, mcp_server_auth_headers=None, @@ -11701,11 +11736,82 @@ class TestConnectChallengeResolver: assert exc.value.status_code == 401 assert (exc.value.headers or {}).get("WWW-Authenticate") == ( - 'Bearer resource_metadata="/.well-known/oauth-protected-resource/mcp/obo", ' + 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, + ) + 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", 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 c57a3b42975..ff94d311f67 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/test_agent_365.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/test_agent_365.py @@ -3,6 +3,7 @@ import time import uuid from types import SimpleNamespace from typing import Any, Final +from unittest.mock import patch import httpx import pytest @@ -23,6 +24,10 @@ from litellm.proxy._experimental.mcp_server.caller_sign_in import ( ) 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 ( + _post_exchange_endpoint, +) +from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchanger import OboTokenExchanger from litellm.proxy._experimental.mcp_server.outbound_credentials.types import ( CredError, ServerSpec, @@ -62,6 +67,25 @@ def _response(status_code: int, payload: Any = None, text: str | None = None) -> return httpx.Response(status_code=status_code, text=text or "", request=request) +_HTTP_CLIENT: Final = "litellm.llms.custom_httpx.http_handler.get_async_httpx_client" + + +def _entra_rejecting_with(body: dict[str, object]) -> object: + """An httpx client whose token POST raises the HTTPStatusError the real exchanger classifies.""" + request: Final = httpx.Request("POST", TOKEN_URL) + response: Final = httpx.Response(401, json=body, request=request) + + class _Resp: + def raise_for_status(self) -> None: + raise httpx.HTTPStatusError("unauthorized", request=request, response=response) + + class _Client: + async def post(self, *args: object, **kwargs: object) -> _Resp: + return _Resp() + + return _Client() + + 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 @@ -764,6 +788,37 @@ class TestUnreachableFallback: assert exc_info.value.status_code == 401 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.""" + exchanger: Final = OboTokenExchanger(_post_exchange_endpoint) + handler: Final = FakeHandler([]) + guardrail: Final = _make_guardrail(handler, exchanger=exchanger, unreachable_fallback="fail_open") + data: Final = _mcp_data() + with ( + patch( + _HTTP_CLIENT, + return_value=_entra_rejecting_with( + { + "error": "invalid_client", + "error_description": "AADSTS5002723: Invalid JWT token. Token is not well formed.", + "error_codes": [5002723], + } + ), + ), + 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 "client_secret" not in info["guardrail_response"]["reason"] + assert handler.calls == [] + @pytest.mark.asyncio async def test_evaluate_4xx_blocks_even_fail_open(self): handler: Final = FakeHandler([_response(403, text="obo token lacks the scope")]) @@ -1218,6 +1273,19 @@ class TestCallerSignIn: 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 @@ -1240,7 +1308,7 @@ class TestCallerSignIn: 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_obo_server_with_provider_advertises_both_issuers_and_scopes(self, monkeypatch): + 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) @@ -1256,7 +1324,7 @@ class TestCallerSignIn: "https://jwt-idp.test", "https://login.microsoftonline.com/tenant-abc/v2.0", ) - assert sign_in.scopes == ("read", "api://client-xyz/access_as_user") + assert sign_in.scopes == ("read",) class TestPreflightCallerSignIn: @@ -1284,6 +1352,20 @@ class TestPreflightCallerSignIn: assert verdict == Rejected(detail="the provided assertion has expired", claims="step-up") + @pytest.mark.asyncio + async def test_malformed_assertion_is_rejected_at_connect_even_fail_open(self): + exchanger: Final = OboTokenExchanger(_post_exchange_endpoint) + guardrail: Final = _make_guardrail(FakeHandler([]), exchanger=exchanger, unreachable_fallback="fail_open") + + with patch( + _HTTP_CLIENT, + return_value=_entra_rejecting_with({"error": "invalid_client", "error_codes": [5002723]}), + ): + verdict: Final = await guardrail.preflight_caller_sign_in(_server(), _user(), FAKE_ASSERTION) + + assert isinstance(verdict, Rejected) + assert "client_secret" not in verdict.detail + @pytest.mark.asyncio async def test_misconfigured_fail_closed_is_unavailable(self): exchanger: Final = StubTokenExchanger([Error(CredError.of_misconfigured("bad client_secret"))]) From 93f3fbe565fa55b0e633fda1afebd748c1abff1c Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 2 Oct 2026 20:53:00 +0000 Subject: [PATCH 32/51] test(mcp): register the configured scopes through credentials in the Agent 365 PRM integration test Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- tests/integration/mcp/test_mcp_caller_sign_in.py | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/tests/integration/mcp/test_mcp_caller_sign_in.py b/tests/integration/mcp/test_mcp_caller_sign_in.py index 53cec8a61e3..ae01857cc91 100644 --- a/tests/integration/mcp/test_mcp_caller_sign_in.py +++ b/tests/integration/mcp/test_mcp_caller_sign_in.py @@ -364,7 +364,12 @@ def test_agent_365_prm_advertises_the_servers_configured_scopes(gateway: Gateway candidate.scenario() as scenario, ): scoped: Final = "a365" + uuid.uuid4().hex[:8] - register_mcp(scenario, peer, scoped, scopes=["https://example/mcp/scoped/access_as_user", "offline_access"]) + 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) From 71b946cd99edd5f3b71a452f0fa77662c3c73643 Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 2 Oct 2026 20:53:25 +0000 Subject: [PATCH 33/51] refactor(mcp): drop the explanatory comment on the AADSTS assertion prefix Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/outbound_credentials/token_exchange_provider.py | 2 -- 1 file changed, 2 deletions(-) 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 ad503996155..7c68224ab1a 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 @@ -35,8 +35,6 @@ 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"} ) -# Entra reports a forged or garbled assertion as ``invalid_client`` with an AADSTS50027xx sub-code, -# the same top-level code as a bad gateway secret; the sub-code is what says the caller has to fix it. _INVALID_ASSERTION_AADSTS_PREFIX: Final = "50027" From ae7b0fb8ea0ffc8033ccbbbe6e936027cf0efaf3 Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 2 Oct 2026 21:52:54 +0000 Subject: [PATCH 34/51] fix(mcp): retry a scoped name as an access group when the server it names is ungranted _scoped_server returned a denial whenever the registry placed a scoped name on a server the caller does not hold, so a key granted only access group docs lost the group's servers when an ungranted server was also named docs. It now returns None there, as it does for a name hidden from the caller's IP, and the caller retries the name as an access group it holds, which is what the merge base did. Regression tests pin the moved oauth_delegate and gateway oauth2 route shapes (alias case, server id, aggregate x-mcp-servers, authorization-server document, authorize relay) to the exact-name route, and the unselected aggregate connect to no challenge Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../_experimental/mcp_server/operations.py | 36 +++---- .../integration/mcp/test_mcp_access_matrix.py | 31 ++++++ .../mcp_server/test_discoverable_endpoints.py | 79 +++++++++++++++ .../test_mcp_server_tool_calls_and_headers.py | 97 ++++++++++++++++--- 4 files changed, 206 insertions(+), 37 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index 283b139320c..9f31263e38e 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -6,7 +6,7 @@ import types import uuid from collections.abc import Mapping, Sequence from datetime import datetime -from typing import Any, Final, Literal, NoReturn, TypeAlias, overload +from typing import Any, Final, NoReturn, TypeAlias, overload from fastapi import HTTPException from mcp import ReadResourceResult, Resource @@ -453,9 +453,7 @@ async def _get_allowed_mcp_servers_from_mcp_server_names( if mcp_servers is not None: for server_or_group in mcp_servers: scoped = _scoped_server(server_or_group, allowed_mcp_servers, client_ip) - if isinstance(scoped, str): - verbose_logger.debug("MCP scope name %s names a server the caller does not hold", server_or_group) - elif scoped is not None: + if scoped is not None: filtered_server[scoped.server_id] = scoped else: try: @@ -494,25 +492,17 @@ def _server_answers_to(server: MCPServer, name: str) -> bool: return server_answers_to_name(server, name) -def _scoped_server( - name: str, allowed_mcp_servers: Sequence[MCPServer], client_ip: str | None -) -> MCPServer | Literal["denied"] | None: - """The server a scoped ``name`` selects for the caller, in this order. ``"denied"`` when the registry's - ``get_mcp_server_answering_to`` pick, made with the same ``client_ip`` the connect preflight and discovery - use, is a server hidden from that IP. Otherwise the granted server answering to ``name``: 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. With no granted server answering: ``"denied"`` when the registry - places the name on a server the caller does not hold, so it is not retried as an access group; ``None`` - when the registry cannot place the name for any caller.""" - registry_pick: Final = global_mcp_server_manager.get_mcp_server_answering_to(name, client_ip=client_ip) - if registry_pick is None and global_mcp_server_manager.get_mcp_server_answering_to(name) is not None: - return "denied" - granted: Final = global_mcp_server_manager.get_mcp_server_answering_to( - name, client_ip=client_ip, among=allowed_mcp_servers - ) - if granted is not None: - return granted - return "denied" if registry_pick is not None else None +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( diff --git a/tests/integration/mcp/test_mcp_access_matrix.py b/tests/integration/mcp/test_mcp_access_matrix.py index 13759ce245c..d199a3c34ff 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 @@ -79,6 +80,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/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index ccf968e8179..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)""" 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 258d24ff619..4f8578e917b 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 @@ -8723,33 +8723,36 @@ async def test_scoped_router_selects_the_server_the_connect_preflight_resolves(a @pytest.mark.asyncio -async def test_scoped_name_of_an_ungranted_server_is_not_retried_as_an_access_group(): +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 - private = MCPServer(server_id="p-id", name="p", server_name="p", alias="shared", transport=MCPTransport.http) + 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({"p-id": private, "m-id": member}) + 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, - return_value=["m-id"], + side_effect=lambda names: [sid for name in names for sid in group_members.get(name, [])], ) as groups: - denied = await _get_allowed_mcp_servers_from_mcp_server_names( - mcp_servers=["shared"], allowed_mcp_servers=[member] + collided = await _get_allowed_mcp_servers_from_mcp_server_names( + mcp_servers=["docs"], allowed_mcp_servers=[member] ) - unknown = await _get_allowed_mcp_servers_from_mcp_server_names( + 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 denied == [], "a denied server name must not widen to an access group of the same name" - assert [s.server_id for s in unknown] == ["m-id"] - assert groups.await_args_list == [call(["team"])] + 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 @@ -8801,7 +8804,6 @@ async def test_scoped_router_hides_a_private_server_from_an_external_ip_like_the 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"), - pytest.param(("p-id", "m-id"), "shared", ("m-id",), [], id="ungranted-only-name-stays-denied"), ], ) async def test_scoped_router_prefers_the_granted_server_answering_to_the_name(registry, scope, granted, expected): @@ -8815,8 +8817,6 @@ async def test_scoped_router_prefers_the_granted_server_answering_to_the_name(re "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), - "p-id": MCPServer(server_id="p-id", name="p", server_name="p", alias="shared", transport=MCPTransport.http), - "m-id": 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({server_id: servers[server_id] for server_id in registry}) @@ -8825,7 +8825,7 @@ async def test_scoped_router_prefers_the_granted_server_answering_to_the_name(re "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp." "MCPRequestHandler._get_mcp_servers_from_access_groups", new_callable=AsyncMock, - return_value=["m-id"], + 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] @@ -10401,6 +10401,75 @@ class TestPreemptive401ModeAware: client_ip=None, ) + 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, + ) + 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 + 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 + + with patch.object(manager, "has_user_oauth_token", new_callable=AsyncMock, return_value=False) as tokens: + 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, + ) + + assert tokens.await_count == 0, "an unselected aggregate connect must not probe any server for a token" + @pytest.mark.asyncio async def test_deferred_discovery_runs_before_delegate_challenge(self): from litellm.proxy._experimental.mcp_server import server as server_module From 5a7ec362a0af8f08ea0edb77109188efbe48db78 Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 2 Oct 2026 22:04:18 +0000 Subject: [PATCH 35/51] test(mcp): assert the aggregate connect passes unchallenged instead of counting token probes The unselected aggregate connect test asserted only on a patched has_user_oauth_token call count, which the test-quality gate flags as mock-echo (TQ002). It now calls the preflight unpatched and asserts it returns without a sign-in challenge, so the mutant that probes every registered server is still killed by the 401 it would raise. Also applies ruff format to test_caller_sign_in.py, a file new in this PR Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/test_caller_sign_in.py | 4 +--- .../test_mcp_server_tool_calls_and_headers.py | 19 +++++++++---------- 2 files changed, 10 insertions(+), 13 deletions(-) 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 index 0e66500fad0..83371400f4c 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_caller_sign_in.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_caller_sign_in.py @@ -105,9 +105,7 @@ def test_provider_returning_none_contributes_nothing(registered): 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 - ) + 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",)) 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 4f8578e917b..673c02af603 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 @@ -10458,17 +10458,16 @@ class TestPreemptive401ModeAware: server = _make_oauth2_server(alias, oauth2_flow="authorization_code", delegate_auth_to_upstream=delegate) manager.registry[server.server_id] = server - with patch.object(manager, "has_user_oauth_token", new_callable=AsyncMock, return_value=False) as tokens: - 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, - ) + 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, + ) - assert tokens.await_count == 0, "an unselected aggregate connect must not probe any server for a token" + 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 01fbd9d6a1fc6105357d3e84c421ead293d84954 Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 2 Oct 2026 22:22:03 +0000 Subject: [PATCH 36/51] fix(mcp): connect preflight resolves the granted server by name so an access group member is not challenged for a colliding plain server Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../proxy/_experimental/mcp_server/server.py | 4 +- .../test_mcp_server_tool_calls_and_headers.py | 51 +++++++++++++++++++ 2 files changed, 54 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index d0efbf7c0fb..4f0e4465ef3 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -1633,7 +1633,9 @@ if MCP_AVAILABLE: if registry_pick and mcp_servers is not None and len(mcp_servers) == 1 else () ) - granted = next(iter(allowed_single), None) + granted = operations.global_mcp_server_manager.get_mcp_server_answering_to( + server_name, client_ip=client_ip, among=allowed_single + ) server = granted if granted is not None else registry_pick granted_single = granted is not None obo_without_subject = ( 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 673c02af603..bfca3a74aad 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 @@ -9084,6 +9084,57 @@ class TestConnectPreflightRoutesLikeTheScopedRouter: 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"}, + ) + + 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(): From 4f2b36323f00eeea8ae69d7ad01d1986ac2b71a2 Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 2 Oct 2026 22:40:49 +0000 Subject: [PATCH 37/51] fix(mcp): connect preflight skips the granted lookup when the caller has no single-server grant Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/_experimental/mcp_server/server.py | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 4f0e4465ef3..861fe36eacc 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -1633,8 +1633,12 @@ if MCP_AVAILABLE: 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 + 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 From 6b8d3cbf155aaca047a92c47b27e965a6965033c Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 2 Oct 2026 22:52:41 +0000 Subject: [PATCH 38/51] test(mcp): assert the connect challenge bytes on the //mcp entry point Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp/test_mcp_agent_365_guardrail.py | 26 ++++++++++++------- 1 file changed, 17 insertions(+), 9 deletions(-) diff --git a/tests/integration/mcp/test_mcp_agent_365_guardrail.py b/tests/integration/mcp/test_mcp_agent_365_guardrail.py index 6f8aea462f8..f3c923cbf10 100644 --- a/tests/integration/mcp/test_mcp_agent_365_guardrail.py +++ b/tests/integration/mcp/test_mcp_agent_365_guardrail.py @@ -38,6 +38,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,17 +140,21 @@ 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) - malformed: Final = rig.caller(entry, "not-a-jws").call( - f"{rig.alias}-add", {"entry": entry}, server_id=rig.server_id - ) - for label, outcome in (("without a bearer", missing), ("opaque bearer", malformed)): - assert outcome.error is not None, f"{entry} {label}: {outcome.raw}" + for label, bearer in (("without a bearer", None), ("opaque bearer", "not-a-jws")): + caller: Final = rig.caller(entry, bearer) if entry == "server_mcp": - assert outcome.status == 401, f"{entry} {label} skips the connect sign-in challenge: {outcome.raw}" - else: - assert REJECTED in outcome.raw, f"{entry} {label}: {outcome.raw}" + challenged: Final = caller.rpc( + "tools/call", {"name": f"{rig.alias}-add", "arguments": {"entry": entry}} + ) + 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}" + 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) - 1) assert rig.guardrail_statuses("call_mcp_tool", expected) == ["guardrail_intervened"] * expected From 3e7efe95d6bf822936b8260b8dc5f69a27409c22 Mon Sep 17 00:00:00 2001 From: yucheng Date: Sat, 3 Oct 2026 00:29:25 +0000 Subject: [PATCH 39/51] fix(mcp): bound the Agent 365 Entra exchange by request_timeout and cover moved OBO and passthrough connect shapes Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../token_exchange_provider.py | 13 +++++-- .../guardrail_hooks/agent_365/agent_365.py | 6 ++- .../test_token_exchange_provider.py | 38 +++++++++++++++++-- .../test_mcp_server_tool_calls_and_headers.py | 37 ++++++++++++++++++ .../guardrail_hooks/test_agent_365.py | 31 +++++++++++++++ 5 files changed, 116 insertions(+), 9 deletions(-) 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 7c68224ab1a..a5f41e540e2 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 @@ -81,7 +81,7 @@ def _oauth_error_fields(response: httpx.Response) -> _OAuthErrorBody: async def _post_exchange_endpoint( - url: str, form: dict[str, str], client_auth_headers: dict[str, str] + url: str, form: dict[str, str], client_auth_headers: dict[str, str], *, timeout: float | None = None ) -> dict[str, object] | None: from litellm.llms.custom_httpx.http_handler import ( # noqa: PLC0415 get_async_httpx_client, # pyright: ignore @@ -95,7 +95,9 @@ 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, timeout=timeout + ) response.raise_for_status() # pyright: ignore parsed: Final[object] = response.json() # pyright: ignore except httpx.HTTPStatusError as status_err: @@ -133,9 +135,12 @@ async def _post_exchange_endpoint( return parsed # pyright: ignore -def build_token_exchanger() -> OboTokenExchanger: +def build_token_exchanger(*, request_timeout: float | None = None) -> OboTokenExchanger: + async def post(url: str, form: dict[str, str], client_auth_headers: dict[str, str]) -> dict[str, object] | None: + return await _post_exchange_endpoint(url, form, client_auth_headers, timeout=request_timeout) + 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/guardrails/guardrail_hooks/agent_365/agent_365.py b/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py index ac09821feb3..7c374bd8306 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py +++ b/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py @@ -185,7 +185,11 @@ class Agent365Guardrail(CustomGuardrail): 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() + self._token_exchanger: Final = ( + token_exchanger + if token_exchanger is not None + else build_token_exchanger(request_timeout=self.request_timeout) + ) verbose_proxy_logger.info("Initialized Microsoft Agent 365 guardrail: %s", guardrail_name) @staticmethod 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 589c9ce16f5..b18a0ae48de 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 @@ -9,7 +9,7 @@ from unittest.mock import patch import pytest from pydantic import SecretStr -from litellm.proxy._experimental.mcp_server.outbound_credentials import Error, ServerSpec +from litellm.proxy._experimental.mcp_server.outbound_credentials import Error, Ok, ServerSpec from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchange_provider import ( _post_exchange_endpoint, build_token_exchanger, @@ -38,7 +38,7 @@ def _client_raising_status(status: int, body: object): 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() @@ -53,6 +53,36 @@ def test_build_gives_each_caller_an_independent_cache(): assert build_token_exchanger() is not build_token_exchanger() +def _recording_client(seen: list[float | None]): + class _Resp: + def raise_for_status(self) -> None: + return None + + def json(self) -> dict[str, object]: + return {"access_token": "x", "expires_in": 60} + + class _Client: + async def post(self, url, headers, data, timeout=None): + seen.append(timeout) + return _Resp() + + return _Client() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("request_timeout", [0.5, None], ids=["bounded", "handler_default"]) +async def test_built_exchanger_posts_with_the_configured_request_timeout(request_timeout): + seen: list[float | None] = [] + 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) + with patch(_HTTP_CLIENT, return_value=_recording_client(seen)): + result = await build_token_exchanger(request_timeout=request_timeout).exchange("jwt", server, config) + assert isinstance(result, Ok) + assert seen == [request_timeout] + + @pytest.mark.asyncio async def test_post_returns_none_on_transport_error(): with patch(_HTTP_CLIENT, side_effect=RuntimeError("boom")): @@ -70,7 +100,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()): @@ -142,7 +172,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()): 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 bfca3a74aad..ed94de2b90e 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 @@ -10499,6 +10499,43 @@ class TestPreemptive401ModeAware: 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 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 ff94d311f67..d5d0fa6c812 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/test_agent_365.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/test_agent_365.py @@ -1352,6 +1352,37 @@ class TestPreflightCallerSignIn: 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): + seen: Final[list[object]] = [] + + class _Resp: + def raise_for_status(self) -> None: + return None + + def json(self) -> dict[str, object]: + return {"access_token": "exchanged", "expires_in": 3600} + + class _Client: + async def post(self, *args: object, **kwargs: object) -> _Resp: + seen.append(kwargs.get("timeout")) + return _Resp() + + guardrail: Final = Agent365Guardrail( + guardrail_name="a365", + tenant_id="tenant-abc", + client_id="cid", + client_secret="csecret", + request_timeout=0.5, + async_handler=FakeHandler([]), + ) + + with patch(_HTTP_CLIENT, return_value=_Client()): + verdict: Final = await guardrail.preflight_caller_sign_in(_server(), _user(), FAKE_ASSERTION) + + assert verdict == SignedIn() + assert seen == [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): exchanger: Final = OboTokenExchanger(_post_exchange_endpoint) From 7603412660b3455239ede2300feedc3b7f92a8fc Mon Sep 17 00:00:00 2001 From: yucheng Date: Sat, 3 Oct 2026 05:34:55 +0000 Subject: [PATCH 40/51] fix(mcp): answer Agent 365 exchange faults inside the tools/call envelope with the merge base's reasons The connect-time sign-in preflight raised its fail-closed 503 on every method of a single-server route, so a post-connect tools/call whose Entra exchange hit a gateway fault got a bare 503 with no JSON-RPC envelope, no hook run and no guardrail Logs row. The gate now reads the body before the preflight and only raises the 503 while the request is an initialize (or the SSE GET); any other request reaches the tool-call hook, which answers 200 isError with the verdict and writes the failure row as the merge base did The guardrail posts to the Entra token endpoint through its own handler again, bounded by request_timeout, and classifies the answer itself through the shared OAuth error reader, so the gateway-fault reason keeps the OAuth code (invalid_client, invalid_scope, ...), a malformed caller assertion (AADSTS 50027xx) stays a 401 caller fault reading rejected (invalid_client), and transport, HTTP 500, non-JSON and missing access_token answers read as the merge base's distinct reasons instead of one collapsed sentence. build_token_exchanger takes the HTTP post as a dependency in place of request_timeout Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/caller_sign_in.py | 11 +- .../token_exchange_provider.py | 22 +- .../proxy/_experimental/mcp_server/server.py | 41 +-- .../guardrail_hooks/agent_365/agent_365.py | 114 +++++++-- .../test_token_exchange_provider.py | 33 +-- .../test_mcp_server_tool_calls_and_headers.py | 110 +++++++- .../guardrail_hooks/test_agent_365.py | 239 ++++++++++++------ 7 files changed, 406 insertions(+), 164 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/caller_sign_in.py b/litellm/proxy/_experimental/mcp_server/caller_sign_in.py index 1d417c01dbd..74637972f75 100644 --- a/litellm/proxy/_experimental/mcp_server/caller_sign_in.py +++ b/litellm/proxy/_experimental/mcp_server/caller_sign_in.py @@ -162,9 +162,12 @@ async def preflight_caller_sign_in( *, root_path: str, resource_metadata: str | None, + connecting: bool, ) -> None: """Run every provider's connect-time check against the subject token, so a bearer the IdP will - reject surfaces as a challenge here rather than a JSON-RPC error at the first tool call.""" + reject surfaces as a challenge here rather than a JSON-RPC error at the first tool call. A fail-closed + provider outage is the connect's 503 only while ``connecting``; on an open session the tool-call hook + answers it 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 @@ -181,9 +184,11 @@ async def preflight_caller_sign_in( raise_token_exchange_challenge( server, root_path=root_path, claims=claims, resource_metadata=resource_metadata ) - case Unavailable(detail=detail, fail_open=True): + case Unavailable(fail_open=True): continue - case Unavailable(detail=detail, fail_open=False): + case Unavailable(detail=detail, fail_open=False) if connecting: raise HTTPException(status_code=503, detail=detail) + case Unavailable(): + continue case _ as verdict: assert_never(verdict) 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 a5f41e540e2..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 @@ -25,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, @@ -39,7 +40,7 @@ _INVALID_ASSERTION_AADSTS_PREFIX: Final = "50027" @dataclass(frozen=True, slots=True) -class _OAuthErrorBody: +class OAuthErrorBody: error: str | None claims: str | None error_codes: tuple[str, ...] @@ -53,7 +54,7 @@ class _OAuthErrorBody: return self.error -def _oauth_error_fields(response: httpx.Response) -> _OAuthErrorBody: +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. @@ -65,13 +66,13 @@ def _oauth_error_fields(response: httpx.Response) -> _OAuthErrorBody: try: body: Final[object] = response.json() except Exception: # noqa: BLE001 - return _OAuthErrorBody(error=None, claims=None, error_codes=()) + return OAuthErrorBody(error=None, claims=None, error_codes=()) if not isinstance(body, dict): - return _OAuthErrorBody(error=None, claims=None, error_codes=()) + return OAuthErrorBody(error=None, claims=None, error_codes=()) code: Final = body.get("error") claims: Final = body.get("claims") raw_codes: Final = body.get("error_codes") - return _OAuthErrorBody( + 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))) @@ -81,7 +82,7 @@ def _oauth_error_fields(response: httpx.Response) -> _OAuthErrorBody: async def _post_exchange_endpoint( - url: str, form: dict[str, str], client_auth_headers: dict[str, str], *, timeout: float | None = None + url: str, form: dict[str, str], client_auth_headers: dict[str, str] ) -> dict[str, object] | None: from litellm.llms.custom_httpx.http_handler import ( # noqa: PLC0415 get_async_httpx_client, # pyright: ignore @@ -96,7 +97,7 @@ async def _post_exchange_endpoint( try: client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.MCP) # pyright: ignore response: Final = await client.post( # pyright: ignore[reportUnknownMemberType] # untyped handler - url, headers=headers, data=form, timeout=timeout + url, headers=headers, data=form ) response.raise_for_status() # pyright: ignore parsed: Final[object] = response.json() # pyright: ignore @@ -108,7 +109,7 @@ async def _post_exchange_endpoint( verbose_logger.warning("MCP token exchange throttled or timed out (HTTP %d)", status_code) return None if 400 <= status_code < 500: - oauth_error: Final = _oauth_error_fields(status_err.response) + 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( @@ -135,10 +136,7 @@ async def _post_exchange_endpoint( return parsed # pyright: ignore -def build_token_exchanger(*, request_timeout: float | None = None) -> OboTokenExchanger: - async def post(url: str, form: dict[str, str], client_auth_headers: dict[str, str]) -> dict[str, object] | None: - return await _post_exchange_endpoint(url, form, client_auth_headers, timeout=request_timeout) - +def build_token_exchanger(*, post: ExchangeHttpPost = _post_exchange_endpoint) -> OboTokenExchanger: return OboTokenExchanger( post, cache=InMemoryTokenCacheBackend(max_size=MCP_TOKEN_EXCHANGE_CACHE_MAX_SIZE), diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 861fe36eacc..42ef852c932 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -1608,6 +1608,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: bool, allowed_server_ids: set[str] | None = None, raw_headers: Mapping[str, str] | None = None, ) -> None: @@ -1779,6 +1780,7 @@ if MCP_AVAILABLE: 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 @@ -2070,6 +2072,22 @@ 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) + consumed_messages, body = ( + await _read_request_body_for_routing(receive) if scope.get("method") == "POST" else ([], b"") + ) + is_initialize: Final = _is_initialize_request(body) + + # 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 + # 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 @@ -2082,6 +2100,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=is_initialize, allowed_server_ids=toolset_allowed_server_ids, raw_headers=raw_headers, ) @@ -2117,8 +2136,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 @@ -2126,8 +2143,7 @@ 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) @@ -2156,11 +2172,6 @@ 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) - use_stateful: Final = bool(session_id or is_initialize) target_manager: Final = session_manager_stateful if use_stateful else session_manager_stateless @@ -2189,17 +2200,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. # @@ -2429,6 +2429,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=scope["method"] == "GET", allowed_server_ids=toolset_allowed_server_ids, raw_headers=raw_headers, ) 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 7c374bd8306..dbc041ec806 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py +++ b/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py @@ -42,8 +42,12 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_sto 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.token_exchanger import TokenExchanger from litellm.proxy._experimental.mcp_server.outbound_credentials.types import ( CredError, ServerSpec, @@ -72,12 +76,12 @@ MCP_SESSION_ID_HEADER: Final = "mcp-session-id" DEFENDER_STATUS_EVALUATED: Final = "Evaluated" 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]) +_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 @@ -130,12 +134,32 @@ class _BlockedDetail(TypedDict): correlation_id: ReadOnly[str | None] +class Agent365TokenExchangeError(Exception): + """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.""" + + def __init__(self, error_code: str) -> None: + super().__init__(error_code) + self.error_code = error_code + + +class Agent365MalformedResponseError(Exception): + pass + + class Agent365ThrottledError(Exception): def __init__(self, status_code: int) -> None: super().__init__(f"HTTP {status_code}") 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. @@ -188,7 +212,7 @@ class Agent365Guardrail(CustomGuardrail): self._token_exchanger: Final = ( token_exchanger if token_exchanger is not None - else build_token_exchanger(request_timeout=self.request_timeout) + else build_token_exchanger(post=self._post_entra_token_endpoint) ) verbose_proxy_logger.info("Initialized Microsoft Agent 365 guardrail: %s", guardrail_name) @@ -230,12 +254,25 @@ class Agent365Guardrail(CustomGuardrail): try: exchange_result: Final = await self._exchange_caller_assertion(assertion) + except Agent365TokenExchangeError as exc: + return self._handle_unavailable( + data=data, tool_name=tool_name, reason=_gateway_fault_reason(exc.error_code) + ) + except Agent365ThrottledError as exc: + self._handle_throttled( + data=data, + tool_name=tool_name, + reason=f"the Entra token endpoint returned HTTP {exc.status_code}", + latency_ms=None, + ) except (httpx.HTTPError, LitellmTimeout, TimeoutError) as exc: return self._handle_unavailable( data=data, tool_name=tool_name, reason=f"the Entra token endpoint could not be reached ({type(exc).__name__})", ) + except Agent365MalformedResponseError as exc: + return self._handle_unavailable(data=data, tool_name=tool_name, reason=str(exc)) match exchange_result: case Ok(token): obo_token: Final = token.access_token @@ -248,15 +285,6 @@ class Agent365Guardrail(CustomGuardrail): status_code=401, reason=f"the Entra On-Behalf-Of token exchange was rejected ({error.unauthorized.detail})", ) - case "misconfigured": - return self._handle_unavailable( - data=data, - tool_name=tool_name, - reason=( - f"Entra rejected the gateway's own Agent 365 credentials ({error.misconfigured}); " - "check the guardrail's client_id and client_secret" - ), - ) case _: return self._handle_unavailable( data=data, @@ -487,6 +515,48 @@ class Agent365Guardrail(CustomGuardrail): 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=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) + if response.status_code >= 500: + raise httpx.HTTPStatusError( + f"Entra token endpoint returned {response.status_code}", + request=response.request, + response=response, + ) + try: + parsed_body: Final[object] = response.json() + except ValueError as exc: + raise Agent365MalformedResponseError("the Entra token endpoint returned a non-JSON body") from exc + 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: + 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") + 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") + return body + async def preflight_caller_sign_in( self, server: MCPServer, user_api_key_auth: "UserAPIKeyAuth | None", subject_token: str ) -> CallerSignInPreflight: @@ -496,25 +566,25 @@ class Agent365Guardrail(CustomGuardrail): 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=self.unreachable_fallback == "fail_open", + 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) - detail: Final = ( - f"Entra rejected the gateway's own Agent 365 credentials ({error.misconfigured}); " - "check the guardrail's client_id and client_secret" - if error.tag == "misconfigured" - else f"the Entra token exchange failed ({error.summary})" - ) - return Unavailable(detail=detail, fail_open=self.unreachable_fallback == "fail_open") + return Unavailable(detail=f"the Entra token exchange failed ({error.summary})", fail_open=fail_open) async def _post_allowing_error_status( self, 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 b18a0ae48de..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 @@ -53,34 +53,23 @@ def test_build_gives_each_caller_an_independent_cache(): assert build_token_exchanger() is not build_token_exchanger() -def _recording_client(seen: list[float | None]): - class _Resp: - def raise_for_status(self) -> None: - return None - - def json(self) -> dict[str, object]: - return {"access_token": "x", "expires_in": 60} - - class _Client: - async def post(self, url, headers, data, timeout=None): - seen.append(timeout) - return _Resp() - - return _Client() - - @pytest.mark.asyncio -@pytest.mark.parametrize("request_timeout", [0.5, None], ids=["bounded", "handler_default"]) -async def test_built_exchanger_posts_with_the_configured_request_timeout(request_timeout): - seen: list[float | None] = [] +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) - with patch(_HTTP_CLIENT, return_value=_recording_client(seen)): - result = await build_token_exchanger(request_timeout=request_timeout).exchange("jwt", server, config) + result = await build_token_exchanger(post=post).exchange("jwt", server, config) assert isinstance(result, Ok) - assert seen == [request_timeout] + 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 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 ed94de2b90e..2d510e70bc6 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 @@ -8933,6 +8933,7 @@ class TestConnectPreflightRoutesLikeTheScopedRouter: 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=True, ) @pytest.mark.asyncio @@ -9021,6 +9022,7 @@ class TestConnectPreflightRoutesLikeTheScopedRouter: 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=True, ) @pytest.mark.asyncio @@ -9128,6 +9130,7 @@ class TestConnectPreflightRoutesLikeTheScopedRouter: 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=True, ) assert outcome is None, ( @@ -10450,6 +10453,7 @@ class TestPreemptive401ModeAware: mcp_server_auth_headers=None, user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key"), client_ip=None, + connecting=True, ) async def _connect_with_a_grant(self, server, requested: str, path: str) -> HTTPException: @@ -10472,6 +10476,7 @@ class TestPreemptive401ModeAware: mcp_server_auth_headers=None, user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key"), client_ip=None, + connecting=True, ) return exc.value @@ -10553,6 +10558,7 @@ class TestPreemptive401ModeAware: mcp_server_auth_headers=None, user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key"), client_ip=None, + connecting=True, ) assert outcome is None, "an unselected aggregate connect must pass without a sign-in challenge" @@ -10655,6 +10661,7 @@ class TestPreemptive401ModeAware: mcp_server_auth_headers=None, user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key"), client_ip=None, + connecting=True, ) assert exc.value.status_code == 401 @@ -10763,6 +10770,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=True, ) @pytest.mark.asyncio @@ -10821,6 +10829,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=True, ) assert exc.value.status_code == 401 @@ -10889,6 +10898,7 @@ class TestOboPreflightScopedToAllowedServers: "x-litellm-api-key": user_api_key_auth.api_key if user_api_key_auth else "", "authorization": self.SUBJECT_HEADERS["Authorization"], }, + connecting=True, ) return allowed_lookup, preflight @@ -10951,6 +10961,7 @@ class TestOboChallengeGateKeepsBaseConnectRules: 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=True, ) @pytest.mark.asyncio @@ -11795,6 +11806,7 @@ class TestConnectChallengeResolver: mcp_server_auth_headers=None, user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"), client_ip=None, + connecting=True, ) finally: litellm.logging_callback_manager.remove_callback_from_list_by_object( @@ -11836,6 +11848,7 @@ class TestConnectChallengeResolver: mcp_server_auth_headers=None, user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"), client_ip=None, + connecting=True, ) finally: litellm.logging_callback_manager.remove_callback_from_list_by_object( @@ -11888,6 +11901,7 @@ class TestConnectChallengeResolver: mcp_server_auth_headers=None, user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"), client_ip=None, + connecting=True, ) assert exc.value.status_code == 401 @@ -11955,6 +11969,7 @@ class TestConnectChallengeResolver: 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=True, ) finally: litellm.logging_callback_manager.remove_callback_from_list_by_object( @@ -12003,6 +12018,7 @@ class TestConnectChallengeResolver: mcp_server_auth_headers=None, user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"), client_ip=None, + connecting=True, ) finally: litellm.logging_callback_manager.remove_callback_from_list_by_object( @@ -12019,7 +12035,7 @@ 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): + async def _connect(self, route_names, guardrail, allowed, raw_headers=None, connecting=True): from litellm.proxy._experimental.mcp_server import server as server_module server = _catalog_server() @@ -12044,6 +12060,7 @@ class TestConnectSignInPreflight: 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 @@ -12091,6 +12108,97 @@ class TestConnectSignInPreflight: 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, so the tools/call + hook 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"] + 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() 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 d5d0fa6c812..0d5f5b3ad8b 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/test_agent_365.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/test_agent_365.py @@ -3,7 +3,6 @@ import time import uuid from types import SimpleNamespace from typing import Any, Final -from unittest.mock import patch import httpx import pytest @@ -24,10 +23,6 @@ from litellm.proxy._experimental.mcp_server.caller_sign_in import ( ) 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 ( - _post_exchange_endpoint, -) -from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchanger import OboTokenExchanger from litellm.proxy._experimental.mcp_server.outbound_credentials.types import ( CredError, ServerSpec, @@ -67,25 +62,6 @@ def _response(status_code: int, payload: Any = None, text: str | None = None) -> return httpx.Response(status_code=status_code, text=text or "", request=request) -_HTTP_CLIENT: Final = "litellm.llms.custom_httpx.http_handler.get_async_httpx_client" - - -def _entra_rejecting_with(body: dict[str, object]) -> object: - """An httpx client whose token POST raises the HTTPStatusError the real exchanger classifies.""" - request: Final = httpx.Request("POST", TOKEN_URL) - response: Final = httpx.Response(401, json=body, request=request) - - class _Resp: - def raise_for_status(self) -> None: - raise httpx.HTTPStatusError("unauthorized", request=request, response=response) - - class _Client: - async def post(self, *args: object, **kwargs: object) -> _Resp: - return _Resp() - - return _Client() - - 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 @@ -230,6 +206,24 @@ def _default_fallback_guardrail(handler: FakeHandler, exchanger: StubTokenExchan ) +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", @@ -793,31 +787,32 @@ class TestUnreachableFallback: """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.""" - exchanger: Final = OboTokenExchanger(_post_exchange_endpoint) - handler: Final = FakeHandler([]) - guardrail: Final = _make_guardrail(handler, exchanger=exchanger, unreachable_fallback="fail_open") - data: Final = _mcp_data() - with ( - patch( - _HTTP_CLIENT, - return_value=_entra_rejecting_with( + handler: Final = FakeHandler( + [ + _response( + 400, { "error": "invalid_client", "error_description": "AADSTS5002723: Invalid JWT token. Token is not well formed.", "error_codes": [5002723], - } - ), - ), - pytest.raises(HTTPException) as exc_info, - ): + }, + ) + ] + ) + 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 "client_secret" not in info["guardrail_response"]["reason"] - assert handler.calls == [] + 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): @@ -847,22 +842,36 @@ class TestUnreachableFallback: @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): - exchanger: Final = StubTokenExchanger([Error(CredError.of_misconfigured(error_code))]) - handler: Final = FakeHandler([]) - guardrail: Final = _make_guardrail(handler, exchanger=exchanger) + 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 async def test_caller_rejection_reason_does_not_blame_the_gateway_credentials(self): @@ -879,16 +888,16 @@ class TestUnreachableFallback: @pytest.mark.asyncio async def test_gateway_credential_rejection_follows_fail_open(self): - exchanger: Final = StubTokenExchanger([Error(CredError.of_misconfigured("invalid_client"))]) - handler: Final = FakeHandler([]) - guardrail: Final = _make_guardrail(handler, exchanger=exchanger, 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): @@ -937,6 +946,92 @@ 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): @@ -1354,48 +1449,24 @@ class TestPreflightCallerSignIn: @pytest.mark.asyncio async def test_configured_timeout_bounds_the_entra_exchange_leg(self): - seen: Final[list[object]] = [] + handler: Final = FakeHandler([_response(200, {"access_token": "exchanged", "expires_in": 3600})]) + guardrail: Final = _entra_driven_guardrail(handler, request_timeout=0.5) - class _Resp: - def raise_for_status(self) -> None: - return None - - def json(self) -> dict[str, object]: - return {"access_token": "exchanged", "expires_in": 3600} - - class _Client: - async def post(self, *args: object, **kwargs: object) -> _Resp: - seen.append(kwargs.get("timeout")) - return _Resp() - - guardrail: Final = Agent365Guardrail( - guardrail_name="a365", - tenant_id="tenant-abc", - client_id="cid", - client_secret="csecret", - request_timeout=0.5, - async_handler=FakeHandler([]), - ) - - with patch(_HTTP_CLIENT, return_value=_Client()): - verdict: Final = await guardrail.preflight_caller_sign_in(_server(), _user(), FAKE_ASSERTION) + verdict: Final = await guardrail.preflight_caller_sign_in(_server(), _user(), FAKE_ASSERTION) assert verdict == SignedIn() - assert seen == [0.5], "the Entra token POST must carry the guardrail's own request_timeout" + 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): - exchanger: Final = OboTokenExchanger(_post_exchange_endpoint) - guardrail: Final = _make_guardrail(FakeHandler([]), exchanger=exchanger, unreachable_fallback="fail_open") + handler: Final = FakeHandler([_response(401, {"error": "invalid_client", "error_codes": [5002723]})]) + guardrail: Final = _entra_driven_guardrail(handler, unreachable_fallback="fail_open") - with patch( - _HTTP_CLIENT, - return_value=_entra_rejecting_with({"error": "invalid_client", "error_codes": [5002723]}), - ): - verdict: Final = await guardrail.preflight_caller_sign_in(_server(), _user(), FAKE_ASSERTION) + verdict: Final = await guardrail.preflight_caller_sign_in(_server(), _user(), FAKE_ASSERTION) - assert isinstance(verdict, Rejected) - assert "client_secret" not in verdict.detail + assert verdict == Rejected(detail="invalid_client", claims=None) @pytest.mark.asyncio async def test_misconfigured_fail_closed_is_unavailable(self): From 10c9fd95b218832429a88e8650aacb8ce344e130 Mon Sep 17 00:00:00 2001 From: yucheng Date: Sat, 3 Oct 2026 06:27:39 +0000 Subject: [PATCH 41/51] fix(mcp): enforce session ownership before reading a session-bearing POST body A POST that names an existing mcp-session-id skips the connect-time body peek, so another caller's request is refused with 403 before any body byte is awaited. Session-bearing bodies are read after the owner check and replayed to the session manager unchanged Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../proxy/_experimental/mcp_server/server.py | 31 +++-- .../test_mcp_server_tool_calls_and_headers.py | 122 ++++++++++++++++++ 2 files changed, 143 insertions(+), 10 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 42ef852c932..9ec3eed32b5 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -2072,21 +2072,23 @@ 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) - consumed_messages, body = ( - await _read_request_body_for_routing(receive) if scope.get("method") == "POST" else ([], b"") + session_header_present: Final = _get_session_id_from_scope(scope) is not None + consumed_messages, connect_body = ( + await _read_request_body_for_routing(receive) + if scope.get("method") == "POST" and not session_header_present + else ([], b"") ) - is_initialize: Final = _is_initialize_request(body) + connecting: Final = _is_initialize_request(connect_body) # 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() + async def wrapped_receive(): + if consumed_messages: + return consumed_messages.pop(0) + return await original_receive() - receive = wrapped_receive + receive = wrapped_receive # https://datatracker.ietf.org/doc/html/rfc9728#name-www-authenticate-response # Must run after toolset scoping so the challenge set is derived @@ -2100,7 +2102,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=is_initialize, + connecting=connecting, allowed_server_ids=toolset_allowed_server_ids, raw_headers=raw_headers, ) @@ -2172,6 +2174,15 @@ if MCP_AVAILABLE: return session_id = _get_session_id_from_scope(scope) + session_messages, session_body = ( + await _read_request_body_for_routing(receive) + if scope.get("method") == "POST" and session_header_present + else ([], b"") + ) + consumed_messages.extend(session_messages) + body: Final = connect_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 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 2d510e70bc6..5dc3e9898db 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 @@ -3754,6 +3754,128 @@ 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_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(): """ From bcfb081279edc1f91ca208b66c8e1f4641ec46bb Mon Sep 17 00:00:00 2001 From: yucheng Date: Sat, 3 Oct 2026 06:43:21 +0000 Subject: [PATCH 42/51] test(mcp): refuse another caller's session before its withheld body arrives The regression test holds the POST body back forever and expects the 403 for a session owned by someone else within two seconds, red on 7603412660 where the connect-time peek waited for the body first Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../test_mcp_server_tool_calls_and_headers.py | 74 +++++++++++++++++++ 1 file changed, 74 insertions(+) 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 5dc3e9898db..b58206701da 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 @@ -3815,6 +3815,80 @@ async def test_stateful_mcp_session_owner_mismatch_is_rejected_before_the_body_i 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_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 From 393853834fed8a23a02968e04c65efa11f26890c Mon Sep 17 00:00:00 2001 From: yucheng Date: Sat, 3 Oct 2026 07:04:46 +0000 Subject: [PATCH 43/51] fix(mcp): keep the fail-closed connect gate for an initialize naming a stale session Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../proxy/_experimental/mcp_server/server.py | 9 ++- .../test_mcp_server_tool_calls_and_headers.py | 79 +++++++++++++++++++ 2 files changed, 85 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 9ec3eed32b5..47d41325e74 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -2072,10 +2072,13 @@ 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) - session_header_present: Final = _get_session_id_from_scope(scope) is not None + 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_session_owners or named_session_id in _stateful_server_instances() + ) consumed_messages, connect_body = ( await _read_request_body_for_routing(receive) - if scope.get("method") == "POST" and not session_header_present + if scope.get("method") == "POST" and not names_live_session else ([], b"") ) connecting: Final = _is_initialize_request(connect_body) @@ -2176,7 +2179,7 @@ if MCP_AVAILABLE: session_messages, session_body = ( await _read_request_body_for_routing(receive) - if scope.get("method") == "POST" and session_header_present + if scope.get("method") == "POST" and names_live_session else ([], b"") ) consumed_messages.extend(session_messages) 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 b58206701da..bfd759bc2a1 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 @@ -3889,6 +3889,85 @@ async def test_stateful_mcp_session_owner_mismatch_is_refused_before_the_body_ar 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_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 From d7a2d2867294e1e62ae44cc09672498085bd02ce Mon Sep 17 00:00:00 2001 From: yucheng Date: Sat, 3 Oct 2026 09:23:03 +0000 Subject: [PATCH 44/51] fix(mcp): answer the connect challenge before reading a session-less POST body The connect gate only needs the body to tell initialize from other methods, so the peek now runs on demand through a callback and the consumed ASGI messages replay to the handler. A gated route answers 401 before the body arrives again, as the merge base did Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/caller_sign_in.py | 11 +- .../proxy/_experimental/mcp_server/server.py | 75 ++++--- .../test_mcp_server_tool_calls_and_headers.py | 191 ++++++++++++++++-- 3 files changed, 227 insertions(+), 50 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/caller_sign_in.py b/litellm/proxy/_experimental/mcp_server/caller_sign_in.py index 74637972f75..1d6d75b438f 100644 --- a/litellm/proxy/_experimental/mcp_server/caller_sign_in.py +++ b/litellm/proxy/_experimental/mcp_server/caller_sign_in.py @@ -15,7 +15,7 @@ never imports a concrete provider. from __future__ import annotations import itertools -from collections.abc import Mapping +from collections.abc import Awaitable, Callable, Mapping from dataclasses import dataclass from typing import TYPE_CHECKING, Final, Protocol, runtime_checkable @@ -162,7 +162,7 @@ async def preflight_caller_sign_in( *, root_path: str, resource_metadata: str | None, - connecting: bool, + connecting: Callable[[], Awaitable[bool]], ) -> None: """Run every provider's connect-time check against the subject token, so a bearer the IdP will reject surfaces as a challenge here rather than a JSON-RPC error at the first tool call. A fail-closed @@ -186,9 +186,8 @@ async def preflight_caller_sign_in( ) case Unavailable(fail_open=True): continue - case Unavailable(detail=detail, fail_open=False) if connecting: - raise HTTPException(status_code=503, detail=detail) - case Unavailable(): - continue + case Unavailable(detail=detail, fail_open=False): + if await connecting(): + raise HTTPException(status_code=503, detail=detail) case _ as verdict: assert_never(verdict) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 47d41325e74..4d9bdf41696 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 @@ -1426,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, @@ -1608,7 +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: bool, + connecting: Callable[[], Awaitable[bool]], allowed_server_ids: set[str] | None = None, raw_headers: Mapping[str, str] | None = None, ) -> None: @@ -2073,25 +2109,13 @@ if MCP_AVAILABLE: 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_session_owners or named_session_id in _stateful_server_instances() + names_live_session: Final = ( + named_session_id is not None and named_session_id in _stateful_server_instances() ) - consumed_messages, connect_body = ( - await _read_request_body_for_routing(receive) - if scope.get("method") == "POST" and not names_live_session - else ([], b"") + connect_peek: Final = _ConnectBodyPeek( + receive, peekable=scope.get("method") == "POST" and not names_live_session ) - connecting: Final = _is_initialize_request(connect_body) - - # Replay body messages if we consumed them for peeking - original_receive: Final = receive - - async def wrapped_receive(): - if consumed_messages: - return consumed_messages.pop(0) - return await original_receive() - - receive = wrapped_receive + 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 @@ -2105,7 +2129,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=connecting, + connecting=connect_peek.connecting, allowed_server_ids=toolset_allowed_server_ids, raw_headers=raw_headers, ) @@ -2177,13 +2201,10 @@ if MCP_AVAILABLE: return session_id = _get_session_id_from_scope(scope) - session_messages, session_body = ( - await _read_request_body_for_routing(receive) - if scope.get("method") == "POST" and names_live_session - else ([], b"") + session_body: Final = ( + await connect_peek.read() if scope.get("method") == "POST" and names_live_session else b"" ) - consumed_messages.extend(session_messages) - body: Final = connect_body or session_body + 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) @@ -2443,7 +2464,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=scope["method"] == "GET", + connecting=_known_connecting(scope["method"] == "GET"), allowed_server_ids=toolset_allowed_server_ids, raw_headers=raw_headers, ) 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 bfd759bc2a1..85b631251ef 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 @@ -53,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:] @@ -3968,6 +3972,159 @@ async def test_initialize_naming_a_stale_session_still_meets_the_fail_closed_con 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 @@ -9208,7 +9365,7 @@ class TestConnectPreflightRoutesLikeTheScopedRouter: 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=True, + connecting=_connecting, ) @pytest.mark.asyncio @@ -9297,7 +9454,7 @@ class TestConnectPreflightRoutesLikeTheScopedRouter: 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=True, + connecting=_connecting, ) @pytest.mark.asyncio @@ -9405,7 +9562,7 @@ class TestConnectPreflightRoutesLikeTheScopedRouter: 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=True, + connecting=_connecting, ) assert outcome is None, ( @@ -10728,7 +10885,7 @@ class TestPreemptive401ModeAware: mcp_server_auth_headers=None, user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key"), client_ip=None, - connecting=True, + connecting=_connecting, ) async def _connect_with_a_grant(self, server, requested: str, path: str) -> HTTPException: @@ -10751,7 +10908,7 @@ class TestPreemptive401ModeAware: mcp_server_auth_headers=None, user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key"), client_ip=None, - connecting=True, + connecting=_connecting, ) return exc.value @@ -10833,7 +10990,7 @@ class TestPreemptive401ModeAware: mcp_server_auth_headers=None, user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key"), client_ip=None, - connecting=True, + connecting=_connecting, ) assert outcome is None, "an unselected aggregate connect must pass without a sign-in challenge" @@ -10936,7 +11093,7 @@ class TestPreemptive401ModeAware: mcp_server_auth_headers=None, user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key"), client_ip=None, - connecting=True, + connecting=_connecting, ) assert exc.value.status_code == 401 @@ -11045,7 +11202,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=True, + connecting=_connecting, ) @pytest.mark.asyncio @@ -11104,7 +11261,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=True, + connecting=_connecting, ) assert exc.value.status_code == 401 @@ -11173,7 +11330,7 @@ class TestOboPreflightScopedToAllowedServers: "x-litellm-api-key": user_api_key_auth.api_key if user_api_key_auth else "", "authorization": self.SUBJECT_HEADERS["Authorization"], }, - connecting=True, + connecting=_connecting, ) return allowed_lookup, preflight @@ -11236,7 +11393,7 @@ class TestOboChallengeGateKeepsBaseConnectRules: 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=True, + connecting=_connecting, ) @pytest.mark.asyncio @@ -12081,7 +12238,7 @@ class TestConnectChallengeResolver: mcp_server_auth_headers=None, user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"), client_ip=None, - connecting=True, + connecting=_connecting, ) finally: litellm.logging_callback_manager.remove_callback_from_list_by_object( @@ -12123,7 +12280,7 @@ class TestConnectChallengeResolver: mcp_server_auth_headers=None, user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"), client_ip=None, - connecting=True, + connecting=_connecting, ) finally: litellm.logging_callback_manager.remove_callback_from_list_by_object( @@ -12176,7 +12333,7 @@ class TestConnectChallengeResolver: mcp_server_auth_headers=None, user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"), client_ip=None, - connecting=True, + connecting=_connecting, ) assert exc.value.status_code == 401 @@ -12244,7 +12401,7 @@ class TestConnectChallengeResolver: 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=True, + connecting=_connecting, ) finally: litellm.logging_callback_manager.remove_callback_from_list_by_object( @@ -12293,7 +12450,7 @@ class TestConnectChallengeResolver: mcp_server_auth_headers=None, user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"), client_ip=None, - connecting=True, + connecting=_connecting, ) finally: litellm.logging_callback_manager.remove_callback_from_list_by_object( @@ -12310,7 +12467,7 @@ 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=True): + async def _connect(self, route_names, guardrail, allowed, raw_headers=None, connecting=_connecting): from litellm.proxy._experimental.mcp_server import server as server_module server = _catalog_server() From 26715ba66766b4e98b56c19837538bcde46aa247 Mon Sep 17 00:00:00 2001 From: yucheng Date: Sat, 3 Oct 2026 09:50:50 +0000 Subject: [PATCH 45/51] fix(mcp): refuse another caller's torn-down session before peeking at its body under a fail-closed outage Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../proxy/_experimental/mcp_server/server.py | 12 +-- .../test_mcp_server_tool_calls_and_headers.py | 87 +++++++++++++++++++ 2 files changed, 94 insertions(+), 5 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 4d9bdf41696..e3fc403302a 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -2112,8 +2112,13 @@ if MCP_AVAILABLE: 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 + receive, peekable=scope.get("method") == "POST" and not names_live_session and not owner_mismatch ) receive = connect_peek.receive @@ -2174,9 +2179,7 @@ if MCP_AVAILABLE: # force-clean another caller's residual tracking entries via a # 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, @@ -2220,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." 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 85b631251ef..2cc066c6f1b 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 @@ -3893,6 +3893,93 @@ async def test_stateful_mcp_session_owner_mismatch_is_refused_before_the_body_ar 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 == ["entra.jwt.token"] + 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 From 03f68acf72d712c9f6669d81907169f5956aed4e Mon Sep 17 00:00:00 2001 From: yucheng Date: Sat, 3 Oct 2026 10:48:01 +0000 Subject: [PATCH 46/51] fix(mcp): hand the caller sign-in provider a custom-auth caller's own bearer as before the subject split Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/mcp_server_manager.py | 27 ++++++++- .../proxy/_experimental/mcp_server/server.py | 2 +- .../mcp_server/test_mcp_server_manager.py | 44 +++++++++++++++ .../test_mcp_server_tool_calls_and_headers.py | 56 ++++++++++++++++++- 4 files changed, 124 insertions(+), 5 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 9e3a9c2194e..af88f0e1eb6 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -3945,6 +3945,26 @@ 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, + user_api_key_auth: UserAPIKeyAuth | None, + ) -> str | None: + """The bearer a caller sign-in provider validates: custom auth admits the caller on its own IdP token + in ``Authorization``, so that token is the subject there and only a virtual key is withheld.""" + admitted_on_own_bearer: Final = ( + user_api_key_auth is not None + and user_api_key_auth.authenticated_by_custom_auth + and not _has_explicit_litellm_admission_header(raw_headers) + ) + if not admitted_on_own_bearer: + return MCPServerManager._extract_subject_token(oauth2_headers, raw_headers, user_api_key_auth) + bearer: Final = MCPServerManager._extract_bearer_token(oauth2_headers, raw_headers) + if bearer is None or bearer.startswith(LITELLM_VIRTUAL_KEY_PREFIX): + return None + return bearer + def _obo_subject_token( self, server: MCPServer, @@ -4262,7 +4282,10 @@ class MCPServerManager: caller_sign_in_for, ) - if subject_token is None and caller_sign_in_for(server, user_api_key_auth) is not None: + sign_in_subject: Final = self._caller_sign_in_subject_token( + oauth2_headers, raw_headers, user_api_key_auth + ) + 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 ) @@ -6019,7 +6042,7 @@ class MCPServerManager: incoming_bearer_token: Final = ( inbound_authorization[len("bearer ") :] if inbound_authorization.lower().startswith("bearer ") else None ) - incoming_subject_token: Final = self._extract_subject_token(None, raw_headers, user_api_key_auth) + incoming_subject_token: Final = self._caller_sign_in_subject_token(None, raw_headers, user_api_key_auth) pre_hook_kwargs: Final = { "guardrail_context": guardrail_context, diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index e3fc403302a..2893e044a27 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -1785,7 +1785,7 @@ if MCP_AVAILABLE: 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._extract_subject_token( # pyright: ignore[reportPrivateUsage] # the manager owns the subject/admission filter shared with the preflight + 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, user_api_key_auth ) if server is not None 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 4c3c10474bf..6fd1f384d52 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 @@ -6985,6 +6985,50 @@ class TestMCPServerManager: assert kwargs["incoming_bearer_token"] == expected_bearer assert kwargs["incoming_subject_token"] == expected_subject + @pytest.mark.asyncio + @pytest.mark.parametrize( + ("raw_headers", "api_key", "expected_subject"), + [ + pytest.param({"authorization": "Bearer eyJ.x.y"}, "eyJ.x.y", "eyJ.x.y", id="idp-bearer-is-the-subject"), + pytest.param({"authorization": "Bearer sk-1234"}, "sk-1234", None, id="virtual-key-is-not-a-subject"), + pytest.param( + {"x-litellm-api-key": "ca-key", "authorization": "Bearer ca-key"}, + "ca-key", + None, + id="explicit-key-admission-repeated-in-authorization-is-not-a-subject", + ), + ], + ) + async def test_pre_call_tool_check_hands_sign_in_the_bearer_custom_auth_admitted( + self, raw_headers, api_key, expected_subject + ): + """Custom auth admits the caller on its own IdP token in ``Authorization`` with no + ``x-litellm-api-key``, so that token is the sign-in subject as it 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 = True + 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 2cc066c6f1b..c061e23159f 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 @@ -12554,7 +12554,9 @@ 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): + 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() @@ -12575,13 +12577,63 @@ class TestConnectSignInPreflight: mcp_servers=list(route_names), oauth2_headers=None, mcp_server_auth_headers=None, - user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"), + 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 + 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_rejected_subject_challenges_at_connect(self): from litellm.proxy._experimental.mcp_server.caller_sign_in import Rejected From ea7666685f97f96ad9e6d101dcfc267a6483b10c Mon Sep 17 00:00:00 2001 From: yucheng Date: Sat, 3 Oct 2026 13:19:48 +0000 Subject: [PATCH 47/51] fix(mcp): hand sign-in the bearer that admitted the caller and exchange only at connect The caller sign-in subject no longer depends on how the caller was admitted: a non-virtual bearer in Authorization is the subject unless it repeats x-litellm-api-key, so built-in OAuth2 and JWT admissions forward the caller's token to Agent 365 as the merge base did. The connect gate pre-flights that subject only while connecting; on an open session the tool-call hook runs the single exchange and answers a rejection inside the JSON-RPC envelope with its guardrail Logs row, again as the merge base did Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/caller_sign_in.py | 14 +++-- .../mcp_server/mcp_server_manager.py | 22 +++----- .../proxy/_experimental/mcp_server/server.py | 2 +- .../mcp_server/test_mcp_server_manager.py | 36 +++++++----- .../test_mcp_server_tool_calls_and_headers.py | 56 +++++++++++++++++-- 5 files changed, 90 insertions(+), 40 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/caller_sign_in.py b/litellm/proxy/_experimental/mcp_server/caller_sign_in.py index 1d6d75b438f..c4b1d0413c5 100644 --- a/litellm/proxy/_experimental/mcp_server/caller_sign_in.py +++ b/litellm/proxy/_experimental/mcp_server/caller_sign_in.py @@ -164,16 +164,19 @@ async def preflight_caller_sign_in( resource_metadata: str | None, connecting: Callable[[], Awaitable[bool]], ) -> None: - """Run every provider's connect-time check against the subject token, so a bearer the IdP will - reject surfaces as a challenge here rather than a JSON-RPC error at the first tool call. A fail-closed - provider outage is the connect's 503 only while ``connecting``; on an open session the tool-call hook - answers it inside the JSON-RPC envelope, with its guardrail Logs row.""" + """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, ) + if not await connecting(): + return for provider in _providers(): if provider.caller_sign_in(server, user_api_key_auth) is None: continue @@ -187,7 +190,6 @@ async def preflight_caller_sign_in( case Unavailable(fail_open=True): continue case Unavailable(detail=detail, fail_open=False): - if await connecting(): - raise HTTPException(status_code=503, detail=detail) + raise HTTPException(status_code=503, detail=detail) case _ as verdict: assert_never(verdict) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index af88f0e1eb6..1331d3b01d5 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -3949,20 +3949,16 @@ class MCPServerManager: def _caller_sign_in_subject_token( oauth2_headers: Mapping[str, str] | None, raw_headers: Mapping[str, str] | None, - user_api_key_auth: UserAPIKeyAuth | None, ) -> str | None: - """The bearer a caller sign-in provider validates: custom auth admits the caller on its own IdP token - in ``Authorization``, so that token is the subject there and only a virtual key is withheld.""" - admitted_on_own_bearer: Final = ( - user_api_key_auth is not None - and user_api_key_auth.authenticated_by_custom_auth - and not _has_explicit_litellm_admission_header(raw_headers) - ) - if not admitted_on_own_bearer: - return MCPServerManager._extract_subject_token(oauth2_headers, raw_headers, user_api_key_auth) + """The bearer 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; only a + virtual key, or a bearer repeating ``x-litellm-api-key``, is withheld.""" bearer: Final = MCPServerManager._extract_bearer_token(oauth2_headers, raw_headers) if bearer is None or bearer.startswith(LITELLM_VIRTUAL_KEY_PREFIX): 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( @@ -4282,9 +4278,7 @@ class MCPServerManager: caller_sign_in_for, ) - sign_in_subject: Final = self._caller_sign_in_subject_token( - oauth2_headers, raw_headers, user_api_key_auth - ) + 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 @@ -6042,7 +6036,7 @@ class MCPServerManager: 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, user_api_key_auth) + incoming_subject_token: Final = self._caller_sign_in_subject_token(None, raw_headers) pre_hook_kwargs: Final = { "guardrail_context": guardrail_context, diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 2893e044a27..ddbbe5c638f 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -1786,7 +1786,7 @@ if MCP_AVAILABLE: 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, user_api_key_auth + oauth2_headers, raw_headers ) if server is not None else None 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 6fd1f384d52..23130f060b1 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 @@ -6936,13 +6936,6 @@ class TestMCPServerManager: @pytest.mark.parametrize( ("raw_headers", "api_key", "expected_bearer", "expected_subject"), [ - pytest.param( - {"authorization": "Bearer eyJ.x.y"}, - "eyJ.x.y", - "eyJ.x.y", - None, - id="idp-token-as-admission-stays-raw-bearer", - ), pytest.param( {"authorization": "Bearer sk-1234"}, "sk-1234", @@ -6987,29 +6980,42 @@ class TestMCPServerManager: @pytest.mark.asyncio @pytest.mark.parametrize( - ("raw_headers", "api_key", "expected_subject"), + ("raw_headers", "api_key", "custom_auth", "expected_subject"), [ - pytest.param({"authorization": "Bearer eyJ.x.y"}, "eyJ.x.y", "eyJ.x.y", id="idp-bearer-is-the-subject"), - pytest.param({"authorization": "Bearer sk-1234"}, "sk-1234", None, id="virtual-key-is-not-a-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_custom_auth_admitted( - self, raw_headers, api_key, expected_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 admits the caller on its own IdP token in ``Authorization`` with no - ``x-litellm-api-key``, so that token is the sign-in subject as it was before the subject split.""" + """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 = True + 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={}) 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 c061e23159f..01aa5a286a7 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 @@ -3974,7 +3974,7 @@ async def test_owner_mismatch_on_a_torn_down_session_is_refused_before_the_body_ litellm.callbacks, guardrail, require_self=False ) - assert guardrail.preflight_calls == ["entra.jwt.token"] + 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] @@ -12634,6 +12634,53 @@ class TestConnectSignInPreflight: 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 @@ -12691,8 +12738,9 @@ class TestConnectSignInPreflight: ) 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, so the tools/call - hook answers the outage inside the result envelope and writes the guardrail Logs row, as base did.""" + 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 @@ -12766,7 +12814,7 @@ class TestConnectSignInPreflight: if session_id: server_module._remove_stateful_session_tracking(session_id) - assert guardrail.preflight_calls == ["entra.jwt.token"] + 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() From e3a56fab9b3a62ed3404ae186120275f3c854d24 Mon Sep 17 00:00:00 2001 From: yucheng Date: Sat, 3 Oct 2026 13:24:22 +0000 Subject: [PATCH 48/51] style(mcp): format the caller sign-in subject parametrization Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/test_mcp_server_manager.py | 10 ++++++---- 1 file changed, 6 insertions(+), 4 deletions(-) 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 23130f060b1..5bfd66116a8 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 @@ -6983,7 +6983,11 @@ class TestMCPServerManager: ("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" + {"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"}, @@ -6992,9 +6996,7 @@ class TestMCPServerManager: "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({"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", From b7aa52159c8d6f558bcdcdbffd272a03286ad960 Mon Sep 17 00:00:00 2001 From: yucheng Date: Sat, 3 Oct 2026 13:34:17 +0000 Subject: [PATCH 49/51] test(mcp): challenge the opaque bearer on initialize and reject it inside the tool-call envelope Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp/test_mcp_agent_365_guardrail.py | 14 +++++++++----- 1 file changed, 9 insertions(+), 5 deletions(-) diff --git a/tests/integration/mcp/test_mcp_agent_365_guardrail.py b/tests/integration/mcp/test_mcp_agent_365_guardrail.py index f3c923cbf10..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, @@ -145,18 +146,21 @@ def test_a_missing_or_malformed_caller_bearer_blocks_on_every_entry_point_whatev 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( - "tools/call", {"name": f"{rig.alias}-add", "arguments": {"entry": entry}} - ) + 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}" - continue + 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) - 1) + expected: Final = 2 * len(ENTRY_POINTS) - 1 assert rig.guardrail_statuses("call_mcp_tool", expected) == ["guardrail_intervened"] * expected From 2e118bd672ee29c3625ec4c05428fd0b39adee97 Mon Sep 17 00:00:00 2001 From: yucheng Date: Sat, 3 Oct 2026 15:16:35 +0000 Subject: [PATCH 50/51] fix(mcp): advertise the Agent 365 sign-in only when the guardrail gates a tagless connect Anonymous protected-resource discovery now runs the same should_run_guardrail probe as a keyed connect, so a default_on guardrail whose Mode only has tags and no default advertises the gateway authorization server as the merge base did instead of the Entra issuer and scope Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../guardrail_hooks/agent_365/agent_365.py | 36 +++++++++---------- .../guardrail_hooks/test_agent_365.py | 20 ++++++++++- 2 files changed, 37 insertions(+), 19 deletions(-) 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 dbc041ec806..ed501b007cd 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py +++ b/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py @@ -480,27 +480,27 @@ class Agent365Guardrail(CustomGuardrail): return str(uuid.uuid4()) def caller_sign_in(self, server: MCPServer, user_api_key_auth: "UserAPIKeyAuth | None") -> CallerSignIn | None: - """The Entra sign-in this guardrail requires of callers: only a ``default_on`` guardrail the caller's - key or team has not opted out of, because the anonymous metadata fetch that follows a challenge cannot - see which key selected a guardrail and would advertise the wrong issuer. Only servers that leave the - caller's top-level ``Authorization`` with the gateway qualify: a forwarded API-key header travels - upstream in its own slot and does not displace the Entra assertion.""" + """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 - if user_api_key_auth is not None: - probe: Final[_AdmissionProbe] = { - "metadata": { - "user_api_key_metadata": user_api_key_auth.metadata, # pyright: ignore[reportUnknownMemberType] # UserAPIKeyAuth.metadata is a raw dict - "user_api_key_team_metadata": user_api_key_auth.team_metadata, # pyright: ignore[reportUnknownMemberType] # UserAPIKeyAuth.team_metadata is a raw dict - } + 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 + } + 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) 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 0d5f5b3ad8b..a225be3cfda 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/test_agent_365.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/test_agent_365.py @@ -39,6 +39,7 @@ from litellm.proxy.utils import ProxyLogging from litellm.types.guardrails import ( GuardrailEventHooks, LitellmParams, + Mode, SupportedGuardrailIntegrations, ) from litellm.types.mcp import MCPAuth, MCPTransport @@ -179,6 +180,7 @@ def _make_guardrail( 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", @@ -188,7 +190,7 @@ def _make_guardrail( unreachable_fallback=unreachable_fallback, async_handler=handler, token_exchanger=exchanger if exchanger is not None else StubTokenExchanger(_obo_ok()), - event_hook="pre_mcp_call", + event_hook=event_hook, default_on=default_on, ) @@ -1403,6 +1405,22 @@ class TestCallerSignIn: 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([])) From 55210dd1af854a234657d4f7878f6f703db49aba Mon Sep 17 00:00:00 2001 From: yucheng Date: Sat, 3 Oct 2026 18:02:18 +0000 Subject: [PATCH 51/51] fix(mcp): sign-in subject only from a non-LiteLLM-key Bearer, OBO preflight answers before the body Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/caller_sign_in.py | 9 ++- .../mcp_server/mcp_server_manager.py | 24 ++++-- tests/integration/mcp/test_mcp_oauth_flows.py | 37 ++++++++++ .../mcp_server/test_caller_sign_in.py | 40 ++++++++++ .../mcp_server/test_mcp_server_manager.py | 74 +++++++++++++++++++ .../test_mcp_server_tool_calls_and_headers.py | 59 +++++++++++++++ 6 files changed, 234 insertions(+), 9 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/caller_sign_in.py b/litellm/proxy/_experimental/mcp_server/caller_sign_in.py index c4b1d0413c5..d948606dcf1 100644 --- a/litellm/proxy/_experimental/mcp_server/caller_sign_in.py +++ b/litellm/proxy/_experimental/mcp_server/caller_sign_in.py @@ -175,11 +175,12 @@ async def preflight_caller_sign_in( raise_token_exchange_challenge, ) - if not await connecting(): + 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 _providers(): - if provider.caller_sign_in(server, user_api_key_auth) is None: - continue + for provider in gating: match await provider.preflight_caller_sign_in(server, user_api_key_auth, subject_token): case SignedIn(): continue diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 1331d3b01d5..6d93756a797 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, @@ -1165,6 +1166,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")) @@ -3950,11 +3955,20 @@ class MCPServerManager: oauth2_headers: Mapping[str, str] | None, raw_headers: Mapping[str, str] | None, ) -> str | None: - """The bearer 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; only a - virtual key, or a bearer repeating ``x-litellm-api-key``, is withheld.""" - bearer: Final = MCPServerManager._extract_bearer_token(oauth2_headers, raw_headers) - if bearer is None or bearer.startswith(LITELLM_VIRTUAL_KEY_PREFIX): + """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: diff --git a/tests/integration/mcp/test_mcp_oauth_flows.py b/tests/integration/mcp/test_mcp_oauth_flows.py index 0e65ee9d2b1..fdedd62abb1 100644 --- a/tests/integration/mcp/test_mcp_oauth_flows.py +++ b/tests/integration/mcp/test_mcp_oauth_flows.py @@ -2,6 +2,7 @@ import base64 import hashlib import json import secrets +import socket import time import uuid from dataclasses import dataclass @@ -268,6 +269,42 @@ def test_a_malformed_caller_assertion_is_a_sign_in_challenge_not_an_outage(gatew 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/test_caller_sign_in.py b/tests/unit/proxy/_experimental/mcp_server/test_caller_sign_in.py index 83371400f4c..edd29bf3ece 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_caller_sign_in.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_caller_sign_in.py @@ -1,3 +1,4 @@ +import asyncio from collections.abc import Iterator, Mapping from typing import Final @@ -11,6 +12,7 @@ from litellm.proxy._experimental.mcp_server.caller_sign_in import ( CallerSignInProvider, SignedIn, caller_sign_in_for, + preflight_caller_sign_in, ) from litellm.proxy._types import UserAPIKeyAuth from litellm.types.mcp import MCPAuth, MCPTransport @@ -137,3 +139,41 @@ def test_oauth_utils_strips_the_route_relative_root_path(): "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_mcp_server_manager.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 5bfd66116a8..ee54f52c2fc 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 @@ -6950,6 +6950,34 @@ class TestMCPServerManager: "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( @@ -6978,6 +7006,52 @@ class TestMCPServerManager: 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"), 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 01aa5a286a7..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 @@ -12584,6 +12584,65 @@ class TestConnectSignInPreflight: 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