mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(chatgpt): share websocket credential selection for call ownership
This commit is contained in:
parent
bba4fd50e8
commit
a9cd1dc56d
4 changed files with 96 additions and 31 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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": ""})
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue