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>
This commit is contained in:
yucheng 2026-10-03 10:48:01 +00:00
parent 26715ba667
commit 03f68acf72
4 changed files with 124 additions and 5 deletions

View file

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

View file

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

View file

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

View file

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