diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index f6f4d7bf1c4..fec5a1268d2 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -487,6 +487,38 @@ 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: + from litellm.proxy.proxy_server import general_settings + + custom_header: Final = general_settings.get("litellm_key_header_name") + if isinstance(custom_header, str): + if not websocket.headers.get(custom_header): + return None + request: Final = Request( + {"type": "http", "headers": websocket.scope.get("headers", [])} # mutable-ok: ASGI request scope + ) + return get_api_key_from_custom_header(request, custom_header) + custom_key: Final = websocket.headers.get("x-litellm-api-key") + if custom_key is not None: + return _get_bearer_token_or_received_api_key(custom_key) + authorization: Final = websocket.headers.get("authorization") + if authorization: + if not authorization.startswith("Bearer "): + raise HTTPException(status_code=403, detail="Invalid Authorization header format") + return authorization[len("Bearer ") :].strip() + api_key: Final = websocket.headers.get("api-key") + if api_key: + return api_key + return next( + ( + protocol.strip().removeprefix("openai-insecure-api-key.") + for protocol in websocket.headers.get("sec-websocket-protocol", "").split(",") + if protocol.strip().startswith("openai-insecure-api-key.") + ), + None, + ) + + async def user_api_key_auth_websocket(websocket: WebSocket): # Accept the WebSocket connection @@ -509,27 +541,14 @@ async def user_api_key_auth_websocket(websocket: WebSocket): request._url = websocket.url - 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 = authorization[len("Bearer ") :].strip() + try: + api_key: Final = get_websocket_api_key(websocket) + except HTTPException: + await websocket.close(code=status.WS_1008_POLICY_VIOLATION) + raise + if not api_key: + await websocket.close(code=status.WS_1008_POLICY_VIOLATION) + raise HTTPException(status_code=403, detail="No API key provided") # Call user_api_key_auth with the extracted API key # Note: You'll need to modify this to work with WebSocket context if needed diff --git a/litellm/proxy/realtime_endpoints/call_sessions.py b/litellm/proxy/realtime_endpoints/call_sessions.py index 308cf389936..e02a66193b8 100644 --- a/litellm/proxy/realtime_endpoints/call_sessions.py +++ b/litellm/proxy/realtime_endpoints/call_sessions.py @@ -23,7 +23,12 @@ from litellm.llms.chatgpt.codex import ( from litellm.llms.chatgpt.realtime import configured_realtime_headers from litellm.proxy._types import ProxyException, UserAPIKeyAuth from litellm.proxy.auth.auth_checks import can_key_call_resolved_model -from litellm.proxy.auth.user_api_key_auth import get_api_key, get_api_key_from_custom_header, user_api_key_auth +from litellm.proxy.auth.user_api_key_auth import ( + get_api_key, + get_api_key_from_custom_header, + get_websocket_api_key, + user_api_key_auth, +) from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper, encrypt_value_helper from litellm.proxy.spend_tracking.budget_reservation import release_or_invalidate_budget_reservation @@ -171,14 +176,13 @@ async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAP protocols: Final = tuple( p.strip() for p in websocket.headers.get("sec-websocket-protocol", "").split(",") if p.strip() ) - alternate_key: Final = websocket.headers.get("api-key") or next( - (p.removeprefix("openai-insecure-api-key.") for p in protocols if p.startswith("openai-insecure-api-key.")), "" - ) - authorization: Final = websocket.headers.get("authorization") or f"Bearer {alternate_key}" logging_obj: Logging | None = None # rebind-ok: cleanup needs the logger only after pre-call succeeds try: try: - call: Final = decode_call(token, authorization) + api_key: Final = get_websocket_api_key(websocket) + if not api_key: + raise HTTPException(403, "No API key provided") + call: Final = decode_call(token, f"Bearer {api_key}") await can_key_call_resolved_model( model=call.alias, llm_model_list=server.llm_model_list, diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index 886d3361c49..f5d802d843c 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -7137,7 +7137,7 @@ def test_user_api_key_auth_opens_a_datadog_span_for_accepted_and_rejected_keys(t @pytest.mark.asyncio @pytest.mark.parametrize("attachment", ["path", "query"]) -@pytest.mark.parametrize("credential", ["authorization", "api-key", "subprotocol"]) +@pytest.mark.parametrize("credential", ["authorization", "api-key", "subprotocol", "x-litellm-api-key", "custom", "custom-mixed"]) @pytest.mark.parametrize("query_model", [b"", b"model=unbudgeted"]) async def test_sideband_auth_uses_encrypted_model_for_budget_checks(monkeypatch, attachment, credential, query_model): import hashlib @@ -7154,6 +7154,8 @@ async def test_sideband_auth_uses_encrypted_model_for_budget_checks(monkeypatch, call_id="rtc_test", model="gpt-live-1-codex", alias="budgeted-voice", owner=hashlib.sha256(b"Bearer owner").hexdigest(), expires_at=time.time() + 300, )) + from litellm.proxy import proxy_server + monkeypatch.setattr(proxy_server, "general_settings", {"litellm_key_header_name": "x-proxy-key"} if credential.startswith("custom") else {}) seen = [] async def authenticate(request, api_key): @@ -7169,6 +7171,9 @@ async def test_sideband_auth_uses_encrypted_model_for_budget_checks(monkeypatch, "headers": { "authorization": [(b"authorization", b"Bearer owner")], "api-key": [(b"api-key", b"owner")], + "x-litellm-api-key": [(b"x-litellm-api-key", b"owner")], + "custom": [(b"x-proxy-key", b"Bearer owner")], + "custom-mixed": [(b"x-proxy-key", b"Bearer owner"), (b"authorization", b"Bearer other-owner")], "subprotocol": [(b"sec-websocket-protocol", b"realtime, openai-insecure-api-key.owner")], }[credential], }, AsyncMock(), AsyncMock()) @@ -7336,3 +7341,38 @@ async def test_sideband_rejects_budget_fallback_before_rerouting(monkeypatch, at assert error.value.status_code == 403 limiter.get_fallback_model_within_budget.assert_not_awaited() send.assert_awaited_once_with({"type": "websocket.close", "code": 1008, "reason": ""}) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("custom_value", [None, b"Bearer different-owner"]) +async def test_sideband_custom_header_cannot_fall_back_to_other_credentials(monkeypatch, custom_value): + import hashlib + import importlib + import time + from unittest.mock import AsyncMock + from fastapi import HTTPException, WebSocket + from litellm.proxy import proxy_server + from litellm.llms.chatgpt.codex import CodexRealtimeCall + from litellm.proxy.realtime_endpoints.call_sessions import encode_call + + auth_module = importlib.import_module("litellm.proxy.auth.user_api_key_auth") + monkeypatch.setenv("LITELLM_SALT_KEY", "test-only-custom-header-salt") + monkeypatch.setattr(proxy_server, "general_settings", {"litellm_key_header_name": "x-proxy-key"}) + token = encode_call(CodexRealtimeCall( + call_id="rtc_test", model="gpt-live-1-codex", alias="voice", + owner=hashlib.sha256(b"Bearer owner").hexdigest(), expires_at=time.time() + 300, + )) + authenticate = AsyncMock() + monkeypatch.setattr(auth_module, "user_api_key_auth", authenticate) + send = AsyncMock() + websocket = WebSocket({ + "type": "websocket", "scheme": "ws", "server": ("localhost", 4000), + "path": "/v1/live/" + token, "path_params": {"call_id": token}, "query_string": b"", + "headers": [(b"authorization", b"Bearer owner")] + + ([(b"x-proxy-key", custom_value)] if custom_value is not None else []), + }, AsyncMock(), send) + with pytest.raises(HTTPException) as error: + await auth_module.user_api_key_auth_websocket(websocket) + assert error.value.status_code == 403 + authenticate.assert_not_awaited() + send.assert_awaited_once_with({"type": "websocket.close", "code": 1008, "reason": ""}) diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py index 5d844b7be5d..32b13ad678a 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py @@ -172,7 +172,7 @@ async def test_realtime_endpoint_rejects_untrusted_call_ids(monkeypatch, call_id @pytest.mark.asyncio @pytest.mark.parametrize("multipart", [False, True]) -@pytest.mark.parametrize("credential", ["authorization", "api-key", "subprotocol"]) +@pytest.mark.parametrize("credential", ["authorization", "api-key", "subprotocol", "x-litellm-api-key", "custom"]) @pytest.mark.parametrize("signaling_credential", ["authorization", "api-key", "x-litellm-api-key", "mixed"]) async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, multipart, credential, signaling_credential): import json @@ -207,12 +207,12 @@ async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, "scheme": "http", "server": ("localhost", 80), "query_string": b"intent=quicksilver&architecture=avas&untrusted=bad", "headers": [(b"content-type", body_request.headers["content-type"].encode()), - *signaling_headers, (b"openai-alpha", b"quicksilver=v2"), + *signaling_headers, *([(b"x-proxy-key", b"Bearer owner")] if credential == "custom" else []), (b"openai-alpha", b"quicksilver=v2"), (b"x-untrusted", b"bad")]}, receive) auth = UserAPIKeyAuth() authorize = AsyncMock() monkeypatch.setattr(proxy_server, "master_key", "owner") - monkeypatch.setattr(proxy_server, "general_settings", {}) + monkeypatch.setattr(proxy_server, "general_settings", {"litellm_key_header_name": "x-proxy-key"} if credential == "custom" else {}) monkeypatch.setattr(codex, "can_key_call_resolved_model", authorize) class Processor: @@ -267,6 +267,8 @@ async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, credential_headers = { "authorization": [(b"authorization", b"Bearer owner")], "api-key": [(b"api-key", b"owner")], + "x-litellm-api-key": [(b"x-litellm-api-key", b"owner")], + "custom": [(b"x-proxy-key", b"Bearer owner")], "subprotocol": [(b"sec-websocket-protocol", b"realtime, openai-insecure-api-key.owner")], } websocket = WebSocket({"type": "websocket", "path": "/v1/live/opaque",