mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
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:
parent
7fd541efb9
commit
f0e9df2934
2 changed files with 152 additions and 26 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue