From fbbb54c841efbf93f106e1a00f5711447cfa311c Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 7 Oct 2026 22:27:46 -0700 Subject: [PATCH] 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 --- litellm/proxy/auth/auth_utils.py | 6 + litellm/proxy/auth/user_api_key_auth.py | 3 +- .../proxy/common_utils/http_parsing_utils.py | 7 +- litellm/proxy/proxy_server.py | 37 +++++-- .../unit/proxy/auth/test_user_api_key_auth.py | 93 +++++++++++++++- .../common_utils/test_http_parsing_utils.py | 34 ++++++ .../proxy_server/test_exception_handlers.py | 104 +++++++++++++++++- 7 files changed, 267 insertions(+), 17 deletions(-) diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index edd6028f07c..2bd10b90101 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -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): diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index afe88fe09a3..66bc5ae44d6 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -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)) diff --git a/litellm/proxy/common_utils/http_parsing_utils.py b/litellm/proxy/common_utils/http_parsing_utils.py index 44360cec2bd..80d80224617 100644 --- a/litellm/proxy/common_utils/http_parsing_utils.py +++ b/litellm/proxy/common_utils/http_parsing_utils.py @@ -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: diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index aa086e25411..b4b23f111d1 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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) diff --git a/tests/unit/proxy/auth/test_user_api_key_auth.py b/tests/unit/proxy/auth/test_user_api_key_auth.py index a58599d28e9..9cc8552efd5 100644 --- a/tests/unit/proxy/auth/test_user_api_key_auth.py +++ b/tests/unit/proxy/auth/test_user_api_key_auth.py @@ -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] diff --git a/tests/unit/proxy/common_utils/test_http_parsing_utils.py b/tests/unit/proxy/common_utils/test_http_parsing_utils.py index 7677837aa0f..f6f148f02a8 100644 --- a/tests/unit/proxy/common_utils/test_http_parsing_utils.py +++ b/tests/unit/proxy/common_utils/test_http_parsing_utils.py @@ -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 diff --git a/tests/unit/proxy/proxy_server/test_exception_handlers.py b/tests/unit/proxy/proxy_server/test_exception_handlers.py index 0aff43057f9..f6efdecc884 100644 --- a/tests/unit/proxy/proxy_server/test_exception_handlers.py +++ b/tests/unit/proxy/proxy_server/test_exception_handlers.py @@ -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"}