mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(chatgpt): pin sideband budgets and preserve OAuth image identity
This commit is contained in:
parent
09f52d9aa4
commit
fe9a4ffca6
5 changed files with 86 additions and 9 deletions
|
|
@ -49,6 +49,12 @@ def encode_reference(
|
|||
}
|
||||
|
||||
|
||||
def without_image_identity_headers(headers: Mapping[str, object]) -> Mapping[str, object]:
|
||||
return MappingProxyType(
|
||||
{key: value for key, value in headers.items() if key.lower() not in ("authorization", "chatgpt-account-id")}
|
||||
)
|
||||
|
||||
|
||||
def image_headers(
|
||||
headers: Mapping[str, object], model: str, params: Mapping[str, object]
|
||||
) -> dict[str, object]: # mutable-ok: image handler requires dictionaries
|
||||
|
|
@ -57,7 +63,11 @@ def image_headers(
|
|||
model=model,
|
||||
litellm_params=GenericLiteLLMParams.model_validate(params),
|
||||
)
|
||||
return {**headers, **auth_headers, "accept": "application/json"} # mutable-ok: JSON request serialization
|
||||
return { # mutable-ok: image handler requires dictionaries
|
||||
**without_image_identity_headers(headers),
|
||||
**auth_headers,
|
||||
"accept": "application/json",
|
||||
}
|
||||
|
||||
|
||||
class ChatGPTImageGenerationConfig(GPTImageGenerationConfig):
|
||||
|
|
|
|||
|
|
@ -6638,6 +6638,14 @@ class BaseLLMHTTPHandler:
|
|||
else:
|
||||
raise Exception(f"Unexpected error while closing WebSocket: {close_error}")
|
||||
|
||||
@staticmethod
|
||||
def _image_extra_headers(custom_llm_provider: str, headers: Mapping[str, object]) -> Mapping[str, object]:
|
||||
if custom_llm_provider == "chatgpt":
|
||||
from litellm.llms.chatgpt.images import without_image_identity_headers
|
||||
|
||||
return without_image_identity_headers(headers)
|
||||
return headers
|
||||
|
||||
def image_edit_handler(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -6694,7 +6702,7 @@ class BaseLLMHTTPHandler:
|
|||
)
|
||||
|
||||
if extra_headers:
|
||||
headers.update(extra_headers)
|
||||
headers.update(self._image_extra_headers(custom_llm_provider, extra_headers))
|
||||
|
||||
api_base: Final = image_edit_provider_config.get_complete_url(
|
||||
model=model,
|
||||
|
|
@ -6793,7 +6801,7 @@ class BaseLLMHTTPHandler:
|
|||
)
|
||||
|
||||
if extra_headers:
|
||||
headers.update(extra_headers)
|
||||
headers.update(self._image_extra_headers(custom_llm_provider, extra_headers))
|
||||
|
||||
api_base: Final = image_edit_provider_config.get_complete_url(
|
||||
model=model,
|
||||
|
|
@ -6910,7 +6918,7 @@ class BaseLLMHTTPHandler:
|
|||
)
|
||||
|
||||
if extra_headers:
|
||||
headers.update(extra_headers)
|
||||
headers.update(self._image_extra_headers(custom_llm_provider, extra_headers))
|
||||
|
||||
api_base: Final = image_generation_provider_config.get_complete_url(
|
||||
model=model,
|
||||
|
|
@ -7017,7 +7025,7 @@ class BaseLLMHTTPHandler:
|
|||
)
|
||||
|
||||
if extra_headers:
|
||||
headers.update(extra_headers)
|
||||
headers.update(self._image_extra_headers(custom_llm_provider, extra_headers))
|
||||
|
||||
api_base: Final = image_generation_provider_config.get_complete_url(
|
||||
model=model,
|
||||
|
|
|
|||
|
|
@ -302,7 +302,7 @@ async def _check_key_model_budget_with_fallback(
|
|||
model=model_name,
|
||||
)
|
||||
except litellm.BudgetExceededError as e:
|
||||
if request_data.get("model") != model_name:
|
||||
if request_data.get("model") != model_name or request.scope.get("litellm_pinned_realtime_model") == model_name:
|
||||
raise e
|
||||
fallback_model: Final = await model_max_budget_limiter.get_fallback_model_within_budget(
|
||||
user_api_key_dict=valid_token,
|
||||
|
|
@ -542,6 +542,8 @@ async def user_api_key_auth_websocket(websocket: WebSocket):
|
|||
if call_token is not None
|
||||
else websocket.query_params.get("model")
|
||||
)
|
||||
if call_token is not None:
|
||||
request.scope["litellm_pinned_realtime_model"] = model
|
||||
|
||||
async def return_body():
|
||||
return _realtime_request_body(model)
|
||||
|
|
|
|||
|
|
@ -27,7 +27,7 @@ def test_generation_routes_with_chatgpt_oauth(chatgpt_tokens, api_base):
|
|||
quality="auto",
|
||||
size="auto",
|
||||
background="auto",
|
||||
extra_headers={"x-gateway-route": "images"},
|
||||
extra_headers={"x-gateway-route": "images", "aUtHoRiZaTiOn": "Bearer wrong", "CHATGPT-ACCOUNT-ID": "wrong"},
|
||||
)
|
||||
assert requests[0].headers["x-gateway-route"] == "images"
|
||||
assert result.data[0].b64_json == "aGVsbG8="
|
||||
|
|
@ -50,6 +50,7 @@ def test_codex_json_edit_survives_sdk_dispatch(chatgpt_tokens):
|
|||
result = litellm.image_edit(
|
||||
model="chatgpt/gpt-image-2",
|
||||
prompt="red circle",
|
||||
extra_headers={"x-gateway-route": "images", "authorization": "Bearer wrong", "CHATGPT-ACCOUNT-ID": "wrong"},
|
||||
images=references,
|
||||
client=client,
|
||||
quality="auto",
|
||||
|
|
@ -61,6 +62,10 @@ def test_codex_json_edit_survives_sdk_dispatch(chatgpt_tokens):
|
|||
|
||||
assert json.loads(requests[0].content)["images"] == references
|
||||
|
||||
assert requests[0].headers["authorization"] == "Bearer test-token-default"
|
||||
assert requests[0].headers["chatgpt-account-id"] == "test-account-default"
|
||||
assert requests[0].headers["x-gateway-route"] == "images"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"references", [[], [{"image_url": "file:///etc/passwd"}], [{}], [{"image_url": "https://example.com/a.png"}] * 6]
|
||||
|
|
@ -82,9 +87,10 @@ def test_edit_converts_multipart_image_bytes():
|
|||
|
||||
def test_image_auth_does_not_accept_inbound_override(chatgpt_tokens):
|
||||
headers = ChatGPTImageGenerationConfig().validate_environment(
|
||||
{"Authorization": "Bearer wrong"}, "gpt-image-2", [], {}, {"chatgpt_token_dir": chatgpt_tokens}
|
||||
{"authorization": "Bearer wrong", "CHATGPT-ACCOUNT-ID": "wrong"}, "gpt-image-2", [], {}, {"chatgpt_token_dir": chatgpt_tokens}
|
||||
)
|
||||
assert headers["Authorization"] == "Bearer test-token-default"
|
||||
assert httpx.Headers(headers)["authorization"] == "Bearer test-token-default"
|
||||
assert httpx.Headers(headers)["chatgpt-account-id"] == "test-account-default"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -100,6 +106,7 @@ async def test_async_codex_edit_without_multipart_image(chatgpt_tokens):
|
|||
response = await litellm.aimage_edit(
|
||||
model="chatgpt/gpt-image-2",
|
||||
prompt="red circle",
|
||||
extra_headers={"x-gateway-route": "images", "authorization": "Bearer wrong", "CHATGPT-ACCOUNT-ID": "wrong"},
|
||||
client=client,
|
||||
images=[{"image_url": "data:image/png;base64,aGVsbG8="}],
|
||||
chatgpt_auth_profile="account3",
|
||||
|
|
@ -109,6 +116,10 @@ async def test_async_codex_edit_without_multipart_image(chatgpt_tokens):
|
|||
assert requests[0].headers["content-type"] == "application/json"
|
||||
await client.client.aclose()
|
||||
|
||||
assert requests[0].headers["authorization"] == "Bearer test-token-default"
|
||||
assert requests[0].headers["chatgpt-account-id"] == "test-account-default"
|
||||
assert requests[0].headers["x-gateway-route"] == "images"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("as_tuple", [False, True])
|
||||
def test_edit_accepts_filesystem_path(tmp_path, as_tuple):
|
||||
|
|
|
|||
|
|
@ -7290,3 +7290,49 @@ async def test_jwt_builder_returns_every_team_grant_the_key_path_gets(is_proxy_a
|
|||
assert token.team_member == Member(user_id="jwt-user", role="admin")
|
||||
assert token.team_member_spend == 1.5
|
||||
assert token.jwt_claims == {"sub": "jwt-user"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("attachment", ["path", "query"])
|
||||
async def test_sideband_rejects_budget_fallback_before_rerouting(monkeypatch, attachment):
|
||||
import hashlib
|
||||
import importlib
|
||||
import time
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock
|
||||
from fastapi import HTTPException, 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,
|
||||
))
|
||||
limiter = SimpleNamespace(
|
||||
is_key_within_model_budget=AsyncMock(side_effect=litellm.BudgetExceededError(current_cost=2, max_budget=1)),
|
||||
get_fallback_model_within_budget=AsyncMock(return_value="cheap-voice"),
|
||||
)
|
||||
auth = UserAPIKeyAuth(models=["budgeted-voice", "cheap-voice"])
|
||||
|
||||
async def authenticate(request, api_key):
|
||||
data = await request.json()
|
||||
await auth_module._check_key_model_budget_with_fallback(auth, limiter, data["model"], data, request)
|
||||
return auth
|
||||
|
||||
monkeypatch.setattr(auth_module, "user_api_key_auth", authenticate)
|
||||
monkeypatch.setattr(auth_module, "can_key_call_model", AsyncMock())
|
||||
send = AsyncMock()
|
||||
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": b"call_id=" + token.encode() if attachment == "query" else b"",
|
||||
"headers": [(b"authorization", b"Bearer owner")],
|
||||
}, AsyncMock(), send)
|
||||
with pytest.raises(HTTPException) as error:
|
||||
await auth_module.user_api_key_auth_websocket(websocket)
|
||||
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": ""})
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue