mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(mcp): diagnose resource-AS rejections per code and SSRF-validate the discovered ID-JAG endpoint
This commit is contained in:
parent
fb18ef4cca
commit
1b108469a0
6 changed files with 153 additions and 42 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue