This commit is contained in:
Lukas 2026-09-27 18:57:40 +08:00 • committed by GitHub
commit 5c8da08c5f
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 114 additions and 14 deletions

View file

@ -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:

View file

@ -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()