diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index c56fe10daa5..9a8dc5e262d 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -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()