mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge 0716a14a8b into 6532dcb73b
This commit is contained in:
commit
cb21b7f2a3
2 changed files with 89 additions and 1 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue