diff --git a/litellm/llms/chatgpt/images.py b/litellm/llms/chatgpt/images.py index 3467b68c31f..82845aa05f4 100644 --- a/litellm/llms/chatgpt/images.py +++ b/litellm/llms/chatgpt/images.py @@ -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): diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 30f12e3a64d..1c6cdefe907 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -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, diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 2a1d0709212..f6f4d7bf1c4 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -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) diff --git a/tests/test_litellm/llms/chatgpt/test_images.py b/tests/test_litellm/llms/chatgpt/test_images.py index 0f781475afe..2c51b932e3e 100644 --- a/tests/test_litellm/llms/chatgpt/test_images.py +++ b/tests/test_litellm/llms/chatgpt/test_images.py @@ -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): diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index 6df7f2118e0..886d3361c49 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -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": ""})