From b232b8cfa725a06a1606d56334a19f3a1fb3ed27 Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 2 Oct 2026 20:40:49 +0000 Subject: [PATCH] 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"))])