mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
2e118bd672
commit
55210dd1af
6 changed files with 234 additions and 9 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue