From 0716a14a8b77ef3f119bb66d1fed0a6ecf69c5b8 Mon Sep 17 00:00:00 2001 From: yassin Date: Sun, 4 Oct 2026 00:34:04 +0000 Subject: [PATCH] fix(proxy): keep the parsed key when the custom key header is absent Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/auth/user_api_key_auth.py | 5 +- .../unit/proxy/auth/test_user_api_key_auth.py | 85 +++++++++++++++++++ 2 files changed, 89 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index cb39801a8e0..c03ea1c59c0 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -1573,7 +1573,10 @@ async def _user_api_key_auth_builder( ) # if user wants to pass LiteLLM_Master_Key as a custom header, example pass litellm keys as X-LiteLLM-Key: Bearer sk-1234 custom_litellm_key_header_name: Final = general_settings.get("litellm_key_header_name") - if custom_litellm_key_header_name is not None: + if ( + custom_litellm_key_header_name is not None + and request.headers.get(custom_litellm_key_header_name) is not None + ): api_key = get_api_key_from_custom_header( request=request, custom_litellm_key_header_name=custom_litellm_key_header_name, diff --git a/tests/unit/proxy/auth/test_user_api_key_auth.py b/tests/unit/proxy/auth/test_user_api_key_auth.py index 1cfef5b3a6c..cc99b5e63fd 100644 --- a/tests/unit/proxy/auth/test_user_api_key_auth.py +++ b/tests/unit/proxy/auth/test_user_api_key_auth.py @@ -1174,6 +1174,91 @@ async def test_x_litellm_api_key(): assert valid_token.token != hash_token(master_key) +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("headers", "api_key", "custom_litellm_key_header"), + [ + ({"Authorization": "Bearer sk-1234"}, "Bearer sk-1234", None), + ({"x-litellm-api-key": "Bearer sk-1234"}, "", "Bearer sk-1234"), + ( + {"X-Custom-Key": "Bearer sk-1234", "Authorization": "Bearer sk-wrong"}, + "Bearer sk-wrong", + None, + ), + ], +) +async def test_auth_custom_key_header_precedence_and_fallback( + headers: dict[str, str], + api_key: str, + custom_litellm_key_header: str | None, + monkeypatch: pytest.MonkeyPatch, +): + from fastapi import Request + from starlette.datastructures import URL + + from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS + from litellm.proxy import proxy_server + + monkeypatch.setattr(proxy_server, "master_key", "sk-1234") + monkeypatch.setattr(proxy_server, "general_settings", {"litellm_key_header_name": "X-Custom-Key"}) + + request_headers = [(name.lower().encode(), value.encode()) for name, value in headers.items()] + request = Request( + scope={ + "type": "http", + "method": "POST", + "path": "/chat/completions", + "headers": request_headers, + }, + receive=AsyncMock(return_value={"type": "http.request", "body": b"", "more_body": False}), + ) + request._url = URL(url="/chat/completions") + + valid_token = await user_api_key_auth( + request=request, + api_key=api_key, + custom_litellm_key_header=custom_litellm_key_header, + ) + + assert valid_token.user_role == LitellmUserRoles.PROXY_ADMIN + assert valid_token.api_key == LITELLM_PROXY_MASTER_KEY_ALIAS + + +@pytest.mark.asyncio +@pytest.mark.parametrize("custom_header_value", [b"Bearer sk-wrong", b""]) +async def test_auth_rejects_wrong_or_empty_present_custom_header( + custom_header_value: bytes, monkeypatch: pytest.MonkeyPatch +): + from fastapi import Request + from starlette.datastructures import URL + + from litellm.proxy import proxy_server + from litellm.proxy._types import ProxyException + + monkeypatch.setattr(proxy_server, "master_key", "sk-1234") + monkeypatch.setattr(proxy_server, "general_settings", {"litellm_key_header_name": "X-Custom-Key"}) + monkeypatch.setattr(proxy_server, "prisma_client", MagicMock(get_data=AsyncMock(return_value=None))) + + request = Request( + scope={ + "type": "http", + "method": "POST", + "path": "/chat/completions", + "headers": [ + (b"x-custom-key", custom_header_value), + (b"authorization", b"Bearer sk-1234"), + ], + }, + receive=AsyncMock(return_value={"type": "http.request", "body": b"", "more_body": False}), + ) + request._url = URL(url="/chat/completions") + + with pytest.raises(ProxyException) as exc_info: + await user_api_key_auth(request=request, api_key="Bearer sk-1234") + + assert exc_info.value.code == "401" + + @pytest.mark.asyncio async def test_user_api_key_from_query_param(): """Ensure user_api_key_auth reads API key from `key` query parameter."""