mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(chatgpt): authenticate sideband models and forward image headers
This commit is contained in:
parent
6e87b4b985
commit
57e831d445
4 changed files with 57 additions and 9 deletions
|
|
@ -402,6 +402,7 @@ def image_generation(
|
|||
model=model,
|
||||
prompt=prompt,
|
||||
image_generation_provider_config=image_generation_config,
|
||||
extra_headers=extra_headers,
|
||||
image_generation_optional_request_params=optional_params,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
litellm_params=litellm_params_dict,
|
||||
|
|
|
|||
|
|
@ -508,15 +508,6 @@ async def user_api_key_auth_websocket(websocket: WebSocket):
|
|||
|
||||
request._url = websocket.url
|
||||
|
||||
query_params: Final = websocket.query_params
|
||||
|
||||
model: Final = query_params.get("model")
|
||||
|
||||
async def return_body():
|
||||
return _realtime_request_body(model)
|
||||
|
||||
request.body = return_body
|
||||
|
||||
authorization: Final = websocket.headers.get("authorization")
|
||||
# If no Authorization header, try the api-key header
|
||||
if not authorization:
|
||||
|
|
@ -542,6 +533,19 @@ async def user_api_key_auth_websocket(websocket: WebSocket):
|
|||
# Call user_api_key_auth with the extracted API key
|
||||
# Note: You'll need to modify this to work with WebSocket context if needed
|
||||
try:
|
||||
from litellm.proxy.realtime_endpoints.call_sessions import decode_call
|
||||
|
||||
call_token: Final = websocket.path_params.get("call_id") or websocket.query_params.get("call_id")
|
||||
model: Final = (
|
||||
decode_call(call_token, f"Bearer {api_key}").alias
|
||||
if call_token is not None
|
||||
else websocket.query_params.get("model")
|
||||
)
|
||||
|
||||
async def return_body():
|
||||
return _realtime_request_body(model)
|
||||
|
||||
request.body = return_body
|
||||
return await user_api_key_auth(request=request, api_key=f"Bearer {api_key}")
|
||||
except Exception as e:
|
||||
if is_invalid_virtual_key_error(e):
|
||||
|
|
|
|||
|
|
@ -25,7 +25,9 @@ def test_generation_routes_with_chatgpt_oauth(chatgpt_tokens):
|
|||
quality="auto",
|
||||
size="auto",
|
||||
background="auto",
|
||||
extra_headers={"x-gateway-route": "images"},
|
||||
)
|
||||
assert requests[0].headers["x-gateway-route"] == "images"
|
||||
assert result.data[0].b64_json == "aGVsbG8="
|
||||
assert str(requests[0].url) == "https://chatgpt.com/backend-api/codex/images/generations"
|
||||
assert requests[0].headers["authorization"] == "Bearer test-token-" + "default"
|
||||
|
|
|
|||
|
|
@ -7133,3 +7133,44 @@ def test_user_api_key_auth_opens_a_datadog_span_for_accepted_and_rejected_keys(t
|
|||
assert report["outcomes"] == ["accepted", "rejected"]
|
||||
auth_span = "litellm.proxy.auth.user_api_key_auth.user_api_key_auth"
|
||||
assert [span for span in report["spans"] if span == auth_span] == [auth_span, auth_span]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("attachment", ["path", "query"])
|
||||
@pytest.mark.parametrize("credential", ["authorization", "api-key", "subprotocol"])
|
||||
@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
|
||||
import importlib
|
||||
import time
|
||||
from unittest.mock import AsyncMock
|
||||
from fastapi import WebSocket
|
||||
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-sideband-budget-salt")
|
||||
token = encode_call(CodexRealtimeCall(
|
||||
call_id="rtc_test", model="gpt-live-1-codex", alias="budgeted-voice",
|
||||
owner=hashlib.sha256(b"Bearer owner").hexdigest(), expires_at=time.time() + 300,
|
||||
))
|
||||
seen = []
|
||||
|
||||
async def authenticate(request, api_key):
|
||||
seen.append((await request.json(), api_key))
|
||||
return "authenticated-with-model"
|
||||
|
||||
monkeypatch.setattr(auth_module, "user_api_key_auth", authenticate)
|
||||
websocket = WebSocket({
|
||||
"type": "websocket", "scheme": "ws", "server": ("localhost", 4000),
|
||||
"path": "/v1/live/" + token if attachment == "path" else "/v1/realtime",
|
||||
"path_params": {"call_id": token} if attachment == "path" else {},
|
||||
"query_string": query_model + (b"&call_id=" + token.encode() if attachment == "query" else b""),
|
||||
"headers": {
|
||||
"authorization": [(b"authorization", b"Bearer owner")],
|
||||
"api-key": [(b"api-key", b"owner")],
|
||||
"subprotocol": [(b"sec-websocket-protocol", b"realtime, openai-insecure-api-key.owner")],
|
||||
}[credential],
|
||||
}, AsyncMock(), AsyncMock())
|
||||
assert await auth_module.user_api_key_auth_websocket(websocket) == "authenticated-with-model"
|
||||
assert seen == [({"model": "budgeted-voice"}, "Bearer owner")]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue