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>
This commit is contained in:
yucheng 2026-10-03 18:02:18 +00:00
parent 2e118bd672
commit 55210dd1af
6 changed files with 234 additions and 9 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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