mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(mcp): preserve explicit caller authorization credentials
This commit is contained in:
parent
84e14789d1
commit
be506936bd
2 changed files with 36 additions and 6 deletions
|
|
@ -36,17 +36,16 @@ def _usable_credential_value(auth_type: MCPAuthType, name: str, value: str) -> b
|
|||
return True
|
||||
|
||||
|
||||
def validate_static_credential(
|
||||
server: MCPServer, headers: Mapping[str, str], *, header_slot: str | None = None, openapi: bool = False
|
||||
) -> Result[None, CredError]:
|
||||
def validate_static_credential(server: MCPServer, headers: Mapping[str, str]) -> Result[None, CredError]:
|
||||
if server.auth_type not in _STATIC_MODES or server.transport == MCPTransport.stdio:
|
||||
return Ok(None)
|
||||
default_slot: Final = "X-API-Key" if server.auth_type == MCPAuth.api_key else "Authorization"
|
||||
slots: Final = frozenset(
|
||||
name.lower()
|
||||
for name in (
|
||||
header_slot or server.upstream_token_header or default_slot,
|
||||
"Authorization" if openapi else default_slot,
|
||||
server.upstream_token_header or default_slot,
|
||||
default_slot,
|
||||
"Authorization",
|
||||
)
|
||||
)
|
||||
values: Final = tuple((name.lower(), value.strip()) for name, value in headers.items() if name.lower() in slots)
|
||||
|
|
@ -75,7 +74,7 @@ def validate_openapi_credentials(
|
|||
headers: Final = merge_openapi_headers(
|
||||
server.static_headers or {}, forwarded_headers, caller_authorization, resolved_headers
|
||||
)
|
||||
match validate_static_credential(server, headers, openapi=True):
|
||||
match validate_static_credential(server, headers):
|
||||
case Error(error):
|
||||
raise_public(error)
|
||||
case Ok():
|
||||
|
|
|
|||
|
|
@ -13693,3 +13693,34 @@ class TestProtectedCredentialPreparation:
|
|||
with pytest.raises(HTTPException) as exc:
|
||||
await MCPServerManager()._create_mcp_client(server, extra_headers={"X-Custom": "", "X-API-Key": ""})
|
||||
assert exc.value.status_code == 500
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("custom_slot", [None, "X-Custom"])
|
||||
@pytest.mark.parametrize("source", ["caller", "forwarded"])
|
||||
async def test_api_key_preserves_explicit_authorization_credential(
|
||||
self, custom_slot: str | None, source: str
|
||||
) -> None:
|
||||
server: Final = MCPServer(
|
||||
server_id="caller-auth", name="caller-auth", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=MCPAuth.api_key, upstream_token_header=custom_slot,
|
||||
)
|
||||
headers: Final = {"Authorization": "Bearer caller-credential", "X-API-Key": ""}
|
||||
client: Final = await MCPServerManager()._create_mcp_client(
|
||||
server, mcp_auth_header=headers if source == "caller" else None,
|
||||
extra_headers=headers if source == "forwarded" else None,
|
||||
)
|
||||
request: Final = await client.prepare_request_auth()
|
||||
assert request.headers["Authorization"] == "Bearer caller-credential"
|
||||
assert request.headers["X-API-Key"] == ""
|
||||
assert custom_slot is None or custom_slot not in request.headers
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("value", ["", " ", "Bearer", "Basic", "token", "ApiKey"])
|
||||
async def test_api_key_rejects_authorization_without_a_credential(self, value: str) -> None:
|
||||
server: Final = MCPServer(
|
||||
server_id="caller-empty", name="caller-empty", url="https://upstream.example/mcp",
|
||||
transport=MCPTransport.http, auth_type=MCPAuth.api_key,
|
||||
)
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await MCPServerManager()._create_mcp_client(server, mcp_auth_header={"Authorization": value})
|
||||
assert exc.value.status_code == 500
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue