feat(mcp): challenge rejected caller sign-in subjects at connect

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-09-27 02:17:42 +00:00
parent fcfd7b4f14
commit 0099b96ddc
6 changed files with 308 additions and 10 deletions

View file

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

View file

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

View file

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

View file

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

View file

@ -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 == []

View file

@ -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 == []