fix(proxy): let the Responses WebSocket follow the proxy keyless policy

user_api_key_auth_websocket rejected a missing key before delegating, so a
proxy running without general_settings.master_key accepted every HTTP route
and refused the WebSocket with 403 No API key provided. Reading the key is
now separate from deciding whether one is required: user_api_key_auth makes
that call, allowing a keyless request when no master key is set and raising
No api key passed in. when one is.

Rejections raise WebSocketException alone. Closing the socket and then
raising an HTTPException asked Starlette to start an HTTP response on a
closed socket, which surfaced as RuntimeError: Unexpected ASGI message
'websocket.http.response.start' on top of the real auth failure.
This commit is contained in:
L4XB 2026-09-15 00:29:03 +02:00
parent 7fd541efb9
commit f0e9df2934
No known key found for this signature in database
2 changed files with 152 additions and 26 deletions

View file

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

View file

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