mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(mcp): strip inbound auth scheme case-insensitively before token exchange (#39346)
* fix(mcp): strip inbound auth scheme case-insensitively before token exchange Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(mcp): type the fake credential provider params in token exchange scheme tests Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yassin <yassin@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
bed4086619
commit
6da516e6f3
4 changed files with 70 additions and 8 deletions
|
|
@ -70,7 +70,7 @@ def to_basic_auth(auth_value: str) -> str:
|
|||
|
||||
|
||||
def strip_auth_scheme(auth_value: str, scheme: str) -> str:
|
||||
"""Return ``auth_value`` with a leading ``<scheme> `` removed, or unchanged when absent.
|
||||
"""Return ``auth_value`` with a leading ``<scheme>`` and separator 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
|
||||
|
|
@ -78,10 +78,9 @@ def strip_auth_scheme(auth_value: str, scheme: str) -> str:
|
|||
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
|
||||
parts: Final = auth_value.split(None, 1)
|
||||
if len(parts) == 2 and parts[0].lower() == scheme.lower():
|
||||
return parts[1]
|
||||
return auth_value
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -3346,9 +3346,7 @@ class MCPServerManager:
|
|||
normalized: Final = {k.lower(): v for k, v in raw_headers.items()}
|
||||
auth_value = normalized.get("authorization")
|
||||
if auth_value:
|
||||
if auth_value.startswith("Bearer "):
|
||||
return auth_value[len("Bearer ") :]
|
||||
return auth_value
|
||||
return strip_auth_scheme(auth_value, "Bearer")
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -1014,6 +1014,8 @@ class TestAuthSchemeNormalization:
|
|||
[
|
||||
("Bearer abc", "Bearer", "abc"),
|
||||
("bearer abc", "Bearer", "abc"),
|
||||
("Bearer\tabc", "Bearer", "abc"),
|
||||
("bearer\t\tabc", "Bearer", "abc"),
|
||||
(" Bearer abc ", "Bearer", "abc "),
|
||||
("abc", "Bearer", "abc"),
|
||||
("Bearerabc", "Bearer", "Bearerabc"),
|
||||
|
|
|
|||
|
|
@ -2616,6 +2616,69 @@ class TestMCPServerManager:
|
|||
await manager.preflight_token_exchange(server=server, oauth2_headers=None, user_api_key_auth=None)
|
||||
assert resolved == ["good-subject"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"authorization",
|
||||
[
|
||||
"Bearer subj-jwt",
|
||||
"bearer subj-jwt",
|
||||
"BEARER subj-jwt",
|
||||
"Bearer\tsubj-jwt",
|
||||
"Bearer subj-jwt",
|
||||
],
|
||||
)
|
||||
async def test_preflight_token_exchange_strips_inbound_authorization_scheme(self, authorization):
|
||||
"""The resolver posts inbound_token verbatim as subject_token."""
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import (
|
||||
StaticHeaderAuth,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Ok
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import CredError, ServerSpec, Subject
|
||||
|
||||
resolved: Final[list[str | None]] = []
|
||||
|
||||
class _FakeProvider:
|
||||
async def resolve_credentials(self, subject: Subject, server: ServerSpec) -> Ok[StaticHeaderAuth, CredError]:
|
||||
resolved.append(subject.inbound_token.get_secret_value() if subject.inbound_token else None)
|
||||
return Ok(StaticHeaderAuth("Bearer MINTED", header_name="Authorization"))
|
||||
|
||||
manager = MCPServerManager(cred_provider=_FakeProvider())
|
||||
server = self._token_exchange_server(f"te-preflight-auth-{authorization!r}")
|
||||
|
||||
await manager.preflight_token_exchange(
|
||||
server=server,
|
||||
oauth2_headers={"Authorization": authorization},
|
||||
user_api_key_auth=None,
|
||||
)
|
||||
|
||||
assert resolved == ["subj-jwt"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_preflight_token_exchange_preserves_authorization_without_separator(self):
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import (
|
||||
StaticHeaderAuth,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Ok
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import CredError, ServerSpec, Subject
|
||||
|
||||
resolved: Final[list[str | None]] = []
|
||||
|
||||
class _FakeProvider:
|
||||
async def resolve_credentials(self, subject: Subject, server: ServerSpec) -> Ok[StaticHeaderAuth, CredError]:
|
||||
resolved.append(subject.inbound_token.get_secret_value() if subject.inbound_token else None)
|
||||
return Ok(StaticHeaderAuth("Bearer MINTED", header_name="Authorization"))
|
||||
|
||||
manager = MCPServerManager(cred_provider=_FakeProvider())
|
||||
server = self._token_exchange_server("te-preflight-auth-no-separator")
|
||||
|
||||
await manager.preflight_token_exchange(
|
||||
server=server,
|
||||
oauth2_headers={"Authorization": "Bearersubj-jwt"},
|
||||
user_api_key_auth=None,
|
||||
)
|
||||
|
||||
assert resolved == ["Bearersubj-jwt"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_preflight_token_exchange_skips_discovery_for_other_auth_modes(self):
|
||||
"""Preflight must not make unrelated auth modes depend on OAuth discovery."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue