mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
42d2158894
commit
fbbb54c841
7 changed files with 267 additions and 17 deletions
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue