diff --git a/litellm/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py index 7bd0a847ad8..11b15a63484 100644 --- a/litellm/experimental_mcp_client/client.py +++ b/litellm/experimental_mcp_client/client.py @@ -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 `` `` removed, or unchanged when absent. + + Callers supply both a bare credential and a complete header value, so prefixing + unconditionally yields ``Bearer Bearer ``. 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 `` 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 diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 26a6f8d1251..dbe97dd5bce 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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( diff --git a/tests/test_litellm/experimental_mcp_client/test_mcp_client.py b/tests/test_litellm/experimental_mcp_client/test_mcp_client.py index 7beb1c43a94..6c3c852395b 100644 --- a/tests/test_litellm/experimental_mcp_client/test_mcp_client.py +++ b/tests/test_litellm/experimental_mcp_client/test_mcp_client.py @@ -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