mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
fcfd7b4f14
commit
0099b96ddc
6 changed files with 308 additions and 10 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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 == []
|
||||
|
|
|
|||
|
|
@ -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 == []
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue