mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
feat(mcp): discover and gate the ID-JAG resource token endpoint and classify resource-AS rejections
This commit is contained in:
parent
869ef0cbfb
commit
fb18ef4cca
7 changed files with 469 additions and 6 deletions
|
|
@ -196,6 +196,10 @@ _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES: tuple[MCPAuth, ...] = (
|
|||
MCPAuth.oauth_delegate,
|
||||
)
|
||||
|
||||
# draft-ietf-oauth-identity-assertion-authz-grant: the grant profile an authorization server
|
||||
# advertises in authorization_grant_profiles_supported when it can redeem ID-JAG assertions.
|
||||
_ID_JAG_GRANT_PROFILE = "urn:ietf:params:oauth:grant-profile:id-jag"
|
||||
|
||||
|
||||
def _blank_to_none(value: str | None) -> str | None:
|
||||
"""Collapse an absent, empty, or whitespace-only string to ``None``.
|
||||
|
|
@ -1055,6 +1059,42 @@ class MCPServerManager:
|
|||
"""
|
||||
return auth_type == MCPAuth.oauth2_token_exchange and not (token_exchange_endpoint or token_url)
|
||||
|
||||
@staticmethod
|
||||
def _id_jag_needs_endpoint_discovery(
|
||||
auth_type: MCPAuthType | None,
|
||||
id_jag_resource_token_endpoint: str | None,
|
||||
) -> bool:
|
||||
"""An ``oauth2_id_jag`` server with no pinned leg-2 endpoint can have its resource
|
||||
authorization server's token endpoint discovered (RFC 9728 -> RFC 8414), the same chain
|
||||
the OBO flow uses; a pinned ``id_jag_resource_token_endpoint`` is the operator's explicit
|
||||
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
|
||||
def _gated_id_jag_endpoint(
|
||||
metadata: MCPOAuthMetadata | None,
|
||||
server_id: str,
|
||||
) -> 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."""
|
||||
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:
|
||||
verbose_logger.warning(
|
||||
"MCP server %s: discovered authorization server %s does not advertise the "
|
||||
"id-jag grant profile (%s); refusing to autofill id_jag_resource_token_endpoint. "
|
||||
"The upstream's authorization server must support enterprise-managed authorization, "
|
||||
"or pin the endpoint explicitly.",
|
||||
server_id,
|
||||
metadata.discovered_issuer or metadata.token_url,
|
||||
_ID_JAG_GRANT_PROFILE,
|
||||
)
|
||||
return None
|
||||
return metadata.token_url
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
cred_provider: Optional[UpstreamCredentialProvider] = None,
|
||||
|
|
@ -1282,7 +1322,12 @@ class MCPServerManager:
|
|||
manual_registration_url,
|
||||
)
|
||||
should_discover = _has_oauth_discovery_source(server_url, use_issuer_anchor) and (
|
||||
is_discovery_auth_type or obo_needs_discovery
|
||||
is_discovery_auth_type
|
||||
or obo_needs_discovery
|
||||
or self._id_jag_needs_endpoint_discovery(
|
||||
auth_type,
|
||||
server_config.get("id_jag_resource_token_endpoint"),
|
||||
)
|
||||
)
|
||||
if not should_discover:
|
||||
mcp_oauth_metadata = None
|
||||
|
|
@ -1413,7 +1458,11 @@ class MCPServerManager:
|
|||
DEFAULT_SUBJECT_TOKEN_TYPE,
|
||||
),
|
||||
# ID-JAG fields
|
||||
id_jag_resource_token_endpoint=server_config.get("id_jag_resource_token_endpoint", None),
|
||||
id_jag_resource_token_endpoint=server_config.get("id_jag_resource_token_endpoint", None)
|
||||
or self._gated_id_jag_endpoint(
|
||||
gated_oauth_metadata if auth_type == MCPAuth.oauth2_id_jag else None,
|
||||
server_name or server_id,
|
||||
),
|
||||
id_jag_resource=server_config.get("id_jag_resource", None),
|
||||
client_private_key=server_config.get("client_private_key", None),
|
||||
client_private_key_id=server_config.get("client_private_key_id", None),
|
||||
|
|
@ -1680,11 +1729,13 @@ class MCPServerManager:
|
|||
use_issuer_anchor: bool,
|
||||
scopes: Optional[list[str]],
|
||||
token_exchange_endpoint: Optional[str],
|
||||
id_jag_resource_token_endpoint: str | None = None,
|
||||
) -> Optional[MCPOAuthMetadata]:
|
||||
has_all_upstream_oauth_fields = bool(manual_authorization_url and manual_token_url and scopes)
|
||||
needs_discovery = _has_oauth_discovery_source(server_url, use_issuer_anchor) and (
|
||||
(is_discovery_auth_type and not has_all_upstream_oauth_fields)
|
||||
or self._obo_needs_endpoint_discovery(auth_type, token_exchange_endpoint, manual_token_url)
|
||||
or self._id_jag_needs_endpoint_discovery(auth_type, id_jag_resource_token_endpoint)
|
||||
)
|
||||
if not needs_discovery:
|
||||
mcp_oauth_metadata: Optional[MCPOAuthMetadata] = None
|
||||
|
|
@ -1801,6 +1852,7 @@ class MCPServerManager:
|
|||
manual_token_url = _blank_to_none(mcp_server.token_url)
|
||||
manual_registration_url = _blank_to_none(mcp_server.registration_url)
|
||||
is_discovery_auth_type = auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES
|
||||
manual_id_jag_endpoint = credentials_dict.get("id_jag_resource_token_endpoint") if credentials_dict else None
|
||||
token_exchange_endpoint = mcp_server.token_exchange_endpoint or (
|
||||
credentials_dict.get("token_exchange_endpoint") if credentials_dict else None
|
||||
)
|
||||
|
|
@ -1823,6 +1875,7 @@ class MCPServerManager:
|
|||
use_issuer_anchor=use_issuer_anchor,
|
||||
scopes=scopes,
|
||||
token_exchange_endpoint=token_exchange_endpoint,
|
||||
id_jag_resource_token_endpoint=manual_id_jag_endpoint,
|
||||
)
|
||||
|
||||
resolved_scopes = scopes or (gated_oauth_metadata.scopes if gated_oauth_metadata else None)
|
||||
|
|
@ -1896,8 +1949,10 @@ class MCPServerManager:
|
|||
or (credentials_dict.get("subject_token_type") if credentials_dict else None)
|
||||
or DEFAULT_SUBJECT_TOKEN_TYPE,
|
||||
# ID-JAG fields — read from credentials JSON blob
|
||||
id_jag_resource_token_endpoint=(
|
||||
credentials_dict.get("id_jag_resource_token_endpoint") if credentials_dict else None
|
||||
id_jag_resource_token_endpoint=manual_id_jag_endpoint
|
||||
or self._gated_id_jag_endpoint(
|
||||
gated_oauth_metadata if auth_type == MCPAuth.oauth2_id_jag else None,
|
||||
mcp_server.server_id,
|
||||
),
|
||||
id_jag_resource=(credentials_dict.get("id_jag_resource") if credentials_dict else None),
|
||||
client_private_key=self._decrypt_credential_field(
|
||||
|
|
@ -1934,8 +1989,59 @@ class MCPServerManager:
|
|||
metadata=gated_oauth_metadata,
|
||||
is_issuer_anchored=use_issuer_anchor,
|
||||
)
|
||||
await self._persist_discovered_id_jag_endpoint(
|
||||
server_id=mcp_server.server_id,
|
||||
auth_type=auth_type,
|
||||
existing_endpoint=manual_id_jag_endpoint,
|
||||
discovered_endpoint=new_server.id_jag_resource_token_endpoint,
|
||||
)
|
||||
return new_server
|
||||
|
||||
async def _persist_discovered_id_jag_endpoint(
|
||||
self,
|
||||
*,
|
||||
server_id: str,
|
||||
auth_type: MCPAuthType | None,
|
||||
existing_endpoint: str | None,
|
||||
discovered_endpoint: str | None,
|
||||
) -> None:
|
||||
"""Write a gate-passing discovered ID-JAG leg-2 endpoint into the credentials blob.
|
||||
|
||||
Same contract as ``_persist_discovered_obo_token_url``: fires at most once per server
|
||||
(skipped once a value exists), best-effort, and makes ``_id_jag_needs_endpoint_discovery``
|
||||
read False on the next build on every pod, so a transient discovery outage cannot strand
|
||||
the server once one build has succeeded. The blob is the field's one home (there is no
|
||||
column), which also keeps it inside the master-key rotation that re-encrypts the blob.
|
||||
"""
|
||||
if auth_type != MCPAuth.oauth2_id_jag:
|
||||
return
|
||||
if existing_endpoint or not discovered_endpoint:
|
||||
return
|
||||
from litellm.proxy._experimental.mcp_server.db import ( # noqa: PLC0415 # db.py imports this module at load
|
||||
update_mcp_server,
|
||||
)
|
||||
from litellm.proxy._types import UpdateMCPServerRequest # noqa: PLC0415 # heavy module; import at call time
|
||||
from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415 # runtime value, set after startup
|
||||
|
||||
if prisma_client is None:
|
||||
return
|
||||
try:
|
||||
await update_mcp_server(
|
||||
prisma_client=prisma_client,
|
||||
data=UpdateMCPServerRequest.model_validate(
|
||||
{
|
||||
"server_id": server_id,
|
||||
"credentials": {"id_jag_resource_token_endpoint": discovered_endpoint},
|
||||
}
|
||||
),
|
||||
touched_by="mcp_oauth_discovery",
|
||||
)
|
||||
verbose_logger.info("Persisted discovered ID-JAG resource token endpoint for MCP server %s", server_id)
|
||||
except Exception as exc: # noqa: BLE001 - best-effort; a failed write re-discovers next build
|
||||
verbose_logger.warning(
|
||||
"Failed to persist discovered ID-JAG resource token endpoint for MCP server %s: %s", server_id, exc
|
||||
)
|
||||
|
||||
async def _persist_discovered_obo_token_url(
|
||||
self,
|
||||
*,
|
||||
|
|
@ -3713,6 +3819,7 @@ class MCPServerManager:
|
|||
token_url=data.get("token_endpoint"),
|
||||
registration_url=data.get("registration_endpoint"),
|
||||
discovered_issuer=claimed_issuer if isinstance(claimed_issuer, str) and claimed_issuer else None,
|
||||
grant_profiles=self._extract_scopes(data.get("authorization_grant_profiles_supported")),
|
||||
)
|
||||
|
||||
if any(
|
||||
|
|
|
|||
|
|
@ -45,6 +45,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.token_endpoint
|
|||
ExchangedToken,
|
||||
ExchangedTokenCache,
|
||||
TokenEndpointClient,
|
||||
TokenEndpointRejection,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchanger import (
|
||||
TokenExchanger,
|
||||
|
|
@ -209,6 +210,7 @@ class UpstreamCredentialProvider:
|
|||
config.client_id,
|
||||
leg2_params,
|
||||
config.client_auth,
|
||||
classify_rejection=partial(_classify_resource_as_rejection, config.client_id),
|
||||
)
|
||||
|
||||
match await self._exchanged_tokens.get_or_compute(cache_key, _exchange):
|
||||
|
|
@ -297,6 +299,28 @@ 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
|
||||
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)"
|
||||
)
|
||||
|
||||
|
||||
def _id_jag_cache_key(subject_token: str, server_id: str, config: IdJagConfig) -> str:
|
||||
"""Bind the cached leg-2 bearer to the caller token, the server, AND the config that minted it.
|
||||
|
||||
|
|
|
|||
|
|
@ -66,6 +66,38 @@ class _TokenEndpointResponse(BaseModel):
|
|||
expires_in: int | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
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)."""
|
||||
|
||||
status_code: int
|
||||
error: str
|
||||
error_description: str | None
|
||||
|
||||
|
||||
RejectionClassifier = Callable[[TokenEndpointRejection], CredError | None]
|
||||
|
||||
|
||||
def _parse_token_endpoint_rejection(response: httpx.Response) -> TokenEndpointRejection | None:
|
||||
try:
|
||||
body = response.json()
|
||||
except (json.JSONDecodeError, ValueError):
|
||||
return None
|
||||
if not isinstance(body, dict):
|
||||
return None
|
||||
error = body.get("error")
|
||||
if not isinstance(error, str) or not error:
|
||||
return None
|
||||
description = body.get("error_description")
|
||||
return TokenEndpointRejection(
|
||||
status_code=response.status_code,
|
||||
error=error,
|
||||
error_description=description if isinstance(description, str) and description else None,
|
||||
)
|
||||
|
||||
|
||||
class TokenEndpointClient:
|
||||
"""One authenticated POST to an OAuth token endpoint, returning the minted token as a value."""
|
||||
|
||||
|
|
@ -75,6 +107,7 @@ class TokenEndpointClient:
|
|||
client_id: str,
|
||||
grant_params: Mapping[str, str],
|
||||
client_auth: ClientAuth,
|
||||
classify_rejection: RejectionClassifier | None = None,
|
||||
) -> Result[ExchangedToken, CredError]:
|
||||
try:
|
||||
data = {**grant_params, **_client_auth_params(endpoint, client_id, client_auth)}
|
||||
|
|
@ -92,6 +125,10 @@ class TokenEndpointClient:
|
|||
verbose_proxy_logger.warning(
|
||||
"MCP token endpoint %s failed with status %s", endpoint, exc.response.status_code
|
||||
)
|
||||
rejection = _parse_token_endpoint_rejection(exc.response) if classify_rejection else None
|
||||
classified = classify_rejection(rejection) if classify_rejection and rejection else None
|
||||
if classified is not None:
|
||||
return Error(classified)
|
||||
return Error(
|
||||
CredError.of_upstream_unavailable(f"token exchange failed with status {exc.response.status_code}")
|
||||
)
|
||||
|
|
|
|||
|
|
@ -35,6 +35,11 @@ class MCPOAuthMetadata(BaseModel):
|
|||
"""True when the metadata came from guessing the resource origin as its authorization
|
||||
server rather than from an RFC 9728/8414-advertised document. Guessed endpoints are
|
||||
usable in memory but must never be persisted as configuration."""
|
||||
grant_profiles: Optional[List[str]] = None
|
||||
"""The authorization server's ``authorization_grant_profiles_supported`` (draft OAuth
|
||||
identity-assertion-authz-grant); the enterprise-managed-authorization gate requires the
|
||||
id-jag profile in it before an autofilled ID-JAG endpoint is trusted. ``None`` means the
|
||||
document did not carry the field, distinct from an empty advertisement."""
|
||||
|
||||
|
||||
class MCPServer(BaseModel):
|
||||
|
|
|
|||
|
|
@ -56,14 +56,16 @@ def _id_jag_config() -> IdJagConfig:
|
|||
|
||||
|
||||
class _FakeTokenEndpoint:
|
||||
"""Records each fetch and returns the next canned Result, leg by leg."""
|
||||
"""Records each fetch (and its rejection classifier) and returns the next canned Result."""
|
||||
|
||||
def __init__(self, results: list[Result[ExchangedToken, CredError]]) -> None:
|
||||
self._results = list(results)
|
||||
self.calls: list[tuple[str, str, dict[str, str]]] = []
|
||||
self.classifiers: list[object] = []
|
||||
|
||||
async def fetch(self, endpoint, client_id, grant_params, client_auth):
|
||||
async def fetch(self, endpoint, client_id, grant_params, client_auth, classify_rejection=None):
|
||||
self.calls.append((endpoint, client_id, dict(grant_params)))
|
||||
self.classifiers.append(classify_rejection)
|
||||
return self._results.pop(0)
|
||||
|
||||
|
||||
|
|
@ -615,3 +617,50 @@ async def test_invalidate_credentials_for_id_jag_is_a_noop_without_a_caller_toke
|
|||
assert isinstance(first, Ok) and isinstance(second, Ok)
|
||||
assert _emitted(second.ok)["Authorization"] == "Bearer cached-bearer"
|
||||
assert len(endpoint.calls) == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_id_jag_leg2_carries_a_resource_as_rejection_classifier_and_leg1_does_not():
|
||||
"""A section 5.2 rejection means different things per leg: at the resource AS redeeming a
|
||||
freshly minted ID-JAG it is a client-registration misconfiguration; at the IdP it keeps the
|
||||
default mapping. The classifier therefore rides only the leg-2 fetch."""
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_endpoint import (
|
||||
TokenEndpointRejection,
|
||||
)
|
||||
|
||||
endpoint = _FakeTokenEndpoint(_two_leg_ok("final-access"))
|
||||
provider = UpstreamCredentialProvider(token_endpoint=endpoint)
|
||||
|
||||
result = await provider.resolve_credentials(_with_inbound("user-id-token"), _spec(_id_jag_config()))
|
||||
|
||||
assert isinstance(result, Ok)
|
||||
leg1_classifier, leg2_classifier = endpoint.classifiers
|
||||
assert leg1_classifier is None
|
||||
assert leg2_classifier is not None
|
||||
classified = leg2_classifier(
|
||||
TokenEndpointRejection(status_code=400, error="invalid_grant", error_description="client mismatch")
|
||||
)
|
||||
assert classified is not None
|
||||
assert classified.tag == "misconfigured"
|
||||
assert "litellm" in classified.summary
|
||||
assert "client mismatch" in classified.summary
|
||||
assert (
|
||||
leg2_classifier(TokenEndpointRejection(status_code=502, error="server_error", error_description=None)) is None
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("code", ["invalid_grant", "invalid_client", "unauthorized_client", "invalid_target"])
|
||||
def test_resource_as_misconfig_codes_classify_as_misconfigured(code):
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.resolver import (
|
||||
_classify_resource_as_rejection,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_endpoint import (
|
||||
TokenEndpointRejection,
|
||||
)
|
||||
|
||||
classified = _classify_resource_as_rejection(
|
||||
"gw-client", TokenEndpointRejection(status_code=400, error=code, error_description=None)
|
||||
)
|
||||
assert classified is not None
|
||||
assert classified.tag == "misconfigured"
|
||||
assert "gw-client" in classified.summary
|
||||
|
|
|
|||
|
|
@ -406,3 +406,81 @@ async def test_cache_does_not_store_a_failed_compute():
|
|||
assert isinstance(first, Error)
|
||||
assert isinstance(second, Ok) and second.ok == "recovered"
|
||||
assert calls == 2
|
||||
|
||||
|
||||
def _oauth_error_resp(status_code=400, body=None):
|
||||
error_resp = MagicMock()
|
||||
error_resp.status_code = status_code
|
||||
error_resp.json.return_value = body if body is not None else {"error": "invalid_grant"}
|
||||
error_resp.raise_for_status.side_effect = httpx.HTTPStatusError(
|
||||
"Bad Request", request=MagicMock(), response=error_resp
|
||||
)
|
||||
return error_resp
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_rejection_classifier_overrides_the_default_mapping():
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_endpoint import (
|
||||
TokenEndpointRejection,
|
||||
)
|
||||
|
||||
seen = []
|
||||
|
||||
def classify(rejection: TokenEndpointRejection):
|
||||
seen.append(rejection)
|
||||
return CredError.of_misconfigured("client registration mismatch")
|
||||
|
||||
with patch(_PATCH_TARGET, return_value=_client(_oauth_error_resp(body={"error": "invalid_grant", "error_description": "aud mismatch"}))):
|
||||
result = await TokenEndpointClient().fetch(
|
||||
_ENDPOINT,
|
||||
_CLIENT_ID,
|
||||
{"grant_type": "g"},
|
||||
ClientSecretAuth(client_secret=SecretStr("s")),
|
||||
classify_rejection=classify,
|
||||
)
|
||||
|
||||
assert isinstance(result, Error)
|
||||
assert result.error.tag == "misconfigured"
|
||||
assert seen[0].error == "invalid_grant"
|
||||
assert seen[0].error_description == "aud mismatch"
|
||||
assert seen[0].status_code == 400
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_classifier_returning_none_keeps_upstream_unavailable():
|
||||
with patch(_PATCH_TARGET, return_value=_client(_oauth_error_resp(body={"error": "server_error"}))):
|
||||
result = await TokenEndpointClient().fetch(
|
||||
_ENDPOINT,
|
||||
_CLIENT_ID,
|
||||
{"grant_type": "g"},
|
||||
ClientSecretAuth(client_secret=SecretStr("s")),
|
||||
classify_rejection=lambda rejection: None,
|
||||
)
|
||||
|
||||
assert isinstance(result, Error)
|
||||
assert result.error.tag == "upstream_unavailable"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("body", [None, "not-a-dict", {}, {"error": ""}, {"error": 42}])
|
||||
async def test_fetch_unparseable_rejection_body_keeps_upstream_unavailable(body):
|
||||
resp = MagicMock()
|
||||
resp.status_code = 400
|
||||
if body is None:
|
||||
import json as _json
|
||||
|
||||
resp.json.side_effect = _json.JSONDecodeError("x", "y", 0)
|
||||
else:
|
||||
resp.json.return_value = body
|
||||
resp.raise_for_status.side_effect = httpx.HTTPStatusError("Bad Request", request=MagicMock(), response=resp)
|
||||
with patch(_PATCH_TARGET, return_value=_client(resp)):
|
||||
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"
|
||||
|
|
|
|||
|
|
@ -9028,3 +9028,166 @@ class TestUrllessIssuerDiscovery:
|
|||
anchored.assert_awaited_once_with("https://idp.example.com", None)
|
||||
resource_rooted.assert_not_awaited()
|
||||
assert built.token_url == "https://idp.example.com/token"
|
||||
class TestIdJagEndpointDiscovery:
|
||||
"""EMA discovery + capability gate: an oauth2_id_jag server with no pinned leg-2 endpoint
|
||||
autofills it from RFC 9728 -> RFC 8414 discovery, released only when the authorization
|
||||
server advertises the id-jag grant profile; a pinned endpoint skips discovery entirely."""
|
||||
|
||||
def _id_jag_row(self, credentials=None):
|
||||
base_credentials = {
|
||||
"client_id": "cid",
|
||||
"client_secret": "csec",
|
||||
"token_exchange_endpoint": "https://idp.example.com/org/token",
|
||||
}
|
||||
return LiteLLM_MCPServerTable(
|
||||
server_id="idjag-db-1",
|
||||
alias="idjag_db",
|
||||
description="ema from db",
|
||||
url="https://up.example.com/mcp",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth2_id_jag,
|
||||
credentials={**base_credentials, **(credentials or {})},
|
||||
created_at=datetime.now(),
|
||||
updated_at=datetime.now(),
|
||||
)
|
||||
|
||||
def _ema_metadata(self, grant_profiles):
|
||||
return MCPOAuthMetadata(
|
||||
token_url="https://ras.example.com/token",
|
||||
discovered_issuer="https://ras.example.com",
|
||||
grant_profiles=grant_profiles,
|
||||
)
|
||||
|
||||
def test_needs_discovery_only_for_unpinned_id_jag(self):
|
||||
needs = MCPServerManager._id_jag_needs_endpoint_discovery
|
||||
assert needs(MCPAuth.oauth2_id_jag, None) is True
|
||||
assert needs(MCPAuth.oauth2_id_jag, "") is True
|
||||
assert needs(MCPAuth.oauth2_id_jag, "https://ras.example.com/token") is False
|
||||
assert needs(MCPAuth.oauth2_token_exchange, None) is False
|
||||
assert needs(MCPAuth.oauth2, None) is False
|
||||
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
|
||||
no_token_url = MCPOAuthMetadata(grant_profiles=[profile])
|
||||
assert gate(no_token_url, "s1") is None
|
||||
guessed = MCPOAuthMetadata(
|
||||
token_url="https://ras.example.com/token", grant_profiles=[profile], from_origin_fallback=True
|
||||
)
|
||||
assert gate(guessed, "s1") is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_build_from_table_autofills_and_persists_gated_id_jag_endpoint(self):
|
||||
manager = MCPServerManager()
|
||||
row = self._id_jag_row()
|
||||
metadata = self._ema_metadata(["urn:ietf:params:oauth:grant-profile:id-jag"])
|
||||
with (
|
||||
patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=metadata)) as mock_discovery,
|
||||
patch.object(manager, "_persist_discovered_id_jag_endpoint", new=AsyncMock()) as mock_persist,
|
||||
):
|
||||
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"
|
||||
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["existing_endpoint"] is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_build_from_table_refuses_autofill_without_the_grant_profile(self):
|
||||
manager = MCPServerManager()
|
||||
row = self._id_jag_row()
|
||||
metadata = self._ema_metadata(["urn:other:profile"])
|
||||
with (
|
||||
patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=metadata)),
|
||||
patch.object(manager, "_persist_discovered_id_jag_endpoint", new=AsyncMock()) as mock_persist,
|
||||
):
|
||||
built = await manager.build_mcp_server_from_table(row, credentials_are_encrypted=False)
|
||||
|
||||
assert built.id_jag_resource_token_endpoint is None
|
||||
persist_kwargs = mock_persist.await_args.kwargs
|
||||
assert persist_kwargs["discovered_endpoint"] is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_build_from_table_pinned_endpoint_skips_discovery_and_gate(self):
|
||||
manager = MCPServerManager()
|
||||
row = self._id_jag_row(credentials={"id_jag_resource_token_endpoint": "https://pinned.example.com/token"})
|
||||
with (
|
||||
patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=None)) as mock_discovery,
|
||||
patch.object(manager, "_persist_discovered_id_jag_endpoint", new=AsyncMock()) as mock_persist,
|
||||
):
|
||||
built = await manager.build_mcp_server_from_table(row, credentials_are_encrypted=False)
|
||||
|
||||
mock_discovery.assert_not_awaited()
|
||||
assert built.id_jag_resource_token_endpoint == "https://pinned.example.com/token"
|
||||
persist_kwargs = mock_persist.await_args.kwargs
|
||||
assert persist_kwargs["existing_endpoint"] == "https://pinned.example.com/token"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_load_servers_from_config_autofills_gated_id_jag_endpoint(self):
|
||||
manager = MCPServerManager()
|
||||
config = {
|
||||
"idjag_cfg": {
|
||||
"url": "https://up.example.com/mcp",
|
||||
"transport": "http",
|
||||
"auth_type": "oauth2_id_jag",
|
||||
"token_exchange_endpoint": "https://idp.example.com/org/token",
|
||||
"client_id": "cid",
|
||||
"client_secret": "csec",
|
||||
}
|
||||
}
|
||||
metadata = self._ema_metadata(["urn:ietf:params:oauth:grant-profile:id-jag"])
|
||||
with patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=metadata)):
|
||||
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"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_load_servers_from_config_pinned_id_jag_endpoint_skips_discovery(self):
|
||||
manager = MCPServerManager()
|
||||
config = {
|
||||
"idjag_cfg": {
|
||||
"url": "https://up.example.com/mcp",
|
||||
"transport": "http",
|
||||
"auth_type": "oauth2_id_jag",
|
||||
"token_exchange_endpoint": "https://idp.example.com/org/token",
|
||||
"id_jag_resource_token_endpoint": "https://pinned.example.com/token",
|
||||
"client_id": "cid",
|
||||
"client_secret": "csec",
|
||||
}
|
||||
}
|
||||
with patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=None)) as mock_discovery:
|
||||
await manager.load_servers_from_config(config)
|
||||
|
||||
mock_discovery.assert_not_awaited()
|
||||
server = next(iter(manager.config_mcp_servers.values()))
|
||||
assert server.id_jag_resource_token_endpoint == "https://pinned.example.com/token"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_as_metadata_parse_carries_grant_profiles(self):
|
||||
manager = MCPServerManager()
|
||||
as_doc = {
|
||||
"issuer": "https://ras.example.com",
|
||||
"token_endpoint": "https://ras.example.com/token",
|
||||
"authorization_grant_profiles_supported": ["urn:ietf:params:oauth:grant-profile:id-jag"],
|
||||
}
|
||||
response = MagicMock()
|
||||
response.raise_for_status = MagicMock()
|
||||
response.json.return_value = as_doc
|
||||
with patch.object(manager, "_fetch_oauth_discovery_url", new=AsyncMock(return_value=response)):
|
||||
metadata = await manager._fetch_single_authorization_server_metadata(
|
||||
"https://ras.example.com", "https://up.example.com/mcp"
|
||||
)
|
||||
|
||||
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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue