From 0099b96ddc62a40513455406d46420c9e7cb9ec4 Mon Sep 17 00:00:00 2001 From: yucheng Date: Sun, 27 Sep 2026 02:17:42 +0000 Subject: [PATCH] feat(mcp): challenge rejected caller sign-in subjects at connect Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/caller_sign_in.py | 65 +++++++++++++- .../proxy/_experimental/mcp_server/server.py | 34 +++++-- .../guardrail_hooks/agent_365/agent_365.py | 56 +++++++++++- .../mcp_server/test_caller_sign_in.py | 7 ++ .../mcp_server/test_mcp_server.py | 89 +++++++++++++++++++ .../guardrail_hooks/test_agent_365.py | 67 +++++++++++++- 6 files changed, 308 insertions(+), 10 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/caller_sign_in.py b/litellm/proxy/_experimental/mcp_server/caller_sign_in.py index 7b1d538df5e..13ef8edfb21 100644 --- a/litellm/proxy/_experimental/mcp_server/caller_sign_in.py +++ b/litellm/proxy/_experimental/mcp_server/caller_sign_in.py @@ -17,7 +17,7 @@ from __future__ import annotations import itertools from collections.abc import Mapping from dataclasses import dataclass -from typing import TYPE_CHECKING, Final, Protocol, cast, runtime_checkable +from typing import TYPE_CHECKING, Final, Protocol, assert_never, cast, runtime_checkable from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError @@ -38,6 +38,30 @@ class CallerSignIn: scopes: tuple[str, ...] +@dataclass(frozen=True, slots=True) +class SignedIn: + """The subject token the caller presented satisfies this provider's sign-in.""" + + +@dataclass(frozen=True, slots=True) +class Rejected: + """The caller's identity provider rejected the presented subject token.""" + + detail: str + claims: str | None = None + + +@dataclass(frozen=True, slots=True) +class Unavailable: + """The provider could not reach a verdict; ``fail_open`` is the provider's own fallback policy.""" + + detail: str + fail_open: bool + + +CallerSignInPreflight = SignedIn | Rejected | Unavailable + + @runtime_checkable class CallerSignInProvider(Protocol): def caller_sign_in(self, server: MCPServer, user_api_key_auth: UserAPIKeyAuth | None) -> CallerSignIn | None: @@ -46,6 +70,13 @@ class CallerSignInProvider(Protocol): a challenge).""" ... + async def preflight_caller_sign_in( + self, server: MCPServer, user_api_key_auth: UserAPIKeyAuth | None, subject_token: str + ) -> CallerSignInPreflight: + """Validate ``subject_token`` against this provider at connect time, where a challenge's + ``WWW-Authenticate`` still reaches the client.""" + ... + class _JwtIssuerEntry(BaseModel): model_config = ConfigDict(extra="ignore") @@ -121,3 +152,35 @@ def caller_sign_in_for(server: MCPServer, user_api_key_auth: UserAPIKeyAuth | No issuers: Final = tuple(dict.fromkeys(itertools.chain.from_iterable(c.issuers for c in contributions))) scopes: Final = tuple(dict.fromkeys(itertools.chain.from_iterable(c.scopes for c in contributions))) return CallerSignIn(issuers=issuers, scopes=scopes) + + +async def preflight_caller_sign_in( + server: MCPServer, + user_api_key_auth: UserAPIKeyAuth | None, + subject_token: str, + *, + root_path: str, + connected_as: str | None, +) -> None: + """Run every provider's connect-time check against the subject token, so a bearer the IdP will + reject surfaces as a challenge here rather than a JSON-RPC error at the first tool call.""" + from fastapi import HTTPException # noqa: PLC0415 # lazy: fastapi import stays off the cold path + + from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import ( # noqa: PLC0415 # lazy: adapter pulls MCP subgraph + raise_token_exchange_challenge, + ) + + for provider in _providers(): + if provider.caller_sign_in(server, user_api_key_auth) is None: + continue + match await provider.preflight_caller_sign_in(server, user_api_key_auth, subject_token): + case SignedIn(): + continue + case Rejected(detail=_, claims=claims): + raise_token_exchange_challenge(server, root_path=root_path, claims=claims, connected_as=connected_as) + case Unavailable(detail=detail, fail_open=True): + continue + case Unavailable(detail=detail, fail_open=False): + raise HTTPException(status_code=503, detail=detail) + case _ as verdict: + assert_never(verdict) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 7570ba36905..a4ba9e4d4bb 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -1718,13 +1718,17 @@ if MCP_AVAILABLE: # JSON-RPC error and the WWW-Authenticate header is lost. OBO keeps its connect gate; # guardrail-only gates fire only on a single-server connect the key's grant admits. sign_in = caller_sign_in_for(server, user_api_key_auth) if server is not None else None + subject_token: Final = ( + operations.global_mcp_server_manager._extract_subject_token( # pyright: ignore[reportPrivateUsage] # the manager owns the subject/admission filter shared with the preflight + oauth2_headers, raw_headers, user_api_key_auth + ) + if server is not None + else None + ) if ( server and sign_in is not None - and operations.global_mcp_server_manager._extract_subject_token( # pyright: ignore[reportPrivateUsage] # the manager owns the subject/admission filter shared with the preflight - oauth2_headers, raw_headers, user_api_key_auth - ) - is None + and subject_token is None and ( (server.auth_type == MCPAuth.oauth2_token_exchange and not oauth2_headers) or await _key_granted_single_server(server, mcp_servers, user_api_key_auth, client_ip) @@ -1737,8 +1741,26 @@ if MCP_AVAILABLE: get_request_root_path, ) - raise_token_exchange_challenge( - server, root_path=get_request_root_path(), connected_as=server_name + raise_token_exchange_challenge(server, root_path=get_request_root_path(), connected_as=server_name) + if ( + server + and sign_in is not None + and subject_token is not None + and await _key_granted_single_server(server, mcp_servers, user_api_key_auth, client_ip) + ): + from litellm.proxy._experimental.mcp_server.caller_sign_in import ( # noqa: PLC0415 # lazy: provider discovery pulls the guardrail registry + preflight_caller_sign_in, + ) + from litellm.proxy.middleware.per_request_root_path_middleware import ( # noqa: PLC0415 # lazy: middleware imports proxy utils + get_request_root_path, + ) + + await preflight_caller_sign_in( + server, + user_api_key_auth, + subject_token, + root_path=get_request_root_path(), + connected_as=server_name, ) # Exchange-backed modes (token_exchange's OBO mint, id_jag's stored-assertion mint): run diff --git a/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py b/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py index b33c61951d0..8bbb4c54b03 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py +++ b/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py @@ -31,13 +31,21 @@ from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, ) -from litellm.proxy._experimental.mcp_server.caller_sign_in import CallerSignIn -from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Error, Ok +from litellm.proxy._experimental.mcp_server.caller_sign_in import ( + CallerSignIn, + CallerSignInPreflight, + Rejected, + SignedIn, + Unavailable, +) +from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import OAuthToken +from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Error, Ok, Result from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchange_provider import ( build_token_exchanger, ) from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchanger import TokenExchanger from litellm.proxy._experimental.mcp_server.outbound_credentials.types import ( + CredError, ServerSpec, TokenExchangeConfig, ) @@ -457,6 +465,50 @@ class Agent365Guardrail(CustomGuardrail): scopes=(GATEWAY_SCOPE_TEMPLATE.format(client_id=self.client_id),), ) + async def _exchange_caller_assertion(self, assertion: str) -> Result[OAuthToken, CredError]: + """The one Entra OBO exchange call both the tool-call path and the connect preflight run; a + successful result is cached by the exchanger, so the session reuses what the preflight minted.""" + return await self._token_exchanger.exchange( + assertion, self._exchange_server, self._exchange_config, tenant_id=self.tenant_id + ) + + async def preflight_caller_sign_in( + self, server: MCPServer, user_api_key_auth: "UserAPIKeyAuth | None", subject_token: str + ) -> CallerSignInPreflight: + """The connect-time check the preemptive gate runs: a bearer Entra rejects gets the sign-in + challenge here, where ``WWW-Authenticate`` still reaches the client, instead of surfacing as a + JSON-RPC error on every tools/call. ``subject_token=None`` stays the challenge gate's job.""" + assertion: Final = entra_assertion(subject_token) + if assertion is None: + return SignedIn() + try: + exchange_result: Final = await self._exchange_caller_assertion(assertion) + except (httpx.HTTPError, LitellmTimeout, TimeoutError) as exc: + return Unavailable( + detail=f"the Entra token endpoint could not be reached ({type(exc).__name__})", + fail_open=self.unreachable_fallback == "fail_open", + ) + match exchange_result: + case Ok(_): + return SignedIn() + case Error(error): + match error.tag: + case "unauthorized": + return Rejected(detail=error.unauthorized.detail, claims=error.unauthorized.claims) + case "misconfigured": + return Unavailable( + detail=( + f"Entra rejected the gateway's own Agent 365 credentials ({error.misconfigured}); " + "check the guardrail's client_id, client_secret and resource_app_id" + ), + fail_open=self.unreachable_fallback == "fail_open", + ) + case _: + return Unavailable( + detail=f"the Entra token exchange failed ({error.summary})", + fail_open=self.unreachable_fallback == "fail_open", + ) + async def _post_allowing_error_status( self, url: str, diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_caller_sign_in.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_caller_sign_in.py index aa853696462..0e66500fad0 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_caller_sign_in.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_caller_sign_in.py @@ -7,7 +7,9 @@ import litellm from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.proxy._experimental.mcp_server.caller_sign_in import ( CallerSignIn, + CallerSignInPreflight, CallerSignInProvider, + SignedIn, caller_sign_in_for, ) from litellm.proxy._types import UserAPIKeyAuth @@ -29,6 +31,11 @@ class _SignInGuardrail(CustomGuardrail): return None return CallerSignIn(issuers=(self.issuer,), scopes=(self.scope,)) + async def preflight_caller_sign_in( + self, server: MCPServer, user_api_key_auth: UserAPIKeyAuth | None, subject_token: str + ) -> CallerSignInPreflight: + return SignedIn() + class _PlainGuardrail(CustomGuardrail): def __init__(self) -> None: diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 0c4c6ce1ab8..bfaa1767456 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -11250,3 +11250,92 @@ class TestConnectChallengeResolver: authenticate: Final = (exc.value.headers or {}).get("WWW-Authenticate") or "" assert authenticate.startswith(f'Bearer resource_metadata="/.well-known/oauth-protected-resource/mcp/{route_name}"') + + +class TestConnectSignInPreflight: + """A subject token the provider rejects must fail at connect with the RFC 9728 challenge, not as a + JSON-RPC error on every tools/call where ``WWW-Authenticate`` is lost.""" + + async def _connect(self, route_names, guardrail, allowed, raw_headers=None): + from litellm.proxy._experimental.mcp_server import server as server_module + + server = _catalog_server() + with ( + patch.object( + mcp_operations.global_mcp_server_manager, + "get_filtered_registry", + return_value={server.server_id: server}, + ), + patch.object( + mcp_operations, + "_get_allowed_mcp_servers", + AsyncMock(return_value=allowed), + ), + ): + await server_module._raise_preemptive_401_for_unauthenticated_servers( + scope={"type": "http", "method": "POST", "path": f"/mcp/{route_names[0]}", "headers": []}, + mcp_servers=list(route_names), + oauth2_headers=None, + mcp_server_auth_headers=None, + user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"), + client_ip=None, + raw_headers=raw_headers + or {"x-litellm-api-key": "sk-litellm-virtual-key", "authorization": "Bearer entra.jwt.token"}, + ) + + @pytest.mark.asyncio + async def test_rejected_subject_challenges_at_connect(self): + from litellm.proxy._experimental.mcp_server.caller_sign_in import Rejected + + server = _catalog_server() + guardrail = _CallerSignInGuardrail( + guardrail_name="sign-in-stub", preflight_result=Rejected("the Entra OBO exchange was rejected (AADSTS70002)") + ) + litellm.logging_callback_manager.add_litellm_callback(guardrail) + try: + with pytest.raises(HTTPException) as exc: + await self._connect(["catalog"], guardrail, [server]) + finally: + litellm.logging_callback_manager.remove_callback_from_list_by_object( + litellm.callbacks, guardrail, require_self=False + ) + + assert exc.value.status_code == 401 + authenticate: Final = (exc.value.headers or {}).get("WWW-Authenticate") or "" + assert 'resource_metadata="/.well-known/oauth-protected-resource/mcp/catalog"' in authenticate + assert 'error="invalid_token"' in authenticate + assert guardrail.preflight_calls == ["entra.jwt.token"] + + @pytest.mark.asyncio + async def test_unavailable_fail_closed_answers_503(self): + from litellm.proxy._experimental.mcp_server.caller_sign_in import Unavailable + + server = _catalog_server() + guardrail = _CallerSignInGuardrail( + guardrail_name="sign-in-stub", preflight_result=Unavailable("the Entra token endpoint could not be reached", fail_open=False) + ) + litellm.logging_callback_manager.add_litellm_callback(guardrail) + try: + with pytest.raises(HTTPException) as exc: + await self._connect(["catalog"], guardrail, [server]) + finally: + litellm.logging_callback_manager.remove_callback_from_list_by_object( + litellm.callbacks, guardrail, require_self=False + ) + + assert exc.value.status_code == 503 + assert exc.value.detail == "the Entra token endpoint could not be reached" + + @pytest.mark.asyncio + async def test_multi_server_connect_never_awaits_the_preflight(self): + server = _catalog_server() + guardrail = _CallerSignInGuardrail(guardrail_name="sign-in-stub") + litellm.logging_callback_manager.add_litellm_callback(guardrail) + try: + await self._connect(["catalog", "other"], guardrail, [server]) + finally: + litellm.logging_callback_manager.remove_callback_from_list_by_object( + litellm.callbacks, guardrail, require_self=False + ) + + assert guardrail.preflight_calls == [] diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py index 47f1d9f5f0c..d64ec31574a 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py @@ -14,7 +14,13 @@ from litellm.exceptions import Timeout as LitellmTimeout from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.litellm_core_utils.secret_redaction import redact_string -from litellm.proxy._experimental.mcp_server.caller_sign_in import CallerSignIn, caller_sign_in_for +from litellm.proxy._experimental.mcp_server.caller_sign_in import ( + CallerSignIn, + Rejected, + SignedIn, + Unavailable, + caller_sign_in_for, +) from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import OAuthToken from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Error, Ok, Result from litellm.proxy._experimental.mcp_server.outbound_credentials.types import ( @@ -1214,3 +1220,62 @@ class TestCallerSignIn: "https://login.microsoftonline.com/tenant-abc/v2.0", ) assert sign_in.scopes == ("read", "api://client-xyz/access_as_user") + + +class TestPreflightCallerSignIn: + """The connect-time check must give the connect gate a verdict it can challenge on: a rejected + subject becomes the RFC 9728 challenge, an unreachable endpoint the guardrail's fallback policy.""" + + @pytest.mark.asyncio + async def test_ok_exchange_signs_in(self): + exchanger: Final = StubTokenExchanger(_obo_ok()) + guardrail: Final = _make_guardrail(FakeHandler([]), exchanger=exchanger) + + verdict: Final = await guardrail.preflight_caller_sign_in(_server(), _user(), FAKE_ASSERTION) + + assert verdict == SignedIn() + assert [call[0] for call in exchanger.calls] == [FAKE_ASSERTION] + + @pytest.mark.asyncio + async def test_unauthorized_error_rejects_with_the_idp_detail(self): + exchanger: Final = StubTokenExchanger( + [Error(CredError.of_unauthorized("the provided assertion has expired", claims="step-up"))] + ) + guardrail: Final = _make_guardrail(FakeHandler([]), exchanger=exchanger) + + verdict: Final = await guardrail.preflight_caller_sign_in(_server(), _user(), FAKE_ASSERTION) + + assert verdict == Rejected(detail="the provided assertion has expired", claims="step-up") + + @pytest.mark.asyncio + async def test_misconfigured_fail_closed_is_unavailable(self): + exchanger: Final = StubTokenExchanger([Error(CredError.of_misconfigured("bad client_secret"))]) + guardrail: Final = _make_guardrail(FakeHandler([]), exchanger=exchanger, unreachable_fallback="fail_closed") + + verdict: Final = await guardrail.preflight_caller_sign_in(_server(), _user(), FAKE_ASSERTION) + + assert isinstance(verdict, Unavailable) + assert verdict.fail_open is False + assert "bad client_secret" in verdict.detail + + @pytest.mark.asyncio + async def test_endpoint_unreachable_fail_open_is_unavailable(self): + exchanger: Final = StubTokenExchanger( + [httpx.ConnectError("refused", request=httpx.Request("POST", "https://example.test"))] + ) + guardrail: Final = _make_guardrail(FakeHandler([]), exchanger=exchanger, unreachable_fallback="fail_open") + + verdict: Final = await guardrail.preflight_caller_sign_in(_server(), _user(), FAKE_ASSERTION) + + assert isinstance(verdict, Unavailable) + assert verdict.fail_open is True + + @pytest.mark.asyncio + async def test_non_assertion_subject_signs_in_without_exchanging(self): + exchanger: Final = StubTokenExchanger() + guardrail: Final = _make_guardrail(FakeHandler([]), exchanger=exchanger) + + verdict: Final = await guardrail.preflight_caller_sign_in(_server(), _user(), "opaque-bearer") + + assert verdict == SignedIn() + assert exchanger.calls == []