From 1b108469a01ad9a96f5133db4836177b398325ad Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Tue, 21 Jul 2026 13:36:12 -0700 Subject: [PATCH] fix(mcp): diagnose resource-AS rejections per code and SSRF-validate the discovered ID-JAG endpoint --- .../mcp_server/mcp_server_manager.py | 31 ++++++-- .../outbound_credentials/resolver.py | 45 +++++++----- .../outbound_credentials/token_endpoint.py | 6 +- .../outbound_credentials/test_resolver.py | 21 +++++- .../test_token_endpoint.py | 20 ++++++ .../mcp_server/test_mcp_server_manager.py | 72 +++++++++++++++---- 6 files changed, 153 insertions(+), 42 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 060b7f32135..e2940455be2 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -45,7 +45,7 @@ from litellm.constants import ( ) from litellm.exceptions import BlockedPiiEntityError, GuardrailRaisedException from litellm.experimental_mcp_client.client import MCPClient, MCPSigV4Auth -from litellm.litellm_core_utils.url_utils import SSRFError, async_safe_get +from litellm.litellm_core_utils.url_utils import SSRFError, async_safe_get, validate_url from litellm.llms.custom_httpx.http_handler import get_async_httpx_client from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( MCPRequestHandler, @@ -1070,16 +1070,22 @@ class MCPServerManager: assertion and skips both the round-trip and the grant-profile gate.""" return auth_type == MCPAuth.oauth2_id_jag and not id_jag_resource_token_endpoint - @staticmethod + @classmethod def _gated_id_jag_endpoint( + cls, metadata: MCPOAuthMetadata | None, server_id: str, + server_url: str | None = None, ) -> str | None: """The discovered leg-2 endpoint, released only when the authorization server passes the enterprise-managed-authorization capability gate: the document must be advertised (never an - origin-fallback guess) and must list the id-jag grant profile, else an endpoint that cannot - serve the jwt-bearer leg would be silently trusted. A rejected discovery leaves the field - unset, so the existing fail-closed misconfigured error at client build names the gap.""" + origin-fallback guess), must list the id-jag grant profile, and a cross-authority endpoint + must clear SSRF validation, since the value came from the upstream's own metadata and will + later receive a POST carrying the gateway's client authentication and a minted assertion. + A same-authority endpoint (the server's own origin) is no pivot beyond the server the + gateway already talks to, which keeps intentional internal deployments working. A rejected + discovery leaves the field unset, so the existing fail-closed misconfigured error at client + build names the gap.""" if metadata is None or metadata.from_origin_fallback or not metadata.token_url: return None if not metadata.grant_profiles or _ID_JAG_GRANT_PROFILE not in metadata.grant_profiles: @@ -1093,6 +1099,19 @@ class MCPServerManager: _ID_JAG_GRANT_PROFILE, ) return None + if server_url and cls._is_same_authority_metadata_url(metadata.token_url, server_url): + return metadata.token_url + try: + validate_url(metadata.token_url) + except SSRFError as exc: + verbose_logger.warning( + "MCP server %s: discovered ID-JAG resource token endpoint was rejected by SSRF " + "validation and will not be autofilled (%s). Pin the endpoint explicitly if this " + "internal authorization server is intentional.", + server_id, + exc, + ) + return None return metadata.token_url def __init__( @@ -1462,6 +1481,7 @@ class MCPServerManager: or self._gated_id_jag_endpoint( gated_oauth_metadata if auth_type == MCPAuth.oauth2_id_jag else None, server_name or server_id, + server_url=server_url, ), id_jag_resource=server_config.get("id_jag_resource", None), client_private_key=server_config.get("client_private_key", None), @@ -1953,6 +1973,7 @@ class MCPServerManager: or self._gated_id_jag_endpoint( gated_oauth_metadata if auth_type == MCPAuth.oauth2_id_jag else None, mcp_server.server_id, + server_url=server_url, ), id_jag_resource=(credentials_dict.get("id_jag_resource") if credentials_dict else None), client_private_key=self._decrypt_credential_field( diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/resolver.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/resolver.py index 0c637cee55c..53aef2d92e8 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/resolver.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/resolver.py @@ -299,26 +299,35 @@ class UpstreamCredentialProvider: return None -# RFC 6749 §5.2 codes that, from a RESOURCE authorization server redeeming a freshly minted -# ID-JAG (RFC 7523 jwt-bearer leg), indicate a registration problem the operator must fix, not -# an outage: the gateway client is unknown there, unauthorized for the grant, the ID-JAG's -# client_id claim does not match the authenticating client, or the target is wrong. -_RESOURCE_AS_MISCONFIG_CODES = frozenset({"invalid_grant", "invalid_client", "unauthorized_client", "invalid_target"}) - - def _classify_resource_as_rejection(client_id: str, rejection: TokenEndpointRejection) -> CredError | None: - """The resource AS rejecting the jwt-bearer leg with a §5.2 code is an actionable - misconfiguration (the assertion was just minted, so expiry is not in play); anything else - keeps the default upstream_unavailable mapping.""" - if rejection.error not in _RESOURCE_AS_MISCONFIG_CODES: - return None + """A resource-AS §5.2 rejection of the jwt-bearer leg, diagnosed per code so the message + never claims more than the code proves: client-authentication codes name the registration, + invalid_target names the audience/resource configuration, and invalid_grant (which the EMA + profile mandates for a client_id-claim mismatch but also covers assertion-validation + failures like clock skew or key propagation) lists both, mismatch first. Anything else, + including a 5xx (never parsed into a rejection), keeps the retryable default mapping.""" described = f" ({rejection.error_description})" if rejection.error_description else "" - return CredError.of_misconfigured( - f"the resource authorization server rejected the ID-JAG with {rejection.error}{described}; " - f"the gateway client {client_id!r} must be registered at the resource authorization server " - "and match the ID-JAG's client_id claim (the IdP and the resource server must share the " - "gateway's client registration)" - ) + prefix = f"the resource authorization server rejected the ID-JAG with {rejection.error}{described}; " + if rejection.error in ("invalid_client", "unauthorized_client"): + return CredError.of_misconfigured( + prefix + f"the gateway client {client_id!r} failed to authenticate or is not authorized " + "for the jwt-bearer grant at the resource authorization server; check its client " + "registration and credentials there" + ) + if rejection.error == "invalid_target": + return CredError.of_misconfigured( + prefix + "the requested target was not accepted; check the server's audience and " + "id_jag_resource configuration against what the resource authorization server serves" + ) + if rejection.error == "invalid_grant": + return CredError.of_misconfigured( + prefix + f"the assertion was rejected; most commonly the gateway client {client_id!r} " + "is not the client named in the ID-JAG's client_id claim (the IdP and the resource " + "server must share the gateway's client registration), though an assertion-validation " + "failure such as clock skew or signing-key propagation produces the same code, so " + "retry once before changing configuration" + ) + return None def _id_jag_cache_key(subject_token: str, server_id: str, config: IdJagConfig) -> str: diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/token_endpoint.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/token_endpoint.py index ba82a3d1e10..91abd890cdb 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/token_endpoint.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/token_endpoint.py @@ -70,7 +70,9 @@ class _TokenEndpointResponse(BaseModel): class TokenEndpointRejection: """The RFC 6749 §5.2 error body of a token-endpoint 4xx, parsed so a caller that knows which leg it posted can classify the rejection (an ``invalid_grant`` from a resource - authorization server means something different from one at the IdP).""" + authorization server means something different from one at the IdP). A 5xx never parses + into one: an upstream fault stays retryable upstream_unavailable regardless of any error + body it happens to carry.""" status_code: int error: str @@ -81,6 +83,8 @@ RejectionClassifier = Callable[[TokenEndpointRejection], CredError | None] def _parse_token_endpoint_rejection(response: httpx.Response) -> TokenEndpointRejection | None: + if not 400 <= response.status_code < 500: + return None try: body = response.json() except (json.JSONDecodeError, ValueError): diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_resolver.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_resolver.py index 66ea2e08869..248159130f9 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_resolver.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_resolver.py @@ -649,8 +649,19 @@ async def test_id_jag_leg2_carries_a_resource_as_rejection_classifier_and_leg1_d ) -@pytest.mark.parametrize("code", ["invalid_grant", "invalid_client", "unauthorized_client", "invalid_target"]) -def test_resource_as_misconfig_codes_classify_as_misconfigured(code): +@pytest.mark.parametrize( + "code,expected_fragment", + [ + ("invalid_client", "client registration and credentials"), + ("unauthorized_client", "client registration and credentials"), + ("invalid_target", "audience and"), + ("invalid_grant", "client_id claim"), + ], +) +def test_resource_as_rejections_get_per_code_diagnoses(code, expected_fragment): + """The message never claims more than the code proves: authn codes name the registration, + invalid_target names the audience config, invalid_grant lists mismatch first and the + assertion-validation alternatives (clock skew, key propagation) second.""" from litellm.proxy._experimental.mcp_server.outbound_credentials.resolver import ( _classify_resource_as_rejection, ) @@ -663,4 +674,8 @@ def test_resource_as_misconfig_codes_classify_as_misconfigured(code): ) assert classified is not None assert classified.tag == "misconfigured" - assert "gw-client" in classified.summary + assert expected_fragment in classified.summary + if code == "invalid_grant": + assert "retry" in classified.summary + if code == "invalid_target": + assert "registration" not in classified.summary diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_token_endpoint.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_token_endpoint.py index 38144d278fe..70f9e005d70 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_token_endpoint.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_token_endpoint.py @@ -484,3 +484,23 @@ async def test_fetch_unparseable_rejection_body_keeps_upstream_unavailable(body) assert isinstance(result, Error) assert result.error.tag == "upstream_unavailable" + + +@pytest.mark.asyncio +async def test_fetch_5xx_with_oauth_error_body_stays_upstream_unavailable(): + """A 5xx is an upstream fault regardless of any error body it carries; only a 4xx is an + RFC 6749 rejection eligible for classification.""" + with patch( + _PATCH_TARGET, + return_value=_client(_oauth_error_resp(status_code=502, body={"error": "invalid_grant"})), + ): + result = await TokenEndpointClient().fetch( + _ENDPOINT, + _CLIENT_ID, + {"grant_type": "g"}, + ClientSecretAuth(client_secret=SecretStr("s")), + classify_rejection=lambda rejection: CredError.of_misconfigured("should not fire"), + ) + + assert isinstance(result, Error) + assert result.error.tag == "upstream_unavailable" 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 da56843efe2..dc707c14414 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 @@ -9053,8 +9053,8 @@ class TestIdJagEndpointDiscovery: def _ema_metadata(self, grant_profiles): return MCPOAuthMetadata( - token_url="https://ras.example.com/token", - discovered_issuer="https://ras.example.com", + token_url="https://up.example.com/ras/token", + discovered_issuer="https://up.example.com/ras", grant_profiles=grant_profiles, ) @@ -9068,20 +9068,24 @@ class TestIdJagEndpointDiscovery: assert needs(None, None) is False def test_gate_releases_only_advertised_id_jag_capable_endpoints(self): - gate = MCPServerManager._gated_id_jag_endpoint profile = "urn:ietf:params:oauth:grant-profile:id-jag" - assert gate(self._ema_metadata([profile]), "s1") == "https://ras.example.com/token" - assert gate(self._ema_metadata([profile, "other"]), "s1") == "https://ras.example.com/token" - assert gate(None, "s1") is None - assert gate(self._ema_metadata(None), "s1") is None - assert gate(self._ema_metadata([]), "s1") is None - assert gate(self._ema_metadata(["urn:other:profile"]), "s1") is None + server_url = "https://up.example.com/mcp" + + def gate(metadata): + return MCPServerManager._gated_id_jag_endpoint(metadata, "s1", server_url=server_url) + + assert gate(self._ema_metadata([profile])) == "https://up.example.com/ras/token" + assert gate(self._ema_metadata([profile, "other"])) == "https://up.example.com/ras/token" + assert gate(None) is None + assert gate(self._ema_metadata(None)) is None + assert gate(self._ema_metadata([])) is None + assert gate(self._ema_metadata(["urn:other:profile"])) is None no_token_url = MCPOAuthMetadata(grant_profiles=[profile]) - assert gate(no_token_url, "s1") is None + assert gate(no_token_url) is None guessed = MCPOAuthMetadata( - token_url="https://ras.example.com/token", grant_profiles=[profile], from_origin_fallback=True + token_url="https://up.example.com/ras/token", grant_profiles=[profile], from_origin_fallback=True ) - assert gate(guessed, "s1") is None + assert gate(guessed) is None @pytest.mark.asyncio async def test_build_from_table_autofills_and_persists_gated_id_jag_endpoint(self): @@ -9095,10 +9099,10 @@ class TestIdJagEndpointDiscovery: built = await manager.build_mcp_server_from_table(row, credentials_are_encrypted=False) mock_discovery.assert_awaited_once() - assert built.id_jag_resource_token_endpoint == "https://ras.example.com/token" + assert built.id_jag_resource_token_endpoint == "https://up.example.com/ras/token" mock_persist.assert_awaited_once() persist_kwargs = mock_persist.await_args.kwargs - assert persist_kwargs["discovered_endpoint"] == "https://ras.example.com/token" + assert persist_kwargs["discovered_endpoint"] == "https://up.example.com/ras/token" assert persist_kwargs["existing_endpoint"] is None @pytest.mark.asyncio @@ -9149,7 +9153,7 @@ class TestIdJagEndpointDiscovery: await manager.load_servers_from_config(config) server = next(iter(manager.config_mcp_servers.values())) - assert server.id_jag_resource_token_endpoint == "https://ras.example.com/token" + assert server.id_jag_resource_token_endpoint == "https://up.example.com/ras/token" @pytest.mark.asyncio async def test_load_servers_from_config_pinned_id_jag_endpoint_skips_discovery(self): @@ -9191,3 +9195,41 @@ class TestIdJagEndpointDiscovery: assert metadata is not None assert metadata.grant_profiles == ["urn:ietf:params:oauth:grant-profile:id-jag"] assert metadata.token_url == "https://ras.example.com/token" + + +class TestIdJagEndpointSSRFGate: + """The discovered leg-2 endpoint later receives a POST with the gateway's client + authentication, so a cross-authority value must clear SSRF validation before release; + the server's own origin is no pivot and passes.""" + + _PROFILE = "urn:ietf:params:oauth:grant-profile:id-jag" + + def _metadata(self, token_url): + return MCPOAuthMetadata( + token_url=token_url, discovered_issuer="https://ras.example.com", grant_profiles=[self._PROFILE] + ) + + def test_cross_authority_internal_endpoint_is_refused(self): + gated = MCPServerManager._gated_id_jag_endpoint( + self._metadata("http://169.254.169.254/token"), "s1", server_url="https://up.example.com/mcp" + ) + assert gated is None + + def test_same_authority_internal_endpoint_passes(self): + gated = MCPServerManager._gated_id_jag_endpoint( + self._metadata("http://127.0.0.1:8952/ras/token"), "s1", server_url="http://127.0.0.1:8952/mcp" + ) + assert gated == "http://127.0.0.1:8952/ras/token" + + def test_cross_authority_public_endpoint_passes_validation(self): + from unittest.mock import patch as _patch + + with _patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.validate_url", + return_value=("https://1.2.3.4/token", "ras.example.com"), + ) as mock_validate: + gated = MCPServerManager._gated_id_jag_endpoint( + self._metadata("https://ras.example.com/token"), "s1", server_url="https://up.example.com/mcp" + ) + mock_validate.assert_called_once_with("https://ras.example.com/token") + assert gated == "https://ras.example.com/token"