diff --git a/litellm/proxy/_experimental/mcp_server/caller_sign_in.py b/litellm/proxy/_experimental/mcp_server/caller_sign_in.py index 1d6d75b438f..c4b1d0413c5 100644 --- a/litellm/proxy/_experimental/mcp_server/caller_sign_in.py +++ b/litellm/proxy/_experimental/mcp_server/caller_sign_in.py @@ -164,16 +164,19 @@ async def preflight_caller_sign_in( resource_metadata: str | None, connecting: Callable[[], Awaitable[bool]], ) -> 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. A fail-closed - provider outage is the connect's 503 only while ``connecting``; on an open session the tool-call hook - answers it inside the JSON-RPC envelope, with its guardrail Logs row.""" + """Run every provider's connect-time check against the subject token while ``connecting``, so a bearer + the IdP will reject surfaces as a challenge here rather than a JSON-RPC error at the first tool call, and + a fail-closed provider outage is the connect's 503. On an open session nothing is exchanged here: the + tool-call hook runs the one exchange and answers inside the JSON-RPC envelope, with its guardrail Logs + row.""" 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, ) + if not await connecting(): + return for provider in _providers(): if provider.caller_sign_in(server, user_api_key_auth) is None: continue @@ -187,7 +190,6 @@ async def preflight_caller_sign_in( case Unavailable(fail_open=True): continue case Unavailable(detail=detail, fail_open=False): - if await connecting(): - raise HTTPException(status_code=503, detail=detail) + raise HTTPException(status_code=503, detail=detail) case _ as verdict: assert_never(verdict) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index af88f0e1eb6..1331d3b01d5 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -3949,20 +3949,16 @@ class MCPServerManager: def _caller_sign_in_subject_token( oauth2_headers: Mapping[str, str] | None, raw_headers: Mapping[str, str] | None, - user_api_key_auth: UserAPIKeyAuth | None, ) -> str | None: - """The bearer a caller sign-in provider validates: custom auth admits the caller on its own IdP token - in ``Authorization``, so that token is the subject there and only a virtual key is withheld.""" - admitted_on_own_bearer: Final = ( - user_api_key_auth is not None - and user_api_key_auth.authenticated_by_custom_auth - and not _has_explicit_litellm_admission_header(raw_headers) - ) - if not admitted_on_own_bearer: - return MCPServerManager._extract_subject_token(oauth2_headers, raw_headers, user_api_key_auth) + """The bearer a caller sign-in provider validates. An admission that consumed ``Authorization`` (custom + auth, built-in OAuth2, JWT) did so on the caller's own IdP token, so that token is the subject; only a + virtual key, or a bearer repeating ``x-litellm-api-key``, is withheld.""" bearer: Final = MCPServerManager._extract_bearer_token(oauth2_headers, raw_headers) if bearer is None or bearer.startswith(LITELLM_VIRTUAL_KEY_PREFIX): return None + admission_header: Final = _raw_header_value(raw_headers, "x-litellm-api-key") + if admission_header and strip_auth_scheme(admission_header, "Bearer") == bearer: + return None return bearer def _obo_subject_token( @@ -4282,9 +4278,7 @@ class MCPServerManager: caller_sign_in_for, ) - sign_in_subject: Final = self._caller_sign_in_subject_token( - oauth2_headers, raw_headers, user_api_key_auth - ) + sign_in_subject: Final = self._caller_sign_in_subject_token(oauth2_headers, raw_headers) if sign_in_subject is None and caller_sign_in_for(server, user_api_key_auth) is not None: raise_token_exchange_challenge( server, root_path=get_request_root_path(), resource_metadata=resource_metadata @@ -6042,7 +6036,7 @@ class MCPServerManager: incoming_bearer_token: Final = ( inbound_authorization[len("bearer ") :] if inbound_authorization.lower().startswith("bearer ") else None ) - incoming_subject_token: Final = self._caller_sign_in_subject_token(None, raw_headers, user_api_key_auth) + incoming_subject_token: Final = self._caller_sign_in_subject_token(None, raw_headers) pre_hook_kwargs: Final = { "guardrail_context": guardrail_context, diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 2893e044a27..ddbbe5c638f 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -1786,7 +1786,7 @@ if MCP_AVAILABLE: resource_metadata = get_passthrough_resource_metadata_url(scope, server_name) subject_token = ( operations.global_mcp_server_manager._caller_sign_in_subject_token( # pyright: ignore[reportPrivateUsage] # the manager owns the subject/admission filter shared with the preflight - oauth2_headers, raw_headers, user_api_key_auth + oauth2_headers, raw_headers ) if server is not None else None diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 6fd1f384d52..23130f060b1 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -6936,13 +6936,6 @@ class TestMCPServerManager: @pytest.mark.parametrize( ("raw_headers", "api_key", "expected_bearer", "expected_subject"), [ - pytest.param( - {"authorization": "Bearer eyJ.x.y"}, - "eyJ.x.y", - "eyJ.x.y", - None, - id="idp-token-as-admission-stays-raw-bearer", - ), pytest.param( {"authorization": "Bearer sk-1234"}, "sk-1234", @@ -6987,29 +6980,42 @@ class TestMCPServerManager: @pytest.mark.asyncio @pytest.mark.parametrize( - ("raw_headers", "api_key", "expected_subject"), + ("raw_headers", "api_key", "custom_auth", "expected_subject"), [ - pytest.param({"authorization": "Bearer eyJ.x.y"}, "eyJ.x.y", "eyJ.x.y", id="idp-bearer-is-the-subject"), - pytest.param({"authorization": "Bearer sk-1234"}, "sk-1234", None, id="virtual-key-is-not-a-subject"), + pytest.param( + {"authorization": "Bearer eyJ.x.y"}, "eyJ.x.y", True, "eyJ.x.y", id="custom-auth-idp-bearer-is-the-subject" + ), + pytest.param( + {"authorization": "Bearer eyJ.x.y"}, + "eyJ.x.y", + False, + "eyJ.x.y", + id="built-in-oauth2-admission-bearer-is-the-subject", + ), + pytest.param( + {"authorization": "Bearer sk-1234"}, "sk-1234", True, None, id="virtual-key-is-not-a-subject" + ), pytest.param( {"x-litellm-api-key": "ca-key", "authorization": "Bearer ca-key"}, "ca-key", + True, None, id="explicit-key-admission-repeated-in-authorization-is-not-a-subject", ), ], ) - async def test_pre_call_tool_check_hands_sign_in_the_bearer_custom_auth_admitted( - self, raw_headers, api_key, expected_subject + async def test_pre_call_tool_check_hands_sign_in_the_bearer_that_admitted_the_caller( + self, raw_headers, api_key, custom_auth, expected_subject ): - """Custom auth admits the caller on its own IdP token in ``Authorization`` with no - ``x-litellm-api-key``, so that token is the sign-in subject as it was before the subject split.""" + """Custom auth and the built-in OAuth2 admission both admit the caller on its own IdP token in + ``Authorization`` with no ``x-litellm-api-key`` and record it as ``api_key``; that token is the + sign-in subject as the raw bearer was before the subject split.""" manager = MCPServerManager() server = MCPServer( server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv", allowed_tools=None ) admitted = UserAPIKeyAuth(api_key=api_key, user_id="u") - admitted.authenticated_by_custom_auth = True + admitted.authenticated_by_custom_auth = custom_auth proxy_logging = MagicMock() proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={}) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py index c061e23159f..01aa5a286a7 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py @@ -3974,7 +3974,7 @@ async def test_owner_mismatch_on_a_torn_down_session_is_refused_before_the_body_ litellm.callbacks, guardrail, require_self=False ) - assert guardrail.preflight_calls == ["entra.jwt.token"] + assert guardrail.preflight_calls == [] handle_request_mock.assert_not_awaited() statuses = [m["status"] for m in sent_messages if m.get("type") == "http.response.start"] assert statuses == [403] @@ -12634,6 +12634,53 @@ class TestConnectSignInPreflight: assert 'error="invalid_token"' in ((exc.value.headers or {}).get("WWW-Authenticate") or "") assert guardrail.preflight_calls == [] + @pytest.mark.asyncio + async def test_built_in_oauth2_admitted_bearer_is_pre_flighted_not_challenged(self): + """The built-in OAuth2 admission records the caller's token as ``api_key`` without the custom-auth + marker; that token is still the sign-in subject, so connect pre-flights it instead of challenging.""" + server = _catalog_server() + guardrail = _CallerSignInGuardrail(guardrail_name="sign-in-stub") + litellm.logging_callback_manager.add_litellm_callback(guardrail) + try: + await self._connect( + ["catalog"], + guardrail, + [server], + raw_headers={"authorization": "Bearer entra.jwt.token"}, + user_api_key_auth=UserAPIKeyAuth(api_key="entra.jwt.token", user_id="u-1"), + ) + finally: + litellm.logging_callback_manager.remove_callback_from_list_by_object( + litellm.callbacks, guardrail, require_self=False + ) + + assert guardrail.preflight_calls == ["entra.jwt.token"] + + @pytest.mark.asyncio + async def test_rejected_subject_on_an_open_session_is_left_to_the_tool_call_hook(self): + """Only the connect pre-flights the subject. On an open session the gate must not exchange at all, so + the tool-call hook runs the one exchange and answers a rejection inside the JSON-RPC envelope with its + guardrail Logs row, as base did.""" + from litellm.proxy._experimental.mcp_server.caller_sign_in import Rejected + + async def _open_session() -> bool: + return False + + server = _catalog_server() + guardrail = _CallerSignInGuardrail( + guardrail_name="sign-in-stub", + preflight_result=Rejected("the Entra OBO exchange was rejected (AADSTS5002723)"), + ) + litellm.logging_callback_manager.add_litellm_callback(guardrail) + try: + await self._connect(["catalog"], guardrail, [server], connecting=_open_session) + finally: + litellm.logging_callback_manager.remove_callback_from_list_by_object( + litellm.callbacks, guardrail, require_self=False + ) + + assert guardrail.preflight_calls == [] + @pytest.mark.asyncio async def test_rejected_subject_challenges_at_connect(self): from litellm.proxy._experimental.mcp_server.caller_sign_in import Rejected @@ -12691,8 +12738,9 @@ class TestConnectSignInPreflight: ) async def test_fail_closed_outage_is_the_connects_503_only_on_initialize(self, rpc_method, session_id, reaches): """Only the ``initialize`` POST turns a fail-closed provider outage into the connect's 503. Every other - JSON-RPC POST on the gated route must reach the session manager with its body intact, so the tools/call - hook answers the outage inside the result envelope and writes the guardrail Logs row, as base did.""" + JSON-RPC POST on the gated route must reach the session manager with its body intact and without a gate + exchange, so the tools/call hook runs the one exchange, answers the outage inside the result envelope and + writes the guardrail Logs row, as base did.""" from litellm.proxy._experimental.mcp_server import server as server_module from litellm.proxy._experimental.mcp_server.caller_sign_in import Unavailable @@ -12766,7 +12814,7 @@ class TestConnectSignInPreflight: if session_id: server_module._remove_stateful_session_tracking(session_id) - assert guardrail.preflight_calls == ["entra.jwt.token"] + assert guardrail.preflight_calls == (["entra.jwt.token"] if reaches is None else []) assert delivered == ({} if reaches is None else {reaches: body}) send.assert_not_awaited()