diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 90b70dd01f2..060b7f32135 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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( diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/resolver.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/resolver.py index 69984a56311..0c637cee55c 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/resolver.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/resolver.py @@ -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. 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 4bc5732ec0e..ba82a3d1e10 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/token_endpoint.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/token_endpoint.py @@ -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}") ) diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index 8ae974b19a6..6df8bd9fdfb 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -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): 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 a710da81962..66ea2e08869 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 @@ -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 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 f100bd56f8f..38144d278fe 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 @@ -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" 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 a5cb16822cf..da56843efe2 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 @@ -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"