fix(chatgpt): pin sideband budgets and preserve OAuth image identity

This commit is contained in:
jibanez-staticduo 2026-09-09 21:48:44 +02:00
parent 09f52d9aa4
commit fe9a4ffca6
No known key found for this signature in database
5 changed files with 86 additions and 9 deletions

View file

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

View file

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

View file

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

View file

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

View file

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