mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(mcp): normalize auth schemes so MCP egress emits exactly one prefix (#37668)
MCP egress prefixed the configured scheme unconditionally, but callers legitimately supply both a bare token (from a stored credential) and an already-schemed value (passed through from the caller's x-mcp-auth or Authorization header). The second shape produced Authorization: Bearer Bearer <jwt>, which upstream servers reject as a malformed token. It presented intermittently because a resolved stored credential arrives via extra_headers and overwrites the doubled header, so only users without one always failed. strip_auth_scheme drops one leading scheme before the header is rebuilt. It matches the scheme case-insensitively per RFC 7235 and requires a credential behind it, so both a token that merely begins with the scheme text and a scheme with nothing behind it are left intact. MCPAuth.authorization stays verbatim because that auth type means the caller owns the whole header value. For MCPAuth.basic the normalization has to happen in update_auth_value rather than at header-build time: to_basic_auth has already encoded the whole "Basic <credentials>" string by then, so no prefix is left to find. A schemed value whose remainder decodes is already encoded and is reused; one that does not decode is the bare pair with the scheme written in front of it, and is encoded rather than forwarded as an invalid header. The same doubling reached OpenAPI-backed servers through _format_byok_openapi_auth_header. A non-BYOK server short-circuits _resolve_byok_mcp_auth_header, so that formatter also receives the deprecated global x-mcp-auth, which is already a complete header value.
This commit is contained in:
parent
f3639a6fb3
commit
2f23cf5701
3 changed files with 216 additions and 13 deletions
|
|
@ -51,6 +51,42 @@ def to_basic_auth(auth_value: str) -> str:
|
|||
return base64.b64encode(auth_value.encode("utf-8")).decode()
|
||||
|
||||
|
||||
def strip_auth_scheme(auth_value: str, scheme: str) -> str:
|
||||
"""Return ``auth_value`` with a leading ``<scheme> `` removed, or unchanged when absent.
|
||||
|
||||
Callers supply both a bare credential and a complete header value, so prefixing
|
||||
unconditionally yields ``Bearer Bearer <jwt>``. Scheme names are case-insensitive per
|
||||
RFC 7235. A credential is required after the scheme, so both a token that merely begins
|
||||
with the scheme text and a scheme with nothing behind it are returned untouched.
|
||||
Surrounding whitespace is left to ``_strip_header_whitespace`` at header-build time.
|
||||
"""
|
||||
scheme_name, _, remainder = auth_value.lstrip().partition(" ")
|
||||
credential: Final = remainder.lstrip()
|
||||
if credential and scheme_name.lower() == scheme.lower():
|
||||
return credential
|
||||
return auth_value
|
||||
|
||||
|
||||
def to_basic_credentials(auth_value: str) -> str:
|
||||
"""Return the base64 credentials for a ``Basic`` header, encoding only when needed.
|
||||
|
||||
``Basic <credentials>`` carries credentials that are already encoded, so encoding the whole
|
||||
value again would bury the scheme inside the payload. This has to run before
|
||||
:func:`to_basic_auth` rather than at header-build time, where no prefix is left to find.
|
||||
A schemed value whose remainder does not decode is the bare ``username:password`` shape with
|
||||
the scheme written in front of it, and is encoded rather than forwarded as an invalid header;
|
||||
a pair always contains ``:``, which is outside the base64 alphabet, so the two never collide.
|
||||
"""
|
||||
credentials: Final = strip_auth_scheme(auth_value, "Basic")
|
||||
if credentials == auth_value:
|
||||
return to_basic_auth(auth_value)
|
||||
try:
|
||||
base64.b64decode(credentials, validate=True)
|
||||
except ValueError:
|
||||
return to_basic_auth(credentials)
|
||||
return credentials
|
||||
|
||||
|
||||
def _strip_header_whitespace(headers: dict[str, str]) -> dict[str, str]:
|
||||
return {
|
||||
(key.strip() if isinstance(key, str) else key): (value.strip() if isinstance(value, str) else value)
|
||||
|
|
@ -441,16 +477,15 @@ class MCPClient:
|
|||
except BaseException as e:
|
||||
verbose_logger.debug("Error during http_client cleanup: %s", e)
|
||||
|
||||
def update_auth_value(self, mcp_auth_value: str | dict[str, str]):
|
||||
def update_auth_value(self, mcp_auth_value: str | dict[str, str]) -> None:
|
||||
"""
|
||||
Set the authentication header for the MCP client.
|
||||
"""
|
||||
if isinstance(mcp_auth_value, dict):
|
||||
self._mcp_auth_value = mcp_auth_value
|
||||
elif self.auth_type == MCPAuth.basic:
|
||||
self._mcp_auth_value = to_basic_credentials(mcp_auth_value)
|
||||
else:
|
||||
if self.auth_type == MCPAuth.basic:
|
||||
# Assuming mcp_auth_value is in format "username:password", convert it when updating
|
||||
mcp_auth_value = to_basic_auth(mcp_auth_value)
|
||||
self._mcp_auth_value = mcp_auth_value
|
||||
|
||||
def _get_auth_headers(self) -> dict:
|
||||
|
|
@ -459,19 +494,20 @@ class MCPClient:
|
|||
if self._mcp_auth_value:
|
||||
if isinstance(self._mcp_auth_value, str):
|
||||
if self.auth_type == MCPAuth.bearer_token:
|
||||
headers["Authorization"] = f"Bearer {self._mcp_auth_value}"
|
||||
headers["Authorization"] = f"Bearer {strip_auth_scheme(self._mcp_auth_value, 'Bearer')}"
|
||||
elif self.auth_type == MCPAuth.basic:
|
||||
headers["Authorization"] = f"Basic {self._mcp_auth_value}"
|
||||
elif self.auth_type == MCPAuth.api_key:
|
||||
headers["X-API-Key"] = self._mcp_auth_value
|
||||
elif self.auth_type == MCPAuth.authorization:
|
||||
# This auth type means the caller owns the whole header value.
|
||||
headers["Authorization"] = self._mcp_auth_value
|
||||
elif self.auth_type == MCPAuth.oauth2:
|
||||
headers["Authorization"] = f"Bearer {self._mcp_auth_value}"
|
||||
headers["Authorization"] = f"Bearer {strip_auth_scheme(self._mcp_auth_value, 'Bearer')}"
|
||||
elif self.auth_type == MCPAuth.token:
|
||||
headers["Authorization"] = f"token {self._mcp_auth_value}"
|
||||
headers["Authorization"] = f"token {strip_auth_scheme(self._mcp_auth_value, 'token')}"
|
||||
elif self.auth_type == MCPAuth.oauth2_token_exchange:
|
||||
headers["Authorization"] = f"Bearer {self._mcp_auth_value}"
|
||||
headers["Authorization"] = f"Bearer {strip_auth_scheme(self._mcp_auth_value, 'Bearer')}"
|
||||
elif isinstance(self._mcp_auth_value, dict):
|
||||
headers.update(self._mcp_auth_value)
|
||||
# Note: aws_sigv4 auth is not handled here — SigV4 requires per-request
|
||||
|
|
|
|||
|
|
@ -46,7 +46,7 @@ from litellm.constants import (
|
|||
MCP_TOOL_LISTING_TIMEOUT,
|
||||
)
|
||||
from litellm.exceptions import BlockedPiiEntityError, GuardrailRaisedException
|
||||
from litellm.experimental_mcp_client.client import MCPClient, MCPSigV4Auth
|
||||
from litellm.experimental_mcp_client.client import MCPClient, MCPSigV4Auth, strip_auth_scheme
|
||||
from litellm.integrations.custom_guardrail import (
|
||||
_sync_guardrail_info_to_logging_obj, # pyright: ignore[reportPrivateUsage] - the same bridge @log_guardrail_information uses; reimplementing it here would fork the metadata-key logic
|
||||
)
|
||||
|
|
@ -841,12 +841,17 @@ def _without_authorization(
|
|||
|
||||
|
||||
def _format_byok_openapi_auth_header(mcp_server: MCPServer, mcp_auth_header: str) -> str:
|
||||
"""Format a raw BYOK credential for OpenAPI tool ``Authorization`` injection."""
|
||||
"""Format a raw BYOK credential for OpenAPI tool ``Authorization`` injection.
|
||||
|
||||
A non-BYOK server short-circuits ``_resolve_byok_mcp_auth_header``, so the value here can also
|
||||
be the deprecated global ``x-mcp-auth``, which is a complete header value and would otherwise
|
||||
be given a second scheme.
|
||||
"""
|
||||
if mcp_server.auth_type == MCPAuth.api_key:
|
||||
return f"ApiKey {mcp_auth_header}"
|
||||
return f"ApiKey {strip_auth_scheme(mcp_auth_header, 'ApiKey')}"
|
||||
if mcp_server.auth_type == MCPAuth.basic:
|
||||
return f"Basic {mcp_auth_header}"
|
||||
return f"Bearer {mcp_auth_header}"
|
||||
return f"Basic {strip_auth_scheme(mcp_auth_header, 'Basic')}"
|
||||
return f"Bearer {strip_auth_scheme(mcp_auth_header, 'Bearer')}"
|
||||
|
||||
|
||||
def _openapi_forwarded_extra_headers(
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
import asyncio
|
||||
import base64
|
||||
import os
|
||||
import sys
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
|
@ -27,11 +28,16 @@ from litellm.experimental_mcp_client.client import (
|
|||
MCPClient,
|
||||
_as_read_timeout,
|
||||
_first_non_cancelled_cause,
|
||||
strip_auth_scheme,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import (
|
||||
classify_list_exception,
|
||||
list_fault_http_status,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
_format_byok_openapi_auth_header,
|
||||
)
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
from litellm.types.mcp import MCPAuth, MCPStdioConfig, MCPTransport
|
||||
|
||||
|
||||
|
|
@ -887,3 +893,159 @@ async def test_read_timeout_logs_an_actionable_line_that_quiet_on_error_cannot_d
|
|||
assert timeout_lines, f"expected an actionable timeout warning, got {warnings}"
|
||||
assert "http://upstream.local/mcp" in timeout_lines[0], "the line must name the server that stopped answering"
|
||||
assert "0.5s" in timeout_lines[0], "the line must name the budget that elapsed"
|
||||
|
||||
|
||||
class TestAuthSchemeNormalization:
|
||||
"""MCP egress must emit exactly one authorization scheme.
|
||||
|
||||
Callers supply both a bare credential and a complete header value (the latter whenever it is
|
||||
passed through from ``x-mcp-auth`` / ``Authorization``), and the second shape used to be given
|
||||
a second scheme, which upstream servers reject as a malformed token.
|
||||
"""
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"auth_type, auth_value",
|
||||
[
|
||||
(MCPAuth.bearer_token, "bare-token"),
|
||||
(MCPAuth.bearer_token, "Bearer bare-token"),
|
||||
(MCPAuth.bearer_token, "bearer bare-token"),
|
||||
(MCPAuth.bearer_token, " BEARER bare-token"),
|
||||
(MCPAuth.oauth2, "bare-token"),
|
||||
(MCPAuth.oauth2, "Bearer bare-token"),
|
||||
(MCPAuth.oauth2_token_exchange, "bare-token"),
|
||||
(MCPAuth.oauth2_token_exchange, "Bearer bare-token"),
|
||||
],
|
||||
)
|
||||
def test_bearer_family_emits_exactly_one_scheme(self, auth_type, auth_value):
|
||||
client = MCPClient(server_url="http://example.com/mcp", auth_type=auth_type, auth_value=auth_value)
|
||||
|
||||
assert client._get_auth_headers()["Authorization"] == "Bearer bare-token"
|
||||
|
||||
@pytest.mark.parametrize("auth_value", ["bare-token", "token bare-token", "TOKEN bare-token"])
|
||||
def test_token_scheme_emits_exactly_one_scheme(self, auth_value):
|
||||
client = MCPClient(server_url="http://example.com/mcp", auth_type=MCPAuth.token, auth_value=auth_value)
|
||||
|
||||
assert client._get_auth_headers()["Authorization"] == "token bare-token"
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"auth_type, auth_value",
|
||||
[
|
||||
(MCPAuth.bearer_token, "Bearertoken"),
|
||||
(MCPAuth.oauth2, "Bearer.eyJzdWIiOiJhYmMifQ.sig"),
|
||||
(MCPAuth.token, "tokenish"),
|
||||
],
|
||||
)
|
||||
def test_a_credential_merely_starting_with_the_scheme_text_is_left_intact(self, auth_type, auth_value):
|
||||
"""RFC 7235 requires whitespace between scheme and credential, so a token whose first
|
||||
characters happen to spell the scheme is a credential, not a schemed value."""
|
||||
client = MCPClient(server_url="http://example.com/mcp", auth_type=auth_type, auth_value=auth_value)
|
||||
|
||||
scheme = "token" if auth_type == MCPAuth.token else "Bearer"
|
||||
assert client._get_auth_headers()["Authorization"] == f"{scheme} {auth_value}"
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"auth_type, auth_value, expected",
|
||||
[
|
||||
(MCPAuth.bearer_token, "Bearer ", "Bearer Bearer"),
|
||||
(MCPAuth.bearer_token, "Bearer ", "Bearer Bearer"),
|
||||
],
|
||||
)
|
||||
def test_a_scheme_with_no_credential_behind_it_still_produces_a_header(self, auth_type, auth_value, expected):
|
||||
"""Treating this as a schemed value would leave nothing to send, and a request with no
|
||||
Authorization at all is harder to diagnose upstream than a visibly wrong one."""
|
||||
client = MCPClient(server_url="http://example.com/mcp", auth_type=auth_type, auth_value=auth_value)
|
||||
|
||||
assert client._get_auth_headers()["Authorization"] == expected
|
||||
|
||||
def test_basic_with_a_scheme_and_no_credential_still_produces_a_header(self):
|
||||
client = MCPClient(server_url="http://example.com/mcp", auth_type=MCPAuth.basic, auth_value="Basic ")
|
||||
|
||||
assert "Authorization" in client._get_auth_headers()
|
||||
|
||||
def test_basic_accepts_an_already_encoded_schemed_value_without_re_encoding_it(self):
|
||||
"""Stripping the scheme at header-build time cannot fix this shape: ``to_basic_auth`` has by
|
||||
then encoded the whole ``Basic ...`` string, leaving no prefix to find."""
|
||||
encoded = base64.b64encode(b"user:pass").decode()
|
||||
|
||||
client = MCPClient(
|
||||
server_url="http://example.com/mcp",
|
||||
auth_type=MCPAuth.basic,
|
||||
auth_value=f"Basic {encoded}",
|
||||
)
|
||||
|
||||
header = client._get_auth_headers()["Authorization"]
|
||||
assert header == f"Basic {encoded}"
|
||||
assert base64.b64decode(header.split(" ", 1)[1]) == b"user:pass"
|
||||
|
||||
@pytest.mark.parametrize("auth_value", ["user:pass", "Basic user:pass", "basic user:pass"])
|
||||
def test_basic_always_emits_encoded_credentials(self, auth_value):
|
||||
"""A schemed value whose remainder is raw rather than encoded is still a username/password
|
||||
pair, so it is encoded rather than forwarded as an invalid RFC 7617 header."""
|
||||
client = MCPClient(server_url="http://example.com/mcp", auth_type=MCPAuth.basic, auth_value=auth_value)
|
||||
|
||||
header = client._get_auth_headers()["Authorization"]
|
||||
assert base64.b64decode(header.split(" ", 1)[1]) == b"user:pass"
|
||||
|
||||
def test_authorization_auth_type_is_passed_through_verbatim(self):
|
||||
"""``MCPAuth.authorization`` means the caller owns the whole header value."""
|
||||
client = MCPClient(
|
||||
server_url="http://example.com/mcp",
|
||||
auth_type=MCPAuth.authorization,
|
||||
auth_value="Bearer Bearer deliberately-doubled",
|
||||
)
|
||||
|
||||
assert client._get_auth_headers()["Authorization"] == "Bearer Bearer deliberately-doubled"
|
||||
|
||||
def test_api_key_credential_is_not_treated_as_a_schemed_value(self):
|
||||
client = MCPClient(
|
||||
server_url="http://example.com/mcp",
|
||||
auth_type=MCPAuth.api_key,
|
||||
auth_value="Bearer looks-schemed",
|
||||
)
|
||||
|
||||
assert client._get_auth_headers()["X-API-Key"] == "Bearer looks-schemed"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"auth_value, scheme, expected",
|
||||
[
|
||||
("Bearer abc", "Bearer", "abc"),
|
||||
("bearer abc", "Bearer", "abc"),
|
||||
(" Bearer abc ", "Bearer", "abc "),
|
||||
("abc", "Bearer", "abc"),
|
||||
("Bearerabc", "Bearer", "Bearerabc"),
|
||||
("Basic abc", "Bearer", "Basic abc"),
|
||||
("token abc", "token", "abc"),
|
||||
("Basic abc", "Basic", "abc"),
|
||||
("Bearer ", "Bearer", "Bearer "),
|
||||
("Bearer ", "Bearer", "Bearer "),
|
||||
],
|
||||
)
|
||||
def test_strip_auth_scheme(auth_value, scheme, expected):
|
||||
assert strip_auth_scheme(auth_value, scheme) == expected
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"auth_type, auth_value, expected",
|
||||
[
|
||||
(MCPAuth.bearer_token, "Bearer jwt", "Bearer jwt"),
|
||||
(MCPAuth.bearer_token, "jwt", "Bearer jwt"),
|
||||
(MCPAuth.api_key, "ApiKey secret", "ApiKey secret"),
|
||||
(MCPAuth.api_key, "secret", "ApiKey secret"),
|
||||
(MCPAuth.basic, "Basic dXNlcjpwYXNz", "Basic dXNlcjpwYXNz"),
|
||||
],
|
||||
)
|
||||
def test_openapi_byok_auth_header_emits_exactly_one_scheme(auth_type, auth_value, expected):
|
||||
"""A non-BYOK server short-circuits ``_resolve_byok_mcp_auth_header``, so this formatter also
|
||||
receives the deprecated global ``x-mcp-auth``, which is already a complete header value."""
|
||||
server = MCPServer(
|
||||
server_id="s1",
|
||||
name="openapi-server",
|
||||
url="http://example.com/mcp",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=auth_type,
|
||||
spec_path="/tmp/spec.json",
|
||||
)
|
||||
|
||||
assert server.is_byok is False
|
||||
assert _format_byok_openapi_auth_header(server, auth_value) == expected
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue