diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 22826f48b52..478f4a7e07d 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -544,6 +544,36 @@ def _apply_budget_limits_to_end_user_params( verbose_proxy_logger.debug("Applied budget limits to end user %s", end_user_id) +def _get_websocket_api_key(websocket: WebSocket) -> str | None: + """Read the API key a WebSocket client presented, or None when it presented none. + + Whether a key is required is decided by ``user_api_key_auth``, which allows a + keyless request when no master key is configured. + """ + authorization: Final = websocket.headers.get("authorization") + if authorization: + if not authorization.startswith("Bearer "): + raise WebSocketException( + code=status.WS_1008_POLICY_VIOLATION, + reason="Invalid Authorization header format", + ) + return authorization[len("Bearer ") :].strip() + + header_key: Final = websocket.headers.get("api-key") + if header_key: + return header_key + + subprotocol_prefix: Final = "openai-insecure-api-key." + return next( + ( + protocol.strip()[len(subprotocol_prefix) :] + for protocol in websocket.headers.get("sec-websocket-protocol", "").split(",") + if protocol.strip().startswith(subprotocol_prefix) + ), + None, + ) + + async def user_api_key_auth_websocket(websocket: WebSocket): # Accept the WebSocket connection @@ -575,38 +605,18 @@ async def user_api_key_auth_websocket(websocket: WebSocket): request.body = return_body - authorization: Final = websocket.headers.get("authorization") - # If no Authorization header, try the api-key header - if not authorization: - api_key = websocket.headers.get("api-key") - if not api_key: - # Try extracting from WebSocket subprotocol (browser clients) - for protocol in websocket.headers.get("sec-websocket-protocol", "").split(","): - protocol = protocol.strip() - if protocol.startswith("openai-insecure-api-key."): - api_key = protocol[len("openai-insecure-api-key.") :] - break - if not api_key: - await websocket.close(code=status.WS_1008_POLICY_VIOLATION) - raise HTTPException(status_code=403, detail="No API key provided") - else: - # Extract the API key from the Bearer token - if not authorization.startswith("Bearer "): - await websocket.close(code=status.WS_1008_POLICY_VIOLATION) - raise HTTPException(status_code=403, detail="Invalid Authorization header format") + api_key: Final = _get_websocket_api_key(websocket) - api_key = authorization[len("Bearer ") :].strip() - - # Call user_api_key_auth with the extracted API key - # Note: You'll need to modify this to work with WebSocket context if needed try: - return await user_api_key_auth(request=request, api_key=f"Bearer {api_key}") + return await user_api_key_auth( + request=request, + api_key=f"Bearer {api_key}" if api_key else None, # pyright: ignore[reportArgumentType] # None = no key + ) except Exception as e: if is_invalid_virtual_key_error(e): raise WebSocketException(code=status.WS_1008_POLICY_VIOLATION) verbose_proxy_logger.exception(e) - await websocket.close(code=status.WS_1008_POLICY_VIOLATION) - raise HTTPException(status_code=403, detail=str(e)) + raise WebSocketException(code=status.WS_1008_POLICY_VIOLATION, reason=str(e)) def update_valid_token_with_end_user_params(valid_token: UserAPIKeyAuth, end_user_params: dict) -> UserAPIKeyAuth: 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 c9ae105d982..b3a6f1286b9 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 @@ -9,6 +9,7 @@ from datetime import datetime, timedelta, timezone from pathlib import Path from textwrap import dedent from types import SimpleNamespace +from typing import Final from unittest.mock import ANY, AsyncMock, MagicMock, patch @@ -48,6 +49,7 @@ from litellm.proxy.auth.user_api_key_auth import ( _user_api_key_auth_builder, get_api_key, user_api_key_auth, + user_api_key_auth_websocket, ) from litellm.proxy.spend_tracking.carried_budget_state import carried_budget_metadata @@ -7851,3 +7853,117 @@ async def test_auth_flow_enters_virtual_key_mapping_when_only_an_issuer_configur assert resolve_mock.await_args.kwargs["jwt_claims"][JWTHandler.LITELLM_JWT_ISSUER_CLAIM] == ISSUER_TWO assert result.api_key == "hashed-mapped-key" assert result.team_id == "svc-team" + + +def _websocket_for_auth(headers: dict | None = None) -> MagicMock: + from fastapi import WebSocket + from starlette.datastructures import URL + + websocket: Final = MagicMock(spec=WebSocket) + websocket.query_params = {"model": "test-model"} + websocket.headers = headers or {} + websocket.scope = { + "type": "websocket", + "path": "/v1/responses", + "headers": [(name.lower().encode(), value.encode()) for name, value in (headers or {}).items()], + } + websocket.url = URL(url="/v1/responses") + websocket.close = AsyncMock() + return websocket + + +_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, +} + + +@pytest.mark.parametrize( + "headers", + [ + pytest.param({}, id="no headers at all"), + pytest.param({"sec-websocket-protocol": "realtime"}, id="subprotocol carrying no key"), + ], +) +@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) + + with patch("litellm.proxy.auth.user_api_key_auth.user_api_key_auth", autospec=True) as mock_auth: + 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("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("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() + + +@pytest.mark.parametrize( + "headers,expected", + [ + 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( + {"sec-websocket-protocol": "realtime, openai-insecure-api-key.sk-abc"}, + "Bearer sk-abc", + 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) + + with patch("litellm.proxy.auth.user_api_key_auth.user_api_key_auth", autospec=True) as mock_auth: + await user_api_key_auth_websocket(websocket) + + assert mock_auth.call_args.kwargs["api_key"] == expected