fix(mcp): diagnose resource-AS rejections per code and SSRF-validate the discovered ID-JAG endpoint

This commit is contained in:
Tin Chi Lo 2026-07-21 13:36:12 -07:00
parent fb18ef4cca
commit 1b108469a0
6 changed files with 153 additions and 42 deletions

View file

@ -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(

View file

@ -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:

View file

@ -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):

View file

@ -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

View file

@ -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"

View file

@ -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"