fix(chatgpt): share websocket credential selection for call ownership

This commit is contained in:
jibanez-staticduo 2026-09-10 06:14:13 +02:00
parent bba4fd50e8
commit a9cd1dc56d
No known key found for this signature in database
4 changed files with 96 additions and 31 deletions

View file

@ -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

View file

@ -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,

View file

@ -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": ""})

View file

@ -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",