fix(proxy): answer a rejected websocket handshake without crashing the HTTP exception handler (#45251)

* fix(proxy): answer a rejected websocket handshake without crashing the HTTP exception handler

* fix(proxy): read the OTLP route check's scope keys with .get and drive the websocket auth tests through the real auth path

* test(proxy): build the websocket auth test keys without reading the clock
This commit is contained in:
Mateo Wang 2026-10-07 22:27:46 -07:00 • committed by GitHub
parent 42d2158894
commit fbbb54c841
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 267 additions and 17 deletions

View file

@ -60,6 +60,12 @@ def is_invalid_virtual_key_error(exception: BaseException | None) -> bool:
return getattr(exception, INVALID_VIRTUAL_KEY_ERROR_MARKER, False) is True
def log_model_access_denial(exc: BaseException) -> None:
if not isinstance(exc, ModelAccessDeniedProxyException):
return
verbose_proxy_logger.warning(exc.sanitized_internal_message())
def mark_invalid_virtual_key_error(exception: ProxyException, is_invalid_virtual_key: bool) -> ProxyException:
"""Return an independently marked malformed-key exception after callback transformations."""
if not is_invalid_virtual_key or str(exception.code) != str(status.HTTP_401_UNAUTHORIZED):

View file

@ -100,6 +100,7 @@ from litellm.proxy.auth.auth_utils import (
get_request_route_template,
is_invalid_virtual_key_error,
iter_request_fallback_targets,
log_model_access_denial,
normalize_request_route,
pre_db_read_auth_checks,
request_dispatched_to_pass_through_endpoint,
@ -732,7 +733,7 @@ async def user_api_key_auth_websocket_for_model(websocket: WebSocket, model: str
except Exception as e:
if is_invalid_virtual_key_error(e):
raise WebSocketException(code=status.WS_1008_POLICY_VIOLATION)
verbose_proxy_logger.exception(e)
log_model_access_denial(e)
await websocket.close(code=status.WS_1008_POLICY_VIOLATION)
raise HTTPException(status_code=403, detail=str(e))

View file

@ -7,6 +7,7 @@ from typing import Annotated, Any, Final, Literal, Union, get_args, get_origin
import orjson
from fastapi import Request, UploadFile, status
from starlette._utils import get_route_path
from starlette.requests import HTTPConnection
from typing_extensions import NotRequired, ReadOnly, Required, assert_never
from litellm._logging import verbose_proxy_logger
@ -182,8 +183,10 @@ def _mark_body_received(byte_count: int | None) -> None:
)
def is_otlp_trace_request(request: Request) -> bool:
return request.method == "POST" and get_route_path(request.scope) in {"/v1/traces", "/v1/logs"}
def is_otlp_trace_request(request: HTTPConnection) -> bool:
if request.scope.get("type") != "http":
return False
return request.scope.get("method") == "POST" and get_route_path(request.scope) in {"/v1/traces", "/v1/logs"}
async def read_request_body(request: Request | None) -> dict:

View file

@ -114,7 +114,6 @@ from litellm.proxy._types import (
LiteLLM_TeamTableCachedObj,
LiteLLM_UserTable,
LitellmUserRoles,
ModelAccessDeniedProxyException,
PassThroughGenericEndpoint,
ProxyErrorTypes,
ProxyException,
@ -356,6 +355,7 @@ from litellm.proxy.auth.auth_object_prefetch import AUTH_OBJECTS_TARGET
from litellm.proxy.auth.auth_utils import (
check_response_size_is_safe,
is_request_body_safe,
log_model_access_denial,
log_once_if_budget_reservation_disabled,
warn_once_if_custom_auth_skips_common_checks,
)
@ -768,6 +768,7 @@ except ImportError:
shutdown_billing_metrics_recorder = None
from fastapi.exception_handlers import http_exception_handler
from starlette.exceptions import HTTPException as StarletteHTTPException
from starlette.websockets import WebSocketState
from litellm.proxy import tracing_endpoints
from litellm.proxy.middleware.admission_control_middleware import (
@ -2057,7 +2058,7 @@ class UserAPIKeyCacheTTLEnum(enum.Enum):
@app.exception_handler(ProxyException)
async def openai_exception_handler(request: Request, exc: ProxyException):
# NOTE: DO NOT MODIFY THIS, its crucial to map to Openai exceptions
_log_model_access_denial(exc)
log_model_access_denial(exc)
headers: Final = exc.headers
error_dict: Final = with_call_id(
JSON_OBJECT.validate_python(exc.to_dict()),
@ -2076,18 +2077,32 @@ async def openai_exception_handler(request: Request, exc: ProxyException):
@app.exception_handler(StarletteHTTPException)
async def otlp_http_exception_handler(request: Request, exc: StarletteHTTPException) -> Response:
response: Final = tracing_endpoints.otlp_error_response(request, exc.status_code, exc.headers)
async def otlp_http_exception_handler(connection: Request | WebSocket, exc: StarletteHTTPException) -> Response | None:
if isinstance(connection, WebSocket):
return await _websocket_http_exception_response(connection, exc)
response: Final = tracing_endpoints.otlp_error_response(connection, exc.status_code, exc.headers)
if response is not None:
_close_dangling_otel_server_span(request, exc.status_code, exc=exc)
_close_dangling_otel_server_span(connection, exc.status_code, exc=exc)
return response
return await http_exception_handler(request, exc)
return await http_exception_handler(connection, exc)
def _log_model_access_denial(exc: ProxyException) -> None:
if not isinstance(exc, ModelAccessDeniedProxyException):
return
verbose_proxy_logger.warning(exc.sanitized_internal_message())
async def _websocket_http_exception_response(websocket: WebSocket, exc: StarletteHTTPException) -> Response | None:
state: Final = websocket.application_state
match state:
case WebSocketState.CONNECTING:
return JSONResponse({"detail": exc.detail}, status_code=exc.status_code, headers=exc.headers)
case WebSocketState.CONNECTED:
await websocket.close(code=_websocket_close_code(exc.status_code))
return None
case WebSocketState.RESPONSE | WebSocketState.DISCONNECTED:
return None
case _:
assert_never(state)
def _websocket_close_code(status_code: int) -> int:
return status.WS_1011_INTERNAL_ERROR if status_code >= 500 else status.WS_1008_POLICY_VIOLATION
def _close_dangling_otel_server_span(request: Request, status_code: int, exc: Exception | None = None) -> None:
@ -13247,7 +13262,7 @@ async def realtime_websocket_endpoint(
llm_router=llm_router,
)
except ProxyException as e:
_log_model_access_denial(e)
log_model_access_denial(e)
await _reject_realtime_session(websocket, user_api_key_dict, code=1008, reason=e.message[:120])
return
await websocket.accept(**accept_kwargs)

View file

@ -5,11 +5,12 @@
import litellm.proxy
import litellm.proxy.proxy_server
from typing import Dict, List, Optional
from unittest.mock import MagicMock, patch, AsyncMock
from typing import Dict, Final, List, Optional
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from starlette.datastructures import URL
from starlette.types import Message
from litellm._logging import verbose_proxy_logger
import logging
import litellm
@ -1811,3 +1812,91 @@ def test_mapped_key_jwt_falls_through_to_the_shared_user_budget_attach():
"the mapped-key branch returns before the shared virtual-key checks, so the "
"user's per-model budget is never attached and never enforced"
)
def _rejected_websocket(sent: list[Message], bearer: str) -> WebSocket:
async def receive() -> Message:
return {"type": "websocket.connect"}
async def send(message: Message) -> None:
sent.append(message)
return WebSocket(
{
"type": "websocket",
"path": "/v1/responses",
"query_string": b"model=gpt-5.4",
"headers": [(b"authorization", f"Bearer {bearer}".encode())],
},
receive,
send,
)
def _serve_virtual_key(monkeypatch: pytest.MonkeyPatch, user_key: str, token: UserAPIKeyAuth) -> None:
from litellm.proxy.proxy_server import hash_token, user_api_key_cache
user_api_key_cache.set_cache(key=hash_token(user_key), value=token)
monkeypatch.setattr(litellm.proxy.proxy_server, "user_api_key_cache", user_api_key_cache)
monkeypatch.setattr(litellm.proxy.proxy_server, "master_key", MASTER_KEY)
monkeypatch.setattr(litellm.proxy.proxy_server, "prisma_client", "connected")
monkeypatch.setattr(litellm, "log_client_error_tracebacks", False)
monkeypatch.setattr(verbose_proxy_logger, "propagate", True)
@pytest.mark.asyncio
async def test_user_api_key_auth_websocket_logs_a_model_access_denial_as_one_warning_without_a_traceback(
caplog: pytest.LogCaptureFixture, monkeypatch: pytest.MonkeyPatch
) -> None:
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_websocket
from litellm.proxy.proxy_server import hash_token
user_key: Final = "sk-websocket-key-limited-to-mini"
_serve_virtual_key(
monkeypatch,
user_key,
UserAPIKeyAuth(token=hash_token(user_key), models=["gpt-5.4-mini"]),
)
sent: Final[list[Message]] = []
with (
caplog.at_level(logging.WARNING, logger=verbose_proxy_logger.name),
pytest.raises(HTTPException) as rejection,
):
await user_api_key_auth_websocket(_rejected_websocket(sent, user_key))
assert rejection.value.status_code == 403
assert sent == [{"type": "websocket.close", "code": status.WS_1008_POLICY_VIOLATION, "reason": ""}]
proxy_records: Final = [record for record in caplog.records if record.name == verbose_proxy_logger.name]
assert [record.getMessage() for record in proxy_records if record.exc_info is not None] == []
assert [record.getMessage() for record in proxy_records if record.levelno == logging.WARNING] == [
"key not allowed to access model. This key can only access models=['gpt-5.4-mini']. Tried to access gpt-5.4"
]
@pytest.mark.asyncio
async def test_user_api_key_auth_websocket_rejection_adds_no_traceback_for_other_auth_errors(
caplog: pytest.LogCaptureFixture, monkeypatch: pytest.MonkeyPatch
) -> None:
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_websocket
from litellm.proxy.proxy_server import hash_token
user_key: Final = "sk-websocket-key-expired-in-2020"
expired_in_2020: Final = "2020-01-01T00:00:00+00:00"
_serve_virtual_key(
monkeypatch,
user_key,
UserAPIKeyAuth(token=hash_token(user_key), expires=expired_in_2020),
)
sent: Final[list[Message]] = []
with (
caplog.at_level(logging.DEBUG, logger=verbose_proxy_logger.name),
pytest.raises(HTTPException) as rejection,
):
await user_api_key_auth_websocket(_rejected_websocket(sent, user_key))
assert rejection.value.status_code == 403
assert "expired key" in str(rejection.value.detail).lower()
assert sent == [{"type": "websocket.close", "code": status.WS_1008_POLICY_VIOLATION, "reason": ""}]
proxy_records: Final = [record for record in caplog.records if record.name == verbose_proxy_logger.name]
assert [record.getMessage() for record in proxy_records if record.exc_info is not None] == []
assert [record.levelno for record in proxy_records if record.levelno >= logging.WARNING] == [logging.ERROR]

View file

@ -34,6 +34,7 @@ from fastapi import Request as Request_http_parsing
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
from litellm.utils import _invalidate_model_cost_lowercase_map
from starlette.types import Message
from starlette.websockets import WebSocket
from tests._vcr_conftest_common import install_live_call_probe, record_vcr_outcome
@ -1565,3 +1566,36 @@ async def test_read_request_body_unexpected_error():
result = await read_request_body(_request(receive))
assert result == {}
async def _never_receive() -> Message:
raise AssertionError("the OTLP route check never reads the body")
async def _never_send(message: Message) -> None:
raise AssertionError("the OTLP route check never sends")
@pytest.mark.parametrize(
("scope", "expected"),
[
({"type": "http", "method": "POST", "path": "/v1/traces"}, True),
({"type": "http", "method": "POST", "path": "/v1/logs"}, True),
({"type": "http", "method": "GET", "path": "/v1/traces"}, False),
({"type": "http", "method": "POST", "path": "/v1/responses"}, False),
({"type": "http", "path": "/v1/traces"}, False),
({"type": "websocket", "path": "/v1/traces"}, False),
({"type": "websocket", "path": "/v1/responses"}, False),
],
)
def test_is_otlp_trace_request_matches_only_http_posts_to_the_otlp_routes(
scope: dict[str, object], expected: bool
) -> None:
full_scope: Final[dict[str, object]] = {**scope, "headers": [], "query_string": b""}
connection: Final = (
WebSocket(full_scope, _never_receive, _never_send)
if scope["type"] == "websocket"
else Request(full_scope, _never_receive)
)
assert http_parsing_utils.is_otlp_trace_request(connection) is expected

View file

@ -5,6 +5,7 @@ Pins covered:
- ``_close_dangling_otel_server_span``
- ``otel_request_validation_exception_handler``
- ``otel_unhandled_exception_handler``
- ``otlp_http_exception_handler``
"""
from __future__ import annotations
@ -16,8 +17,12 @@ from unittest.mock import MagicMock
import httpx
import pytest
from fastapi import HTTPException, Request
from fastapi import Depends, FastAPI, HTTPException, Request
from fastapi.exceptions import RequestValidationError
from fastapi.testclient import TestClient
from starlette.testclient import WebSocketDenialResponse
from starlette.types import Message
from starlette.websockets import WebSocket, WebSocketDisconnect
from litellm.proxy._types import ProxyException
from litellm.proxy.proxy_server import (
@ -25,6 +30,7 @@ from litellm.proxy.proxy_server import (
openai_exception_handler,
otel_request_validation_exception_handler,
otel_unhandled_exception_handler,
otlp_http_exception_handler,
)
from .conftest import normalize
@ -507,6 +513,7 @@ async def test_otlp_auth_errors_hide_internal_details_and_survive_missing_native
if isinstance(error, ProxyException)
else await otlp_http_exception_handler(request, error)
)
assert response is not None
assert response.status_code == (401 if isinstance(error, ProxyException) else 403)
assert response.headers["content-type"].startswith(media_type)
message: Final = (
@ -516,3 +523,98 @@ async def test_otlp_auth_errors_hide_internal_details_and_survive_missing_native
)
expected: Final = "Unauthorized" if isinstance(error, ProxyException) else "Forbidden"
assert message == (expected if native_available or media_type == "application/json" else "")
def _websocket(sent: list[Message]) -> WebSocket:
async def receive() -> Message:
return {"type": "websocket.connect"}
async def send(message: Message) -> None:
sent.append(message)
return WebSocket({"type": "websocket", "path": "/v1/responses", "headers": [], "query_string": b""}, receive, send)
@pytest.mark.asyncio
async def test_http_exception_on_a_plain_http_request_keeps_the_default_json_body() -> None:
request: Final = _make_request(path="/v1/models")
response: Final = await otlp_http_exception_handler(request, HTTPException(404, "not found"))
assert response is not None
assert response.status_code == 404
assert json.loads(bytes(response.body)) == {"detail": "not found"}
@pytest.mark.asyncio
async def test_http_exception_on_a_websocket_closed_before_accept_sends_nothing_more() -> None:
sent: Final[list[Message]] = []
websocket: Final = _websocket(sent)
await websocket.close(code=1008)
response: Final = await otlp_http_exception_handler(websocket, HTTPException(403, "No API key provided"))
assert response is None
assert sent == [{"type": "websocket.close", "code": 1008, "reason": ""}]
@pytest.mark.asyncio
async def test_http_exception_on_a_connecting_websocket_denies_the_upgrade_with_its_status() -> None:
websocket: Final = _websocket([])
response: Final = await otlp_http_exception_handler(websocket, HTTPException(403, "No API key provided"))
assert response is not None
assert response.status_code == 403
assert json.loads(bytes(response.body)) == {"detail": "No API key provided"}
@pytest.mark.asyncio
@pytest.mark.parametrize(("status_code", "close_code"), [(403, 1008), (429, 1008), (500, 1011), (503, 1011)])
async def test_http_exception_on_an_accepted_websocket_closes_it_with_a_matching_code(
status_code: int, close_code: int
) -> None:
sent: Final[list[Message]] = []
websocket: Final = _websocket(sent)
await websocket.accept()
response: Final = await otlp_http_exception_handler(websocket, HTTPException(status_code, "late failure"))
assert response is None
assert sent[-1] == {"type": "websocket.close", "code": close_code, "reason": ""}
def test_websocket_auth_rejection_reaches_the_client_as_a_denial_instead_of_a_server_error() -> None:
async def reject_like_user_api_key_auth_websocket(websocket: WebSocket) -> None:
await websocket.close(code=1008)
raise HTTPException(status_code=403, detail="No API key provided")
async def responses(websocket: WebSocket, _: None = Depends(reject_like_user_api_key_auth_websocket)) -> None:
await websocket.accept()
app: Final = FastAPI()
app.exception_handler(HTTPException)(otlp_http_exception_handler)
app.add_api_websocket_route("/v1/responses", responses)
with pytest.raises(WebSocketDisconnect) as disconnect, TestClient(app).websocket_connect("/v1/responses"):
pass
assert disconnect.value.code == 1008
def test_websocket_rejected_before_any_close_reaches_the_client_as_an_http_denial_with_the_status() -> None:
async def reject_without_closing(websocket: WebSocket) -> None:
raise HTTPException(status_code=403, detail="No API key provided")
async def responses(websocket: WebSocket, _: None = Depends(reject_without_closing)) -> None:
await websocket.accept()
app: Final = FastAPI()
app.exception_handler(HTTPException)(otlp_http_exception_handler)
app.add_api_websocket_route("/v1/responses", responses)
with pytest.raises(WebSocketDenialResponse) as denial, TestClient(app).websocket_connect("/v1/responses"):
pass
assert denial.value.status_code == 403
assert denial.value.json() == {"detail": "No API key provided"}