From 55210dd1af854a234657d4f7878f6f703db49aba Mon Sep 17 00:00:00 2001 From: yucheng Date: Sat, 3 Oct 2026 18:02:18 +0000 Subject: [PATCH] fix(mcp): sign-in subject only from a non-LiteLLM-key Bearer, OBO preflight answers before the body Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/caller_sign_in.py | 9 ++- .../mcp_server/mcp_server_manager.py | 24 ++++-- tests/integration/mcp/test_mcp_oauth_flows.py | 37 ++++++++++ .../mcp_server/test_caller_sign_in.py | 40 ++++++++++ .../mcp_server/test_mcp_server_manager.py | 74 +++++++++++++++++++ .../test_mcp_server_tool_calls_and_headers.py | 59 +++++++++++++++ 6 files changed, 234 insertions(+), 9 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/caller_sign_in.py b/litellm/proxy/_experimental/mcp_server/caller_sign_in.py index c4b1d0413c5..d948606dcf1 100644 --- a/litellm/proxy/_experimental/mcp_server/caller_sign_in.py +++ b/litellm/proxy/_experimental/mcp_server/caller_sign_in.py @@ -175,11 +175,12 @@ async def preflight_caller_sign_in( raise_token_exchange_challenge, ) - if not await connecting(): + gating: Final = tuple( + provider for provider in _providers() if provider.caller_sign_in(server, user_api_key_auth) is not None + ) + if not gating or not await connecting(): return - for provider in _providers(): - if provider.caller_sign_in(server, user_api_key_auth) is None: - continue + for provider in gating: match await provider.preflight_caller_sign_in(server, user_api_key_auth, subject_token): case SignedIn(): continue diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 1331d3b01d5..6d93756a797 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -13,6 +13,7 @@ import json import math import os import re +import secrets import time from collections.abc import ( AsyncIterator, @@ -1165,6 +1166,10 @@ def _raw_header_value(raw_headers: Mapping[str, str] | None, name: str) -> str | return next((v for k, v in (raw_headers or {}).items() if isinstance(k, str) and k.lower() == name), None) +def _is_master_key(bearer: str, master_key: str | None) -> bool: + return bool(master_key) and secrets.compare_digest(bearer.encode(), (master_key or "").encode()) + + def _has_explicit_litellm_admission_header(raw_headers: Mapping[str, str] | None) -> bool: """Admission only consumes a non-empty ``x-litellm-api-key``; an empty one falls back to ``Authorization``.""" return bool(_raw_header_value(raw_headers, "x-litellm-api-key")) @@ -3950,11 +3955,20 @@ class MCPServerManager: oauth2_headers: Mapping[str, str] | None, raw_headers: Mapping[str, str] | None, ) -> str | None: - """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): + """The ``Bearer`` credential 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; any other scheme, a LiteLLM key (virtual or master) and a bearer repeating + ``x-litellm-api-key`` are withheld.""" + from litellm.proxy.proxy_server import master_key # noqa: PLC0415 # circular import + + authorization: Final = (oauth2_headers or {}).get("Authorization") or _raw_header_value( + raw_headers, "authorization" + ) + scheme_and_credential: Final = (authorization or "").split(None, 1) + if len(scheme_and_credential) != 2 or scheme_and_credential[0].lower() != "bearer": + return None + bearer: Final = scheme_and_credential[1] + if bearer.startswith(LITELLM_VIRTUAL_KEY_PREFIX) or _is_master_key(bearer, master_key): 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: diff --git a/tests/integration/mcp/test_mcp_oauth_flows.py b/tests/integration/mcp/test_mcp_oauth_flows.py index 0e65ee9d2b1..fdedd62abb1 100644 --- a/tests/integration/mcp/test_mcp_oauth_flows.py +++ b/tests/integration/mcp/test_mcp_oauth_flows.py @@ -2,6 +2,7 @@ import base64 import hashlib import json import secrets +import socket import time import uuid from dataclasses import dataclass @@ -268,6 +269,42 @@ def test_a_malformed_caller_assertion_is_a_sign_in_challenge_not_an_outage(gatew assert tool_calls(peer.drain()) == () +def test_a_refused_token_exchange_answers_before_the_request_body_arrives(gateway: Gateway) -> None: + """The exchange preflight runs on the headers alone, so a client whose body is still in flight gets the + challenge at once instead of the gateway waiting for bytes it will never use.""" + + def entra_like_idp(request: Request) -> Reply: + body: Final = {"error": "invalid_client", "error_codes": [5002723], "error_description": "Invalid JWT token"} + return Reply(status=401, body=json.dumps(body).encode()) + + with mcp_peer() as peer, wire_server(entra_like_idp) as idp, gateway.scenario() as scenario: + alias: Final = "te" + uuid.uuid4().hex[:8] + identity: Final = register_mcp( + scenario, + peer, + alias, + auth_type="oauth2_token_exchange", + token_exchange_endpoint=idp.url + "/token", + credentials={"client_id": "te-client", "client_secret": "te-secret"}, + ) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + origin: Final = urlsplit(str(gateway.client.base_url)) + head: Final = ( + f"POST /mcp/{alias} HTTP/1.1\r\nHost: {origin.netloc}\r\nx-litellm-api-key: {key}\r\n" + "Authorization: Bearer eyJhbGciOiJSUzI1NiJ9.eyJhdWQiOiJ3cm9uZyJ9.c2ln\r\n" + "Accept: application/json, text/event-stream\r\nContent-Type: application/json\r\n" + "Content-Length: 4096\r\n\r\n" + ) + started: Final = time.monotonic() + with socket.create_connection((origin.hostname or "127.0.0.1", origin.port or 80), timeout=10) as raw: + raw.sendall(head.encode()) + status_line: Final = raw.recv(4096).split(b"\r\n", 1)[0] + assert status_line == b"HTTP/1.1 401 Unauthorized", status_line + assert time.monotonic() - started < 5 + assert len(idp.drain()) == 1 + assert tool_calls(peer.drain()) == () + + def _assert_subject_token_challenge(response: httpx.Response, alias: str) -> None: assert response.status_code == 401, response.text challenge: Final = response.headers["www-authenticate"] diff --git a/tests/unit/proxy/_experimental/mcp_server/test_caller_sign_in.py b/tests/unit/proxy/_experimental/mcp_server/test_caller_sign_in.py index 83371400f4c..edd29bf3ece 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_caller_sign_in.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_caller_sign_in.py @@ -1,3 +1,4 @@ +import asyncio from collections.abc import Iterator, Mapping from typing import Final @@ -11,6 +12,7 @@ from litellm.proxy._experimental.mcp_server.caller_sign_in import ( CallerSignInProvider, SignedIn, caller_sign_in_for, + preflight_caller_sign_in, ) from litellm.proxy._types import UserAPIKeyAuth from litellm.types.mcp import MCPAuth, MCPTransport @@ -137,3 +139,41 @@ def test_oauth_utils_strips_the_route_relative_root_path(): "app_root_path": "", } assert get_route_relative_request_path(scope) == "/catalog" # pyright: ignore[reportArgumentType] + + +@pytest.mark.asyncio +async def test_preflight_with_no_gating_provider_never_reads_the_request_body(monkeypatch): + """A plain OBO connect has nothing to pre-flight, so the exchange answers before the body arrives, as it + did before the sign-in seam; reading the body first would stall a client that sends its headers early.""" + monkeypatch.setenv("JWT_ISSUER", "https://jwt-idp.test") + body_read: Final = asyncio.Event() + + async def connecting() -> bool: + body_read.set() + return True + + await preflight_caller_sign_in( + _server(auth_type=MCPAuth.oauth2_token_exchange), + None, + "sub.ject.jws", + root_path="", + resource_metadata=None, + connecting=connecting, + ) + + assert not body_read.is_set() + + +@pytest.mark.asyncio +async def test_preflight_with_a_gating_provider_reads_the_body_to_tell_a_connect_apart(registered): + body_read: Final = asyncio.Event() + + async def connecting() -> bool: + body_read.set() + return True + + await preflight_caller_sign_in( + _server(), None, "sub.ject.jws", root_path="", resource_metadata=None, connecting=connecting + ) + + assert body_read.is_set() 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 5bfd66116a8..ee54f52c2fc 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 @@ -6950,6 +6950,34 @@ class TestMCPServerManager: "eyJ.x.y", id="key-admission-plus-idp-bearer-subject", ), + pytest.param( + {"x-litellm-api-key": "sk-1234", "authorization": "bearer eyJ.x.y"}, + "sk-1234", + "eyJ.x.y", + "eyJ.x.y", + id="lowercase-bearer-scheme-is-the-subject", + ), + pytest.param( + {"x-litellm-api-key": "sk-1234", "authorization": "Basic a.b.c"}, + "sk-1234", + None, + None, + id="non-bearer-scheme-is-not-a-subject", + ), + pytest.param( + {"x-litellm-api-key": "sk-1234", "authorization": "Digest x.y.z"}, + "sk-1234", + None, + None, + id="digest-scheme-is-not-a-subject", + ), + pytest.param( + {"x-litellm-api-key": "sk-1234", "authorization": "eyJ.x.y"}, + "sk-1234", + None, + None, + id="scheme-less-value-is-not-a-subject", + ), ], ) async def test_pre_call_tool_check_separates_raw_bearer_from_subject( @@ -6978,6 +7006,52 @@ class TestMCPServerManager: assert kwargs["incoming_bearer_token"] == expected_bearer assert kwargs["incoming_subject_token"] == expected_subject + @pytest.mark.asyncio + @pytest.mark.parametrize( + ("raw_headers", "expected_subject"), + [ + pytest.param({"authorization": "Bearer gw.master.key"}, None, id="master-key-alone-is-not-a-subject"), + pytest.param( + {"x-litellm-api-key": "sk-1234", "authorization": "Bearer gw.master.key"}, + None, + id="master-key-next-to-a-virtual-key-is-not-a-subject", + ), + pytest.param( + {"x-litellm-api-key": "sk-1234", "authorization": "Bearer gw.other.jws"}, + "gw.other.jws", + id="a-dotted-bearer-that-is-not-the-master-key-is-the-subject", + ), + ], + ) + async def test_pre_call_tool_check_withholds_a_dotted_master_key_from_the_subject( + self, raw_headers, expected_subject + ): + """A master key is a LiteLLM credential whatever its shape, so even one with the two dots of a + compact JWS never becomes the sign-in subject a provider would send to its IdP.""" + manager = MCPServerManager() + server = MCPServer( + server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv", allowed_tools=None + ) + 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) + + with patch("litellm.proxy.proxy_server.master_key", "gw.master.key"): + await manager.pre_call_tool_check( + server_name="srv", + name="turn", + arguments={}, + user_api_key_auth=UserAPIKeyAuth(api_key="sk-1234", user_id="u"), + 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"] == raw_headers["authorization"].removeprefix("Bearer ") + assert kwargs["incoming_subject_token"] == expected_subject + @pytest.mark.asyncio @pytest.mark.parametrize( ("raw_headers", "api_key", "custom_auth", "expected_subject"), 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 01aa5a286a7..7eae64bd160 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 @@ -12584,6 +12584,65 @@ class TestConnectSignInPreflight: connecting=connecting, ) + @pytest.mark.asyncio + @pytest.mark.parametrize( + "authorization", + ["Basic a.b.c", "Digest x.y.z", "eyJ.pay.sig"], + ids=["basic_scheme", "digest_scheme", "scheme_less"], + ) + async def test_authorization_without_a_bearer_scheme_is_challenged_not_pre_flighted(self, authorization): + """Only a ``Bearer`` credential is a sign-in subject; a Basic or Digest value, or a bare string that + merely has two dots, is challenged locally and never handed to a provider's IdP.""" + server = _catalog_server() + guardrail = _CallerSignInGuardrail(guardrail_name="sign-in-stub") + litellm.logging_callback_manager.add_litellm_callback(guardrail) + try: + with pytest.raises(HTTPException) as exc: + await self._connect( + ["catalog"], + guardrail, + [server], + raw_headers={"x-litellm-api-key": "sk-litellm-virtual-key", "authorization": authorization}, + ) + 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 + @pytest.mark.parametrize( + "raw_headers", + [ + {"authorization": "Bearer gw.master.key"}, + {"x-litellm-api-key": "sk-litellm-virtual-key", "authorization": "Bearer gw.master.key"}, + ], + ids=["master_key_alone", "master_key_next_to_a_virtual_key"], + ) + async def test_dotted_master_key_bearer_is_challenged_never_pre_flighted(self, raw_headers): + """The master key is a LiteLLM credential even when it has the two dots of a compact JWS, so the + connect answers the local challenge instead of sending the key to a provider's IdP.""" + server = _catalog_server() + guardrail = _CallerSignInGuardrail(guardrail_name="sign-in-stub") + litellm.logging_callback_manager.add_litellm_callback(guardrail) + try: + with ( + patch("litellm.proxy.proxy_server.master_key", "gw.master.key"), + pytest.raises(HTTPException) as exc, + ): + await self._connect(["catalog"], guardrail, [server], raw_headers=raw_headers) + 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_custom_auth_admitted_bearer_is_pre_flighted_not_challenged(self): """Custom auth admits the caller on its own IdP token in ``Authorization`` with no