mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
test(proxy): assert websocket auth outcomes instead of delegate arguments
The websocket auth tests now run the real user_api_key_auth against a keyless and a master-key proxy and check what the client gets back, in place of the arguments handed to a mocked delegate
This commit is contained in:
parent
ea0e577385
commit
4147d2a55d
1 changed files with 65 additions and 83 deletions
|
|
@ -4,7 +4,7 @@ import logging
|
|||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
from contextlib import contextmanager
|
||||
from contextlib import AbstractContextManager, contextmanager
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from functools import partial
|
||||
from pathlib import Path
|
||||
|
|
@ -8558,32 +8558,36 @@ async def test_router_settings_model_group_alias_authorizes_target_for_team(monk
|
|||
assert get_client_requested_model(request) == "AgentX-LLM"
|
||||
|
||||
|
||||
def _websocket_for_auth(headers: dict | None = None) -> MagicMock:
|
||||
def _websocket_for_auth(headers: dict[str, str]) -> tuple[MagicMock, AsyncMock]:
|
||||
from fastapi import WebSocket
|
||||
from starlette.datastructures import URL
|
||||
|
||||
close: Final = AsyncMock()
|
||||
websocket: Final = MagicMock(spec=WebSocket)
|
||||
websocket.query_params = {"model": "test-model"}
|
||||
websocket.headers = headers or {}
|
||||
websocket.headers = headers
|
||||
websocket.scope = {
|
||||
"type": "websocket",
|
||||
"path": "/v1/responses",
|
||||
"headers": [(name.lower().encode(), value.encode()) for name, value in (headers or {}).items()],
|
||||
"headers": [(name.lower().encode(), value.encode()) for name, value in headers.items()],
|
||||
}
|
||||
websocket.url = URL(url="/v1/responses")
|
||||
websocket.close = AsyncMock()
|
||||
return websocket
|
||||
websocket.close = close
|
||||
return websocket, close
|
||||
|
||||
|
||||
_KEYLESS_PROXY_STATE: Final = {
|
||||
"prisma_client": None,
|
||||
"user_custom_auth": None,
|
||||
"general_settings": {},
|
||||
"llm_model_list": [],
|
||||
"llm_router": None,
|
||||
"jwt_handler": None,
|
||||
"open_telemetry_logger": None,
|
||||
}
|
||||
def _proxy_state(master_key: str | None) -> AbstractContextManager[object]:
|
||||
return patch.multiple( # test-quality-ok: master_key is a proxy_server module global with no injection seam
|
||||
"litellm.proxy.proxy_server",
|
||||
master_key=master_key,
|
||||
prisma_client=None,
|
||||
user_custom_auth=None,
|
||||
general_settings={},
|
||||
llm_model_list=[],
|
||||
llm_router=None,
|
||||
jwt_handler=None,
|
||||
open_telemetry_logger=None,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
|
|
@ -8594,87 +8598,65 @@ _KEYLESS_PROXY_STATE: Final = {
|
|||
],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_websocket_auth_forwards_a_missing_key_as_none(headers):
|
||||
"""A missing key must reach user_api_key_auth as None, the value
|
||||
APIKeyHeader(auto_error=False) gives the HTTP routes for an absent header.
|
||||
Rejecting it here meant a proxy with no master key refused the WebSocket
|
||||
while accepting every HTTP route."""
|
||||
websocket: Final = _websocket_for_auth(headers)
|
||||
async def test_websocket_auth_without_master_key_accepts_a_keyless_client(headers: dict[str, str]) -> None:
|
||||
websocket, _ = _websocket_for_auth(headers)
|
||||
|
||||
with patch( # test-quality-ok: what is under test is the value handed to the delegate, so it has to be observed
|
||||
"litellm.proxy.auth.user_api_key_auth.user_api_key_auth", autospec=True
|
||||
) as mock_auth:
|
||||
await user_api_key_auth_websocket(websocket)
|
||||
with _proxy_state(master_key=None):
|
||||
result: Final = await user_api_key_auth_websocket(websocket)
|
||||
|
||||
assert mock_auth.call_args.kwargs["api_key"] is None
|
||||
websocket.close.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_websocket_auth_without_master_key_returns_an_internal_user():
|
||||
"""With no master key configured, a keyless connection authenticates."""
|
||||
websocket: Final = _websocket_for_auth({})
|
||||
|
||||
with patch.multiple( # test-quality-ok: master_key is a proxy_server module global with no injection seam
|
||||
"litellm.proxy.proxy_server", master_key=None, **_KEYLESS_PROXY_STATE
|
||||
):
|
||||
result = await user_api_key_auth_websocket(websocket)
|
||||
|
||||
assert isinstance(result, UserAPIKeyAuth)
|
||||
assert result.user_role == LitellmUserRoles.INTERNAL_USER
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_websocket_auth_with_master_key_still_refuses_a_keyless_client():
|
||||
"""Delegating the decision is only correct if the delegate still says no."""
|
||||
from starlette.exceptions import WebSocketException
|
||||
|
||||
websocket: Final = _websocket_for_auth({})
|
||||
|
||||
with patch.multiple( # test-quality-ok: master_key is a proxy_server module global with no injection seam
|
||||
"litellm.proxy.proxy_server", master_key="sk-master-key", **_KEYLESS_PROXY_STATE
|
||||
):
|
||||
with pytest.raises(WebSocketException):
|
||||
await user_api_key_auth_websocket(websocket)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_websocket_auth_rejects_a_malformed_header_without_closing_first():
|
||||
"""Closing the socket and then raising an HTTPException makes Starlette
|
||||
start an HTTP response on a closed socket, which surfaces as a RuntimeError
|
||||
on top of the real auth failure."""
|
||||
from starlette.exceptions import WebSocketException
|
||||
|
||||
websocket: Final = _websocket_for_auth({"authorization": "Token sk-1234"})
|
||||
|
||||
with pytest.raises(WebSocketException) as exc_info:
|
||||
await user_api_key_auth_websocket(websocket)
|
||||
|
||||
assert exc_info.value.code == status.WS_1008_POLICY_VIOLATION
|
||||
websocket.close.assert_not_called()
|
||||
assert result.api_key is None
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"headers,expected",
|
||||
"headers",
|
||||
[
|
||||
pytest.param({"authorization": "Bearer sk-abc"}, "Bearer sk-abc", id="bearer token"),
|
||||
pytest.param({"api-key": "sk-abc"}, "Bearer sk-abc", id="api-key header"),
|
||||
pytest.param({"authorization": "Bearer sk-master-key"}, id="bearer token"),
|
||||
pytest.param({"api-key": "sk-master-key"}, id="api-key header"),
|
||||
pytest.param(
|
||||
{"sec-websocket-protocol": "realtime, openai-insecure-api-key.sk-abc"},
|
||||
"Bearer sk-abc",
|
||||
{"sec-websocket-protocol": "realtime, openai-insecure-api-key.sk-master-key"},
|
||||
id="browser subprotocol",
|
||||
),
|
||||
],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_websocket_auth_still_reads_every_key_source(headers, expected):
|
||||
"""Accept control: forwarding None for every request would satisfy the
|
||||
keyless assertions above and drop real keys on the floor."""
|
||||
websocket: Final = _websocket_for_auth(headers)
|
||||
async def test_websocket_auth_accepts_the_master_key_from_every_key_source(headers: dict[str, str]) -> None:
|
||||
websocket, _ = _websocket_for_auth(headers)
|
||||
|
||||
with patch( # test-quality-ok: what is under test is the value handed to the delegate, so it has to be observed
|
||||
"litellm.proxy.auth.user_api_key_auth.user_api_key_auth", autospec=True
|
||||
) as mock_auth:
|
||||
with _proxy_state(master_key="sk-master-key"):
|
||||
result: Final = await user_api_key_auth_websocket(websocket)
|
||||
|
||||
assert result.user_role == LitellmUserRoles.PROXY_ADMIN
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"headers",
|
||||
[
|
||||
pytest.param({}, id="no key"),
|
||||
pytest.param({"authorization": "Bearer sk-wrong-key"}, id="wrong key"),
|
||||
],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_websocket_auth_with_master_key_refuses_a_client_without_it(headers: dict[str, str]) -> None:
|
||||
from starlette.exceptions import WebSocketException
|
||||
|
||||
websocket, close = _websocket_for_auth(headers)
|
||||
|
||||
with _proxy_state(master_key="sk-master-key"), pytest.raises(WebSocketException) as exc_info:
|
||||
await user_api_key_auth_websocket(websocket)
|
||||
|
||||
assert mock_auth.call_args.kwargs["api_key"] == expected
|
||||
assert exc_info.value.code == status.WS_1008_POLICY_VIOLATION
|
||||
close.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_websocket_auth_rejects_a_malformed_header_without_closing_first() -> None:
|
||||
from starlette.exceptions import WebSocketException
|
||||
|
||||
websocket, close = _websocket_for_auth({"authorization": "Token sk-1234"})
|
||||
|
||||
with _proxy_state(master_key=None), pytest.raises(WebSocketException) as exc_info:
|
||||
await user_api_key_auth_websocket(websocket)
|
||||
|
||||
assert exc_info.value.code == status.WS_1008_POLICY_VIOLATION
|
||||
close.assert_not_called()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue