From 03f68acf72d712c9f6669d81907169f5956aed4e Mon Sep 17 00:00:00 2001 From: yucheng Date: Sat, 3 Oct 2026 10:48:01 +0000 Subject: [PATCH] fix(mcp): hand the caller sign-in provider a custom-auth caller's own bearer as before the subject split Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/mcp_server_manager.py | 27 ++++++++- .../proxy/_experimental/mcp_server/server.py | 2 +- .../mcp_server/test_mcp_server_manager.py | 44 +++++++++++++++ .../test_mcp_server_tool_calls_and_headers.py | 56 ++++++++++++++++++- 4 files changed, 124 insertions(+), 5 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 9e3a9c2194e..af88f0e1eb6 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -3945,6 +3945,26 @@ class MCPServerManager: return None return bearer + @staticmethod + 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) + bearer: Final = MCPServerManager._extract_bearer_token(oauth2_headers, raw_headers) + if bearer is None or bearer.startswith(LITELLM_VIRTUAL_KEY_PREFIX): + return None + return bearer + def _obo_subject_token( self, server: MCPServer, @@ -4262,7 +4282,10 @@ class MCPServerManager: caller_sign_in_for, ) - if subject_token is None and caller_sign_in_for(server, user_api_key_auth) is not None: + sign_in_subject: Final = self._caller_sign_in_subject_token( + oauth2_headers, raw_headers, user_api_key_auth + ) + 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 ) @@ -6019,7 +6042,7 @@ class MCPServerManager: incoming_bearer_token: Final = ( inbound_authorization[len("bearer ") :] if inbound_authorization.lower().startswith("bearer ") else None ) - incoming_subject_token: Final = self._extract_subject_token(None, raw_headers, user_api_key_auth) + incoming_subject_token: Final = self._caller_sign_in_subject_token(None, raw_headers, user_api_key_auth) 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 e3fc403302a..2893e044a27 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -1785,7 +1785,7 @@ if MCP_AVAILABLE: sign_in = caller_sign_in_for(server, user_api_key_auth) if server is not None else None resource_metadata = get_passthrough_resource_metadata_url(scope, server_name) subject_token = ( - operations.global_mcp_server_manager._extract_subject_token( # pyright: ignore[reportPrivateUsage] # the manager owns the subject/admission filter shared with the preflight + 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 ) if server is not 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 4c3c10474bf..6fd1f384d52 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 @@ -6985,6 +6985,50 @@ class TestMCPServerManager: assert kwargs["incoming_bearer_token"] == expected_bearer assert kwargs["incoming_subject_token"] == expected_subject + @pytest.mark.asyncio + @pytest.mark.parametrize( + ("raw_headers", "api_key", "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( + {"x-litellm-api-key": "ca-key", "authorization": "Bearer ca-key"}, + "ca-key", + 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 + ): + """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.""" + 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 + proxy_logging = MagicMock() + proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={}) + proxy_logging.pre_call_hook = AsyncMock(return_value=None) + + await manager.pre_call_tool_check( + server_name="srv", + name="turn", + arguments={}, + user_api_key_auth=admitted, + proxy_logging_obj=proxy_logging, + server=server, + raw_headers=raw_headers, + ) + + kwargs: Final = proxy_logging._create_mcp_request_object_from_kwargs.call_args.args[0] + assert kwargs["incoming_bearer_token"] == api_key + assert kwargs["incoming_subject_token"] == expected_subject + @pytest.mark.asyncio async def test_check_tool_permission_for_key_team_allows_permitted_tool(self): """ 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 2cc066c6f1b..c061e23159f 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 @@ -12554,7 +12554,9 @@ 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, connecting=_connecting): + async def _connect( + self, route_names, guardrail, allowed, raw_headers=None, connecting=_connecting, user_api_key_auth=None + ): from litellm.proxy._experimental.mcp_server import server as server_module server = _catalog_server() @@ -12575,13 +12577,63 @@ class TestConnectSignInPreflight: 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"), + user_api_key_auth=user_api_key_auth or 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"}, connecting=connecting, ) + @pytest.mark.asyncio + async def test_custom_auth_admitted_bearer_is_pre_flighted_not_challenged(self): + """Custom auth admits the caller on its own IdP token in ``Authorization`` with no + ``x-litellm-api-key``; that token is the sign-in subject, so connect pre-flights it as a tool + call forwarded it before the subject split, instead of challenging for a missing subject.""" + server = _catalog_server() + guardrail = _CallerSignInGuardrail(guardrail_name="sign-in-stub") + admitted = UserAPIKeyAuth(api_key="entra.jwt.token", user_id="u-1") + admitted.authenticated_by_custom_auth = True + 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=admitted, + ) + 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_custom_auth_admitted_virtual_key_bearer_is_still_challenged(self): + server = _catalog_server() + guardrail = _CallerSignInGuardrail(guardrail_name="sign-in-stub") + admitted = UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1") + admitted.authenticated_by_custom_auth = True + litellm.logging_callback_manager.add_litellm_callback(guardrail) + try: + with pytest.raises(HTTPException) as exc: + await self._connect( + ["catalog"], + guardrail, + [server], + raw_headers={"authorization": "Bearer sk-litellm-virtual-key"}, + user_api_key_auth=admitted, + ) + finally: + litellm.logging_callback_manager.remove_callback_from_list_by_object( + litellm.callbacks, guardrail, require_self=False + ) + + assert exc.value.status_code == 401 + assert 'error="invalid_token"' in ((exc.value.headers or {}).get("WWW-Authenticate") or "") + 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