diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index e3ce9bcd850..d9d4912ad6c 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -669,37 +669,30 @@ async def user_api_key_auth_websocket_for_model(websocket: WebSocket, model: str 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") + raise WebSocketException(code=status.WS_1008_POLICY_VIOLATION, reason="Invalid Authorization header format") 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 470db99108a..5d0b2a0bb58 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 @@ -6,12 +6,13 @@ import subprocess import sys import time from collections.abc import Mapping -from contextlib import contextmanager +from contextlib import AbstractContextManager, contextmanager from datetime import datetime, timedelta, timezone from functools import partial from pathlib import Path from textwrap import dedent from types import SimpleNamespace +from typing import Final from unittest.mock import ANY, AsyncMock, MagicMock, patch @@ -61,6 +62,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, user_api_key_auth_websocket_for_model, ) from litellm.proxy.spend_tracking.carried_budget_state import carried_budget_metadata @@ -9290,3 +9292,108 @@ async def test_websocket_auth_hands_the_reservation_to_the_socket_state(): assert result.budget_reservation == reservation assert websocket.state.budget_reservation is reservation assert websocket.scope["state"]["budget_reservation"] is reservation + + +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 + websocket.scope = { + "type": "websocket", + "path": "/v1/responses", + "headers": [(name.lower().encode(), value.encode()) for name, value in headers.items()], + } + websocket.url = URL(url="/v1/responses") + websocket.close = close + return websocket, close + + +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( + "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_without_master_key_accepts_a_keyless_client(headers: dict[str, str]) -> None: + websocket, _ = _websocket_for_auth(headers) + + with _proxy_state(master_key=None): + result: Final = await user_api_key_auth_websocket(websocket) + + assert result.user_role == LitellmUserRoles.INTERNAL_USER + assert result.api_key is None + + +@pytest.mark.parametrize( + "headers", + [ + pytest.param({"authorization": "Bearer sk-master-key"}, id="bearer token"), + pytest.param({"authorization": "Bearer sk-master-key"}, id="bearer token with extra space"), + pytest.param({"api-key": "sk-master-key"}, id="api-key header"), + pytest.param( + {"sec-websocket-protocol": "realtime, openai-insecure-api-key.sk-master-key"}, + id="browser subprotocol", + ), + ], +) +@pytest.mark.asyncio +async def test_websocket_auth_accepts_the_master_key_from_every_key_source(headers: dict[str, str]) -> None: + websocket, _ = _websocket_for_auth(headers) + + 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 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()