diff --git a/litellm/litellm_core_utils/core_helpers.py b/litellm/litellm_core_utils/core_helpers.py index 3afa6a913b5..b095b4b12c6 100644 --- a/litellm/litellm_core_utils/core_helpers.py +++ b/litellm/litellm_core_utils/core_helpers.py @@ -765,3 +765,30 @@ def set_response_cost_in_hidden_params(response: _CarriesHiddenParams, cost: flo RESPONSE_COST_HEADER: cost, } hidden_params["additional_headers"] = merged + + +_HIDDEN_PARAMS_ADAPTER: Final = TypeAdapter(Mapping[str, object]) +_PROVIDER_HEADERS_ADAPTER: Final = TypeAdapter(Mapping[str, str]) + + +def set_provider_response_headers_in_hidden_params( + response: _CarriesHiddenParams, headers: httpx.Headers | Mapping[str, str] +) -> None: + hidden_params: Final = response._hidden_params # pyright: ignore[reportPrivateUsage] # no public accessor + existing_additional_headers: Final[object] = hidden_params.get("additional_headers") + raw_headers: Final[dict[str, str]] = dict(headers) # mutable-ok: stored as the plain-dict hidden param + additional_headers: Final[dict[str, object]] = { # mutable-ok: assigned into the plain-dict hidden params + **process_response_headers(raw_headers), + **(existing_additional_headers if isinstance(existing_additional_headers, Mapping) else _NO_HEADERS), + } + hidden_params["headers"] = raw_headers + hidden_params["additional_headers"] = additional_headers + + +def get_provider_response_headers_from_hidden_params(response: object) -> Mapping[str, str] | None: + hidden_params: Final[object] = getattr(response, "_hidden_params", None) + try: + validated: Final = _HIDDEN_PARAMS_ADAPTER.validate_python(hidden_params) + return _PROVIDER_HEADERS_ADAPTER.validate_python(validated.get("headers")) + except ValidationError: + return None diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index a6391a2ae27..3b9419d483b 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -72,6 +72,7 @@ from litellm.litellm_core_utils.classifier_logging import ( is_classifier_call, ) from litellm.litellm_core_utils.core_helpers import ( + get_provider_response_headers_from_hidden_params, is_expected_client_error, reconstruct_model_name, set_response_cost_in_hidden_params, @@ -2353,6 +2354,15 @@ class Logging(LiteLLMLoggingBaseClass): ) return logging_result + def _surface_response_headers_from_result(self, logging_result: object) -> None: + existing: Final[object] = self.model_call_details.get("response_headers") + if existing is not None: + return + headers: Final = get_provider_response_headers_from_hidden_params(logging_result) + if headers is None: + return + self.model_call_details["response_headers"] = headers + def _merge_hidden_params_from_response_into_metadata(self, logging_result: object) -> None: """ Copy response._hidden_params into litellm_params.metadata['hidden_params']. @@ -2386,6 +2396,7 @@ class Logging(LiteLLMLoggingBaseClass): build_logging_payload: bool = True, ): """Resolve hidden params, compute response cost, and emit the standard logging payload.""" + self._surface_response_headers_from_result(logging_result) hidden_params: Final = getattr(logging_result, "_hidden_params", {}) if hidden_params: if self.model_call_details.get("litellm_params") is not None: @@ -2788,6 +2799,7 @@ class Logging(LiteLLMLoggingBaseClass): if complete_streaming_response is not None: verbose_logger.debug("Logging Details LiteLLM-Success Call streaming complete") self.model_call_details["complete_streaming_response"] = complete_streaming_response + self._surface_response_headers_from_result(complete_streaming_response) self.model_call_details["response_cost"] = self._response_cost_calculator( result=complete_streaming_response ) @@ -3302,6 +3314,7 @@ class Logging(LiteLLMLoggingBaseClass): print_verbose("Async success callbacks: Got a complete streaming response") self.model_call_details["async_complete_streaming_response"] = complete_streaming_response + self._surface_response_headers_from_result(complete_streaming_response) try: if self.model_call_details.get("cache_hit", False) is True: @@ -6362,12 +6375,15 @@ def _extract_response_obj_and_hidden_params( original_exception: Exception | None, ) -> tuple[dict, dict | None]: """Extract response_obj and hidden_params from init_response_obj.""" - hidden_params: dict | None = None + hidden_params: dict | None = ( + getattr(init_response_obj, "_hidden_params", None) + if isinstance(init_response_obj, BaseModel | HttpxBinaryResponseContent) + else None + ) if init_response_obj is None: response_obj = {} elif isinstance(init_response_obj, BaseModel): response_obj = init_response_obj.model_dump() - hidden_params = getattr(init_response_obj, "_hidden_params", None) elif isinstance(init_response_obj, dict): response_obj = init_response_obj elif isinstance(init_response_obj, HttpxBinaryResponseContent): diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index dba0dee38fc..4f31742aaaa 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -45,6 +45,7 @@ from litellm.litellm_core_utils.audio_utils.subtitle_utils import ( SUBTITLE_RESPONSE_FORMATS, synthesize_subtitle_document, ) +from litellm.litellm_core_utils.core_helpers import set_provider_response_headers_in_hidden_params from litellm.litellm_core_utils.get_litellm_params import AWS_CREDENTIAL_KWARGS_KEYS from litellm.litellm_core_utils.llm_request_utils import serialize_multipart_form_fields from litellm.litellm_core_utils.realtime_errors import ( @@ -1461,6 +1462,7 @@ class BaseLLMHTTPHandler: transformed: Final = provider_config.transform_audio_transcription_response( raw_response=response, ) + set_provider_response_headers_in_hidden_params(transformed, response.headers) if not provider_config.supports_subtitle_synthesis: return transformed requested_format: Final = optional_params.get("response_format") @@ -6960,11 +6962,13 @@ class BaseLLMHTTPHandler: provider_config=image_edit_provider_config, ) - return image_edit_provider_config.transform_image_edit_response( + image_edit_response: Final = image_edit_provider_config.transform_image_edit_response( model=model, raw_response=response, logging_obj=logging_obj, ) + set_provider_response_headers_in_hidden_params(image_edit_response, response.headers) + return image_edit_response async def async_image_edit_handler( self, @@ -7059,11 +7063,13 @@ class BaseLLMHTTPHandler: provider_config=image_edit_provider_config, ) - return image_edit_provider_config.transform_image_edit_response( + image_edit_response: Final = image_edit_provider_config.transform_image_edit_response( model=model, raw_response=response, logging_obj=logging_obj, ) + set_provider_response_headers_in_hidden_params(image_edit_response, response.headers) + return image_edit_response def image_generation_handler( self, @@ -7186,6 +7192,7 @@ class BaseLLMHTTPHandler: litellm_params=dict(litellm_params), encoding=None, ) + set_provider_response_headers_in_hidden_params(model_response, response.headers) return model_response @@ -7293,6 +7300,7 @@ class BaseLLMHTTPHandler: litellm_params=dict(litellm_params), encoding=None, ) + set_provider_response_headers_in_hidden_params(model_response, response.headers) return model_response @@ -12077,11 +12085,13 @@ class BaseLLMHTTPHandler: provider_config=text_to_speech_provider_config, ) - return text_to_speech_provider_config.transform_text_to_speech_response( + speech_response: Final = text_to_speech_provider_config.transform_text_to_speech_response( model=model, raw_response=response, logging_obj=logging_obj, ) + set_provider_response_headers_in_hidden_params(speech_response, response.headers) + return speech_response async def async_text_to_speech_handler( self, @@ -12176,11 +12186,13 @@ class BaseLLMHTTPHandler: provider_config=text_to_speech_provider_config, ) - return text_to_speech_provider_config.transform_text_to_speech_response( + speech_response: Final = text_to_speech_provider_config.transform_text_to_speech_response( model=model, raw_response=response, logging_obj=logging_obj, ) + set_provider_response_headers_in_hidden_params(speech_response, response.headers) + return speech_response ######################################################### ########## SKILLS API HANDLERS ########################## diff --git a/litellm/llms/openai/openai.py b/litellm/llms/openai/openai.py index 63874ca9619..d6340d182ae 100644 --- a/litellm/llms/openai/openai.py +++ b/litellm/llms/openai/openai.py @@ -27,6 +27,7 @@ from litellm import LlmProviders from litellm._logging import verbose_logger from litellm.constants import DEFAULT_MAX_RETRIES from litellm.files.types import FileContentStreamingResult +from litellm.litellm_core_utils.core_helpers import set_provider_response_headers_in_hidden_params from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.litellm_core_utils.logging_utils import speech_request_body, track_llm_api_timing from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator @@ -1404,7 +1405,6 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): organization: str | None = None, headers: dict | None = None, ): - response = None try: openai_aclient: Final = self._get_openai_client( is_async=True, @@ -1428,8 +1428,10 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): ) request_data: Final = {**data, "extra_headers": headers} if headers else data - response = await openai_aclient.images.generate(**request_data, timeout=timeout) - stringified_response: Final = response.model_dump() + raw_response: Final = await openai_aclient.images.with_raw_response.generate( + **request_data, timeout=timeout + ) + stringified_response: Final = raw_response.parse().model_dump() ## LOGGING logging_obj.post_call( input=prompt, @@ -1437,11 +1439,13 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): additional_args={"complete_input_dict": data}, original_response=stringified_response, ) - return convert_to_model_response_object( + image_response: Final[ImageResponse] = convert_to_model_response_object( response_object=stringified_response, model_response_object=model_response, response_type="image_generation", ) + set_provider_response_headers_in_hidden_params(image_response, raw_response.headers) + return image_response except Exception as e: ## LOGGING logging_obj.post_call( @@ -1512,9 +1516,9 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): ## COMPLETION CALL request_data: Final = {**data, "extra_headers": headers} if headers else data - _response: Final = openai_client.images.generate(**request_data, timeout=timeout) + raw_response: Final = openai_client.images.with_raw_response.generate(**request_data, timeout=timeout) - response: Final = _response.model_dump() + response: Final = raw_response.parse().model_dump() ## LOGGING logging_obj.post_call( input=prompt, @@ -1522,11 +1526,13 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): additional_args={"complete_input_dict": data}, original_response=response, ) - return convert_to_model_response_object( + image_response: Final[ImageResponse] = convert_to_model_response_object( response_object=response, model_response_object=model_response, response_type="image_generation", ) + set_provider_response_headers_in_hidden_params(image_response, raw_response.headers) + return image_response except OpenAIError as e: ## LOGGING logging_obj.post_call( @@ -1609,7 +1615,9 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): input=input, **optional_params, ) - return HttpxBinaryResponseContent(response=response.response) + speech_response: Final = HttpxBinaryResponseContent(response=response.response) + set_provider_response_headers_in_hidden_params(speech_response, response.response.headers) + return speech_response async def async_audio_speech( self, @@ -1655,8 +1663,9 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): input=input, **optional_params, ) - - return HttpxBinaryResponseContent(response=response.response) + speech_response: Final = HttpxBinaryResponseContent(response=response.response) + set_provider_response_headers_in_hidden_params(speech_response, response.response.headers) + return speech_response class OpenAIFilesAPI(BaseLLM): diff --git a/litellm/llms/openai/transcriptions/handler.py b/litellm/llms/openai/transcriptions/handler.py index 701b3d30362..014251db821 100644 --- a/litellm/llms/openai/transcriptions/handler.py +++ b/litellm/llms/openai/transcriptions/handler.py @@ -4,11 +4,10 @@ import httpx from openai import AsyncOpenAI, OpenAI from pydantic import BaseModel -import litellm - if TYPE_CHECKING: from aiohttp import ClientSession from litellm.litellm_core_utils.audio_utils.utils import get_audio_file_name +from litellm.litellm_core_utils.core_helpers import set_provider_response_headers_in_hidden_params from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.base_llm.audio_transcription.transformation import ( BaseAudioTranscriptionConfig, @@ -31,11 +30,6 @@ class OpenAIAudioTranscription(OpenAIChatCompletion): data: dict, timeout: float | httpx.Timeout, ): - """ - Helper to: - - call openai_aclient.audio.transcriptions.with_raw_response when litellm.return_response_headers is True - - call openai_aclient.audio.transcriptions.create by default - """ try: raw_response = await openai_aclient.audio.transcriptions.with_raw_response.create(**data, timeout=timeout) headers: Final = dict(raw_response.headers) @@ -51,20 +45,11 @@ class OpenAIAudioTranscription(OpenAIChatCompletion): data: dict, timeout: float | httpx.Timeout, ): - """ - Helper to: - - call openai_aclient.audio.transcriptions.with_raw_response when litellm.return_response_headers is True - - call openai_aclient.audio.transcriptions.create by default - """ try: - if litellm.return_response_headers is True: - raw_response = openai_client.audio.transcriptions.with_raw_response.create(**data, timeout=timeout) - headers: Final = dict(raw_response.headers) - response = raw_response.parse() - return headers, response - else: - response = openai_client.audio.transcriptions.create(**data, timeout=timeout) - return None, response + raw_response: Final = openai_client.audio.transcriptions.with_raw_response.create(**data, timeout=timeout) + headers: Final = dict(raw_response.headers) + response: Final = raw_response.parse() + return headers, response except Exception as e: raise e @@ -133,11 +118,12 @@ class OpenAIAudioTranscription(OpenAIChatCompletion): "complete_input_dict": data, }, ) - _, response = self.make_sync_openai_audio_transcriptions_request( + headers, response = self.make_sync_openai_audio_transcriptions_request( openai_client=openai_client, data=data, timeout=timeout, ) + logging_obj.model_call_details["response_headers"] = headers if isinstance(response, BaseModel): stringified_response = response.model_dump() @@ -158,6 +144,7 @@ class OpenAIAudioTranscription(OpenAIChatCompletion): hidden_params=hidden_params, response_type="audio_transcription", ) + set_provider_response_headers_in_hidden_params(final_response, headers) return final_response async def async_audio_transcriptions( @@ -217,12 +204,14 @@ class OpenAIAudioTranscription(OpenAIChatCompletion): actual_model: Final = data.get("model", "whisper-1") hidden_params: Final = {"model": actual_model, "custom_llm_provider": "openai"} - return convert_to_model_response_object( + final_response: Final[TranscriptionResponse] = convert_to_model_response_object( response_object=stringified_response, model_response_object=model_response, hidden_params=hidden_params, response_type="audio_transcription", ) + set_provider_response_headers_in_hidden_params(final_response, headers) + return final_response except Exception as e: ## LOGGING logging_obj.post_call( diff --git a/tests/image_gen_tests/test_image_edits.py b/tests/image_gen_tests/test_image_edits.py index 0c2f57066e8..36fd65ba71b 100644 --- a/tests/image_gen_tests/test_image_edits.py +++ b/tests/image_gen_tests/test_image_edits.py @@ -250,6 +250,7 @@ async def test_azure_image_edit_litellm_sdk(): self._json_data = json_data self.status_code = status_code self.text = json.dumps(json_data) + self.headers = {} def json(self): return self._json_data @@ -370,6 +371,7 @@ async def test_openai_image_edit_cost_tracking(): self._json_data = json_data self.status_code = status_code self.text = json.dumps(json_data) + self.headers = {} def json(self): return self._json_data @@ -460,6 +462,7 @@ async def test_azure_image_edit_cost_tracking(): self._json_data = json_data self.status_code = status_code self.text = json.dumps(json_data) + self.headers = {} def json(self): return self._json_data @@ -737,6 +740,7 @@ async def test_image_edit_array_handling(): self._json_data = json_data self.status_code = status_code self.text = json.dumps(json_data) + self.headers = {} def json(self): return self._json_data diff --git a/tests/image_gen_tests/test_xinference.py b/tests/image_gen_tests/test_xinference.py index 3dc4fee85da..76cae593e41 100644 --- a/tests/image_gen_tests/test_xinference.py +++ b/tests/image_gen_tests/test_xinference.py @@ -24,9 +24,14 @@ async def test_xinference_image_generation(): def model_dump(self): return mock_openai_response - # Create a mock client with the images.generate method + class MockRawResponse: + headers = {} + + def parse(self): + return MockResponse() + mock_client = AsyncMock() - mock_client.images.generate = AsyncMock(return_value=MockResponse()) + mock_client.images.with_raw_response.generate = AsyncMock(return_value=MockRawResponse()) # Capture the actual arguments sent to OpenAI client captured_args = None @@ -36,9 +41,9 @@ async def test_xinference_image_generation(): nonlocal captured_args, captured_kwargs captured_args = args captured_kwargs = kwargs - return MockResponse() + return MockRawResponse() - mock_client.images.generate.side_effect = capture_generate_call + mock_client.images.with_raw_response.generate.side_effect = capture_generate_call # Mock the _get_openai_client method to return our mock client with patch.object( @@ -65,7 +70,7 @@ async def test_xinference_image_generation(): assert response.data[0].url == "https://example.com/image.png" # Validate that the OpenAI client was called with correct parameters - mock_client.images.generate.assert_called_once() + mock_client.images.with_raw_response.generate.assert_called_once() assert captured_kwargs is not None assert ( captured_kwargs["model"] == "stabilityai/stable-diffusion-3.5-large" @@ -97,9 +102,14 @@ async def test_xinference_image_generation_with_response_format(): def model_dump(self): return mock_openai_response - # Create a mock client with the images.generate method + class MockRawResponse: + headers = {} + + def parse(self): + return MockResponse() + mock_client = AsyncMock() - mock_client.images.generate = AsyncMock(return_value=MockResponse()) + mock_client.images.with_raw_response.generate = AsyncMock(return_value=MockRawResponse()) # Capture the actual arguments sent to OpenAI client captured_args = None @@ -109,9 +119,9 @@ async def test_xinference_image_generation_with_response_format(): nonlocal captured_args, captured_kwargs captured_args = args captured_kwargs = kwargs - return MockResponse() + return MockRawResponse() - mock_client.images.generate.side_effect = capture_generate_call + mock_client.images.with_raw_response.generate.side_effect = capture_generate_call # Mock the _get_openai_client method to return our mock client with patch.object( @@ -141,7 +151,7 @@ async def test_xinference_image_generation_with_response_format(): assert response.data[0].b64_json is not None # Validate that the OpenAI client was called with correct parameters - mock_client.images.generate.assert_called_once() + mock_client.images.with_raw_response.generate.assert_called_once() assert captured_kwargs is not None assert ( captured_kwargs["model"] == "stabilityai/stable-diffusion-3.5-large" diff --git a/tests/llm_translation/test_litellm_proxy_provider.py b/tests/llm_translation/test_litellm_proxy_provider.py index a10fc55ecc5..a7a2848a514 100644 --- a/tests/llm_translation/test_litellm_proxy_provider.py +++ b/tests/llm_translation/test_litellm_proxy_provider.py @@ -210,11 +210,14 @@ async def test_litellm_gateway_image_generation_direct(is_async): "created": 1, "data": [{"url": "https://example.com/image.png"}], } + mock_raw_response = MagicMock() + mock_raw_response.parse.return_value = mock_openai_response + mock_raw_response.headers = {} if is_async: # Mock the AsyncOpenAI client that gets created inside _get_openai_client mock_async_client = AsyncMock() - mock_async_client.images.generate = AsyncMock(return_value=mock_openai_response) + mock_async_client.images.with_raw_response.generate = AsyncMock(return_value=mock_raw_response) with patch( "litellm.llms.openai.openai.AsyncOpenAI", return_value=mock_async_client @@ -234,14 +237,14 @@ async def test_litellm_gateway_image_generation_direct(is_async): assert constructor_kwargs["base_url"] == "http://my-proxy" # Verify the AsyncOpenAI client was called correctly - mock_async_client.images.generate.assert_awaited_once() - call_kwargs = mock_async_client.images.generate.call_args.kwargs + mock_async_client.images.with_raw_response.generate.assert_awaited_once() + call_kwargs = mock_async_client.images.with_raw_response.generate.call_args.kwargs assert call_kwargs["model"] == "dall-e-3" assert call_kwargs["prompt"] == "A beautiful sunset over mountains" else: # Mock the sync OpenAI client that gets created inside _get_openai_client mock_sync_client = MagicMock() - mock_sync_client.images.generate.return_value = mock_openai_response + mock_sync_client.images.with_raw_response.generate.return_value = mock_raw_response with patch( "litellm.llms.openai.openai.OpenAI", return_value=mock_sync_client @@ -260,8 +263,8 @@ async def test_litellm_gateway_image_generation_direct(is_async): assert constructor_kwargs["base_url"] == "http://my-proxy" # Verify the OpenAI client was called correctly - mock_sync_client.images.generate.assert_called_once() - call_kwargs = mock_sync_client.images.generate.call_args.kwargs + mock_sync_client.images.with_raw_response.generate.assert_called_once() + call_kwargs = mock_sync_client.images.with_raw_response.generate.call_args.kwargs assert call_kwargs["model"] == "dall-e-3" assert call_kwargs["prompt"] == "A beautiful sunset over mountains" @@ -285,6 +288,7 @@ async def test_litellm_gateway_from_sdk_image_edit(is_async): self._json_data = json_data self.status_code = status_code self.text = json.dumps(json_data) + self.headers = {} def json(self): return self._json_data diff --git a/tests/llm_translation/test_openai.py b/tests/llm_translation/test_openai.py index af4ba85d58e..0488c4c68e6 100644 --- a/tests/llm_translation/test_openai.py +++ b/tests/llm_translation/test_openai.py @@ -313,7 +313,7 @@ def test_openai_max_retries_0(mock_get_openai_client): def test_openai_image_generation_forwards_organization(mock_get_openai_client): """Ensure organization flows to OpenAI client for image generation.""" - class _DummyImages: + class _DummyRawImages: def generate(self, **kwargs): # type: ignore class _Resp: def model_dump(self_inner): # minimal OpenAI ImagesResponse shape @@ -327,7 +327,16 @@ def test_openai_image_generation_forwards_organization(mock_get_openai_client): }, } - return _Resp() + class _RawResp: + headers = {} + + def parse(self_inner): + return _Resp() + + return _RawResp() + + class _DummyImages: + with_raw_response = _DummyRawImages() class _DummyClient: def __init__(self): diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_sources_of_truth.py b/tests/test_litellm/integrations/otel/test_otel_v2_sources_of_truth.py index 7e93d3d67a7..57f4557c6f7 100644 --- a/tests/test_litellm/integrations/otel/test_otel_v2_sources_of_truth.py +++ b/tests/test_litellm/integrations/otel/test_otel_v2_sources_of_truth.py @@ -1207,14 +1207,18 @@ def test_speech_response_without_a_byte_count_produces_no_output() -> None: def test_speech_binary_response_is_logged_as_its_summary_not_dropped() -> None: import httpx + from litellm.litellm_core_utils.core_helpers import set_provider_response_headers_in_hidden_params from litellm.litellm_core_utils.litellm_logging import _extract_response_obj_and_hidden_params from litellm.types.llms.openai import HttpxBinaryResponseContent raw: Final = httpx.Response(200, headers={"content-type": "audio/mpeg"}, content=b"\x00" * 1234) - response_obj, hidden_params = _extract_response_obj_and_hidden_params(HttpxBinaryResponseContent(raw), None) + speech: Final = HttpxBinaryResponseContent(raw) + set_provider_response_headers_in_hidden_params(speech, raw.headers) + response_obj, hidden_params = _extract_response_obj_and_hidden_params(speech, None) assert response_obj == {"object": "binary", "content_type": "audio/mpeg", "num_bytes": 1234} - assert hidden_params is None + assert hidden_params is not None + assert hidden_params["headers"]["content-type"] == "audio/mpeg" def test_speech_binary_response_still_streaming_reports_the_bytes_downloaded_so_far() -> None: diff --git a/tests/test_litellm/litellm_core_utils/test_core_helpers.py b/tests/test_litellm/litellm_core_utils/test_core_helpers.py index 6eeea271127..2a6dd347d5f 100644 --- a/tests/test_litellm/litellm_core_utils/test_core_helpers.py +++ b/tests/test_litellm/litellm_core_utils/test_core_helpers.py @@ -2,22 +2,27 @@ import logging +import httpx import pytest from litellm.litellm_core_utils.core_helpers import ( _FINISH_REASON_MAP, + RESPONSE_COST_HEADER, bind_budget_reservation_to_callbacks, budget_reservation_from_metadata, drop_params_env_flag, drop_params_flag, get_or_create_metadata_bucket, + get_provider_response_headers_from_hidden_params, map_finish_reason, normalize_drop_params, reconstruct_model_name, redact_nested_match_and_regex_keys, + set_provider_response_headers_in_hidden_params, unbind_budget_reservation_from_callbacks, ) from litellm.proxy._types import UserAPIKeyAuth +from litellm.types.utils import ImageResponse, TranscriptionResponse class TestBudgetReservationBinding: @@ -489,3 +494,66 @@ class TestIsExpectedClientError: category=RateLimitErrorCategory.VENDOR_RATE_LIMIT, ) assert is_expected_client_error(vendor_limit) is False + + +class TestProviderResponseHeadersInHiddenParams: + def test_records_raw_headers_and_the_processed_additional_headers(self): + response = ImageResponse() + response._hidden_params = {"additional_headers": {RESPONSE_COST_HEADER: 0.04}} + + set_provider_response_headers_in_hidden_params( + response, httpx.Headers({"X-Request-Id": "req_img", "x-ratelimit-remaining-requests": "41"}) + ) + + assert response._hidden_params["headers"] == { + "x-request-id": "req_img", + "x-ratelimit-remaining-requests": "41", + } + additional_headers = response._hidden_params["additional_headers"] + assert additional_headers["llm_provider-x-request-id"] == "req_img" + assert additional_headers["x-ratelimit-remaining-requests"] == "41" + assert additional_headers[RESPONSE_COST_HEADER] == 0.04 + + def test_litellm_owned_additional_headers_win_over_provider_headers(self): + response = TranscriptionResponse(text="hi") + response._hidden_params = {"additional_headers": {"llm_provider-x-request-id": "kept"}} + + set_provider_response_headers_in_hidden_params(response, {"x-request-id": "provider"}) + + assert response._hidden_params["additional_headers"]["llm_provider-x-request-id"] == "kept" + assert response._hidden_params["headers"] == {"x-request-id": "provider"} + + def test_getter_returns_the_recorded_headers(self): + response = ImageResponse() + + set_provider_response_headers_in_hidden_params(response, {"x-request-id": "req_img"}) + + assert get_provider_response_headers_from_hidden_params(response) == {"x-request-id": "req_img"} + + @pytest.mark.parametrize( + "hidden_params", + [ + None, + "headers", + {"additional_headers": {}}, + {"headers": "x-request-id: req_img"}, + {"headers": {"x-request-id": 7}}, + ], + ) + def test_getter_returns_none_without_a_string_header_mapping(self, hidden_params): + response = ImageResponse() + response._hidden_params = hidden_params + + assert get_provider_response_headers_from_hidden_params(response) is None + + def test_getter_returns_none_for_an_object_without_hidden_params(self): + assert get_provider_response_headers_from_hidden_params(object()) is None + + def test_headers_never_leak_into_a_sibling_response(self): + recorded = TranscriptionResponse() + sibling = TranscriptionResponse() + + set_provider_response_headers_in_hidden_params(recorded, {"x-request-id": "req_stt"}) + + assert get_provider_response_headers_from_hidden_params(sibling) is None + assert "additional_headers" not in sibling._hidden_params diff --git a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py index 23c01841b1b..bb8098e6e23 100644 --- a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py +++ b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py @@ -24,6 +24,7 @@ from litellm.cost_calculator import ocr_batch_cost from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.litellm_logging import Logging as LitellmLogging from litellm.litellm_core_utils.litellm_logging import ( + _extract_response_obj_and_hidden_params, _get_status_fields, set_callbacks, ) @@ -32,6 +33,7 @@ from litellm.proxy._types import UserAPIKeyAuth from litellm.types.llms.openai import ResponseAPIUsage, ResponseCompletedEvent, ResponsesAPIResponse from litellm.types.utils import ( CallTypes, + ImageResponse, LiteLLMRealtimeStreamLoggingObject, ModelResponse, TextCompletionResponse, @@ -8694,3 +8696,88 @@ async def test_async_failure_handler_delivers_failure_payload_to_custom_logger() assert "smoke-failure" in payload["error_str"] assert payload["model"] == "openai/gpt-5.6" assert events.empty() + + +def _image_logging_obj() -> LitellmLogging: + logging_obj = LitellmLogging( + model="gpt-image-2", + messages="a cat", + stream=False, + call_type="aimage_generation", + start_time=time.time(), + litellm_call_id="response-headers-test", + function_id="response-headers-test", + ) + logging_obj.model_call_details["litellm_params"] = {"metadata": {}} + logging_obj.optional_params = {} + return logging_obj + + +def _image_result_with_headers(request_id: str) -> ImageResponse: + result = ImageResponse(created=1, data=[]) + result._hidden_params = {"headers": {"x-request-id": request_id}} + return result + + +def test_process_hidden_params_surfaces_response_headers_from_the_result(): + logging_obj = _image_logging_obj() + + logging_obj._process_hidden_params_and_response_cost( + _image_result_with_headers("req_img"), datetime.datetime.now(), datetime.datetime.now() + ) + + assert logging_obj.model_call_details["response_headers"] == {"x-request-id": "req_img"} + + +def test_process_hidden_params_keeps_handler_set_response_headers(): + logging_obj = _image_logging_obj() + logging_obj.model_call_details["response_headers"] = {"x-request-id": "from-handler"} + + logging_obj._process_hidden_params_and_response_cost( + _image_result_with_headers("from-result"), datetime.datetime.now(), datetime.datetime.now() + ) + + assert logging_obj.model_call_details["response_headers"] == {"x-request-id": "from-handler"} + + +def _assembled_stream_result_with_headers() -> ModelResponse: + result = _assembled_stream_result() + result._hidden_params = {"headers": {"x-request-id": "req_stream"}} + return result + + +@pytest.mark.asyncio +async def test_async_streaming_success_passes_result_headers_to_callback_kwargs(): + releasing = CustomLogger() + releasing.async_log_success_event = AsyncMock() + patcher, logging_obj = _streaming_logging_obj_with_callbacks([releasing]) + + with patcher: + await logging_obj.async_success_handler(result=_assembled_stream_result_with_headers()) + + kwargs = releasing.async_log_success_event.await_args.kwargs["kwargs"] + assert kwargs["response_headers"] == {"x-request-id": "req_stream"} + + +def test_sync_streaming_success_passes_result_headers_to_callback_kwargs(): + releasing = CustomLogger() + releasing.log_success_event = MagicMock() + patcher, logging_obj = _streaming_logging_obj_with_callbacks([releasing]) + + with patcher: + logging_obj.success_handler(result=_assembled_stream_result_with_headers()) + + kwargs = releasing.log_success_event.call_args.kwargs["kwargs"] + assert kwargs["response_headers"] == {"x-request-id": "req_stream"} + + +def test_extract_response_obj_and_hidden_params_reads_binary_content_hidden_params(): + from litellm.types.llms.openai import HttpxBinaryResponseContent as LiteLLMBinaryResponseContent + + result = LiteLLMBinaryResponseContent(response=httpx.Response(status_code=200, content=b"audio bytes")) + result._hidden_params = {"headers": {"x-request-id": "req_tts"}} + + response_obj, hidden_params = _extract_response_obj_and_hidden_params(result, None) + + assert hidden_params == {"headers": {"x-request-id": "req_tts"}} + assert response_obj["object"] == "binary" diff --git a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py index 68f37c8ffcc..75b6ce4c626 100644 --- a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py +++ b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py @@ -26,6 +26,8 @@ from litellm.llms.base_llm.search.transformation import BaseSearchConfig, Search from litellm.llms.bedrock.base_aws_llm import SignsRequestsWithAWS from litellm.llms.brave.search.transformation import BraveSearchConfig from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig +from litellm.llms.base_llm.image_generation.transformation import BaseImageGenerationConfig +from litellm.llms.base_llm.text_to_speech.transformation import BaseTextToSpeechConfig from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.llms.custom_httpx.llm_http_handler import ( BaseLLMHTTPHandler, @@ -41,7 +43,7 @@ from litellm.llms.bedrock.messages.invoke_transformations.anthropic_claude3_tran from litellm.llms.mistral.ocr.transformation import MistralOCRConfig from litellm.llms.openai.videos.transformation import OpenAIVideoConfig from litellm.llms.tinyfish.search.transformation import TinyfishSearchConfig -from litellm.types.llms.openai import ResponsesAPIResponse +from litellm.types.llms.openai import HttpxBinaryResponseContent, ResponsesAPIResponse from litellm.types.router import GenericLiteLLMParams from litellm.types.utils import ImageObject, ImageResponse, ModelResponse, TranscriptionResponse from tests.test_litellm.llms.bedrock.event_loop_probe import EventLoopProbe @@ -4166,3 +4168,202 @@ async def test_chat_completion_agentic_followup_does_not_repeat_request_params_f assert followup_calls[0]["temperature"] == 0.2 assert followup_calls[0]["api_base"] == "https://a" assert followup_calls[0]["model"] == "openai/gpt-5" + + +_UPSTREAM_HEADERS: Final = {"x-request-id": "req_upstream", "x-ratelimit-remaining-requests": "41"} + + +def _assert_upstream_headers_recorded(response) -> None: + assert response._hidden_params["headers"]["x-request-id"] == "req_upstream" + assert response._hidden_params["additional_headers"]["llm_provider-x-request-id"] == "req_upstream" + assert response._hidden_params["additional_headers"]["x-ratelimit-remaining-requests"] == "41" + + +def _json_with_upstream_headers(payload: dict) -> httpx.MockTransport: + return httpx.MockTransport(lambda request: httpx.Response(200, json=payload, headers=_UPSTREAM_HEADERS)) + + +def _binary_with_upstream_headers() -> httpx.MockTransport: + return httpx.MockTransport( + lambda request: httpx.Response( + 200, content=b"audio-bytes", headers={**_UPSTREAM_HEADERS, "content-type": "audio/mpeg"} + ) + ) + + +def test_audio_transcriptions_records_upstream_response_headers(): + client = HTTPHandler(client=httpx.Client(transport=_json_with_upstream_headers({"text": "transcribed"}))) + + response = BaseLLMHTTPHandler().audio_transcriptions( + client=client, + atranscription=False, + **_json_transcription_call_kwargs(_JSONBodyAudioTranscriptionConfig()), + ) + + _assert_upstream_headers_recorded(response) + + +@pytest.mark.asyncio +async def test_async_audio_transcriptions_records_upstream_response_headers(): + client = AsyncHTTPHandler() + client.client = httpx.AsyncClient(transport=_json_with_upstream_headers({"text": "transcribed"})) + + response = await BaseLLMHTTPHandler().async_audio_transcriptions( + client=client, + **_json_transcription_call_kwargs(_JSONBodyAudioTranscriptionConfig()), + ) + + _assert_upstream_headers_recorded(response) + + +def _image_edit_call_kwargs() -> dict: + return { + "model": "edit-model", + "image": b"raw-image", + "prompt": "add a hat", + "image_edit_provider_config": _ImageEditRecordingConfig(), + "image_edit_optional_request_params": {}, + "custom_llm_provider": "openai", + "litellm_params": GenericLiteLLMParams(), + "logging_obj": Mock(), + "timeout": 10.0, + } + + +def test_image_edit_handler_records_upstream_response_headers(): + client = HTTPHandler() + client.client = httpx.Client(transport=_json_with_upstream_headers({"transformed_by": "sync"})) + + response = BaseLLMHTTPHandler().image_edit_handler(client=client, **_image_edit_call_kwargs()) + + _assert_upstream_headers_recorded(response) + + +@pytest.mark.asyncio +async def test_async_image_edit_handler_records_upstream_response_headers(): + client = AsyncHTTPHandler() + client.client = httpx.AsyncClient(transport=_json_with_upstream_headers({"transformed_by": "async"})) + + response = await BaseLLMHTTPHandler().async_image_edit_handler(client=client, **_image_edit_call_kwargs()) + + _assert_upstream_headers_recorded(response) + + +class _HeaderImageGenerationConfig(BaseImageGenerationConfig): + def get_supported_openai_params(self, model): + return [] + + def map_openai_params(self, non_default_params, optional_params, model, drop_params): + return optional_params + + def get_complete_url(self, api_base, api_key, model, optional_params, litellm_params, stream=None): + return "https://images.example/v1/generations" + + def transform_image_generation_request(self, model, prompt, optional_params, litellm_params, headers): + return {"prompt": prompt} + + def transform_image_generation_response( + self, + model, + raw_response, + model_response, + logging_obj, + request_data, + optional_params, + litellm_params, + encoding, + api_key=None, + json_mode=None, + ): + return ImageResponse(data=[ImageObject(b64_json=raw_response.json()["b64_json"])]) + + +def _image_generation_call_kwargs() -> dict: + return { + "model": "image-model", + "prompt": "a cat", + "image_generation_provider_config": _HeaderImageGenerationConfig(), + "image_generation_optional_request_params": {}, + "custom_llm_provider": "openai", + "litellm_params": {}, + "logging_obj": Mock(), + "timeout": 10.0, + } + + +def test_image_generation_handler_records_upstream_response_headers(): + client = HTTPHandler() + client.client = httpx.Client(transport=_json_with_upstream_headers({"b64_json": "abc"})) + + response = BaseLLMHTTPHandler().image_generation_handler(client=client, **_image_generation_call_kwargs()) + + assert response.data[0].b64_json == "abc" + _assert_upstream_headers_recorded(response) + + +@pytest.mark.asyncio +async def test_async_image_generation_handler_records_upstream_response_headers(): + client = AsyncHTTPHandler() + client.client = httpx.AsyncClient(transport=_json_with_upstream_headers({"b64_json": "abc"})) + + response = await BaseLLMHTTPHandler().async_image_generation_handler( + client=client, **_image_generation_call_kwargs() + ) + + assert response.data[0].b64_json == "abc" + _assert_upstream_headers_recorded(response) + + +class _HeaderTextToSpeechConfig(BaseTextToSpeechConfig): + def get_supported_openai_params(self, model): + return [] + + def map_openai_params(self, model, optional_params, voice=None, drop_params=False, kwargs=None): + return voice, optional_params + + def validate_environment(self, headers, model, api_key=None, api_base=None): + return {} + + def get_complete_url(self, model, api_base, litellm_params): + return "https://tts.example/v1/speech" + + def transform_text_to_speech_request(self, model, input, voice, optional_params, litellm_params, headers): + return {"dict_body": {"input": input}} + + def transform_text_to_speech_response(self, model, raw_response, logging_obj): + return HttpxBinaryResponseContent(response=raw_response) + + +def _text_to_speech_call_kwargs() -> dict: + return { + "model": "tts-model", + "input": "hello", + "voice": "alloy", + "text_to_speech_provider_config": _HeaderTextToSpeechConfig(), + "text_to_speech_optional_params": {}, + "custom_llm_provider": "openai", + "litellm_params": {}, + "logging_obj": Mock(), + "timeout": 10.0, + } + + +def test_text_to_speech_handler_records_upstream_response_headers(): + client = HTTPHandler() + client.client = httpx.Client(transport=_binary_with_upstream_headers()) + + response = BaseLLMHTTPHandler().text_to_speech_handler(client=client, **_text_to_speech_call_kwargs()) + + assert response.content == b"audio-bytes" + _assert_upstream_headers_recorded(response) + + +@pytest.mark.asyncio +async def test_async_text_to_speech_handler_records_upstream_response_headers(): + client = AsyncHTTPHandler() + client.client = httpx.AsyncClient(transport=_binary_with_upstream_headers()) + + response = await BaseLLMHTTPHandler().async_text_to_speech_handler(client=client, **_text_to_speech_call_kwargs()) + + assert response.content == b"audio-bytes" + _assert_upstream_headers_recorded(response) diff --git a/tests/test_litellm/llms/openai/test_openai.py b/tests/test_litellm/llms/openai/test_openai.py index 9539e13a802..2be691fa65b 100644 --- a/tests/test_litellm/llms/openai/test_openai.py +++ b/tests/test_litellm/llms/openai/test_openai.py @@ -1,13 +1,15 @@ import asyncio import json from typing import Final +from unittest.mock import Mock import httpx import pytest -from openai import AsyncOpenAI +from openai import AsyncOpenAI, OpenAI import litellm from litellm.llms.openai.openai import OpenAIChatCompletion +from litellm.types.utils import ImageResponse @pytest.mark.parametrize( @@ -253,3 +255,98 @@ async def test_acompletion_streams_tool_call_arguments_over_injected_transport() assert tool_call.function.name == "get_weather" assert json.loads(tool_call.function.arguments) == {"city": "Paris"} assert rebuilt.choices[0].finish_reason == "tool_calls" + + +_PROVIDER_HEADERS: Final = {"x-request-id": "req_openai", "x-ratelimit-remaining-requests": "41"} + + +def _image_generation_transport() -> httpx.MockTransport: + return httpx.MockTransport( + lambda request: httpx.Response( + 200, json={"created": 1, "data": [{"b64_json": "abc"}]}, headers=_PROVIDER_HEADERS + ) + ) + + +def _speech_transport() -> httpx.MockTransport: + return httpx.MockTransport( + lambda request: httpx.Response( + 200, content=b"audio-bytes", headers={**_PROVIDER_HEADERS, "content-type": "audio/mpeg"} + ) + ) + + +def _assert_provider_headers_recorded(response) -> None: + assert response._hidden_params["headers"]["x-request-id"] == "req_openai" + assert response._hidden_params["additional_headers"]["llm_provider-x-request-id"] == "req_openai" + assert response._hidden_params["additional_headers"]["x-ratelimit-remaining-requests"] == "41" + + +def _image_generation_kwargs() -> dict: + return { + "model": "gpt-image-2", + "prompt": "a cat", + "timeout": 10, + "optional_params": {}, + "logging_obj": Mock(), + "api_key": "transport-only", + "model_response": ImageResponse(), + } + + +def test_image_generation_records_provider_response_headers(): + with httpx.Client(transport=_image_generation_transport()) as http_client: + response = OpenAIChatCompletion().image_generation( + client=OpenAI(api_key="transport-only", http_client=http_client), **_image_generation_kwargs() + ) + + _assert_provider_headers_recorded(response) + + +@pytest.mark.asyncio +async def test_aimage_generation_records_provider_response_headers(): + async with httpx.AsyncClient(transport=_image_generation_transport()) as http_client: + response = await OpenAIChatCompletion().image_generation( + client=AsyncOpenAI(api_key="transport-only", http_client=http_client), + aimg_generation=True, + **_image_generation_kwargs(), + ) + + _assert_provider_headers_recorded(response) + + +def _audio_speech_kwargs() -> dict: + return { + "model": "gpt-4o-mini-tts", + "input": "hello", + "voice": "alloy", + "optional_params": {}, + "api_key": "transport-only", + "api_base": None, + "organization": None, + "project": None, + "max_retries": 0, + "timeout": 10, + "logging_obj": Mock(), + } + + +def test_audio_speech_records_provider_response_headers(): + with httpx.Client(transport=_speech_transport()) as http_client: + response = OpenAIChatCompletion().audio_speech( + client=OpenAI(api_key="transport-only", http_client=http_client), **_audio_speech_kwargs() + ) + + _assert_provider_headers_recorded(response) + + +@pytest.mark.asyncio +async def test_async_audio_speech_records_provider_response_headers(): + async with httpx.AsyncClient(transport=_speech_transport()) as http_client: + response = await OpenAIChatCompletion().audio_speech( + client=AsyncOpenAI(api_key="transport-only", http_client=http_client), + aspeech=True, + **_audio_speech_kwargs(), + ) + + _assert_provider_headers_recorded(response) diff --git a/tests/test_litellm/llms/openai/transcriptions/test_openai_transcriptions_handler.py b/tests/test_litellm/llms/openai/transcriptions/test_openai_transcriptions_handler.py new file mode 100644 index 00000000000..f2dbb71fea1 --- /dev/null +++ b/tests/test_litellm/llms/openai/transcriptions/test_openai_transcriptions_handler.py @@ -0,0 +1,71 @@ +from typing import Final +from unittest.mock import Mock + +import httpx +import pytest +from openai import AsyncOpenAI, OpenAI + +from litellm.llms.openai.transcriptions.handler import OpenAIAudioTranscription +from litellm.types.utils import TranscriptionResponse + +_PROVIDER_HEADERS: Final = {"x-request-id": "req_stt", "x-ratelimit-remaining-requests": "41"} + + +def _transcription_transport() -> httpx.MockTransport: + return httpx.MockTransport(lambda request: httpx.Response(200, json={"text": "hello"}, headers=_PROVIDER_HEADERS)) + + +def _logging_obj() -> Mock: + logging_obj = Mock() + logging_obj.model_call_details = {} + return logging_obj + + +def _call_kwargs(logging_obj: Mock) -> dict: + return { + "model": "gpt-4o-mini-transcribe", + "audio_file": ("audio.wav", b"riff-bytes", "audio/wav"), + "optional_params": {}, + "litellm_params": {}, + "model_response": TranscriptionResponse(), + "timeout": 10.0, + "max_retries": 0, + "logging_obj": logging_obj, + "api_key": "transport-only", + "api_base": None, + } + + +def _assert_headers_recorded(response: TranscriptionResponse, logging_obj: Mock) -> None: + assert response.text == "hello" + assert response._hidden_params["headers"]["x-request-id"] == "req_stt" + assert response._hidden_params["additional_headers"]["llm_provider-x-request-id"] == "req_stt" + assert response._hidden_params["additional_headers"]["x-ratelimit-remaining-requests"] == "41" + assert logging_obj.model_call_details["response_headers"]["x-request-id"] == "req_stt" + + +def test_audio_transcriptions_records_provider_response_headers(): + logging_obj = _logging_obj() + + with httpx.Client(transport=_transcription_transport()) as http_client: + response = OpenAIAudioTranscription().audio_transcriptions( + client=OpenAI(api_key="transport-only", http_client=http_client), + atranscription=False, + **_call_kwargs(logging_obj), + ) + + _assert_headers_recorded(response, logging_obj) + + +@pytest.mark.asyncio +async def test_async_audio_transcriptions_records_provider_response_headers(): + logging_obj = _logging_obj() + + async with httpx.AsyncClient(transport=_transcription_transport()) as http_client: + response = await OpenAIAudioTranscription().audio_transcriptions( + client=AsyncOpenAI(api_key="transport-only", http_client=http_client), + atranscription=True, + **_call_kwargs(logging_obj), + ) + + _assert_headers_recorded(response, logging_obj) diff --git a/tests/test_litellm/test_non_chat_routes_open_llm_spans.py b/tests/test_litellm/test_non_chat_routes_open_llm_spans.py index d62959ccd43..02c95c4bb2a 100644 --- a/tests/test_litellm/test_non_chat_routes_open_llm_spans.py +++ b/tests/test_litellm/test_non_chat_routes_open_llm_spans.py @@ -46,9 +46,9 @@ class _FakeSpeech: )() -class _FakeImages: +class _FakeRawImages: async def generate(self, **kwargs: Any) -> Any: - return type( + parsed: Final = type( "_Images", (), { @@ -58,6 +58,16 @@ class _FakeImages: } }, )() + return type( + "_RawImages", + (), + {"parse": lambda self: parsed, "headers": httpx.Headers({"x-request-id": "req-image"})}, + )() + + +class _FakeImages: + def __init__(self) -> None: + self.with_raw_response = _FakeRawImages() class _FakeModerations: diff --git a/tests/unit/llms/openai/image_generation/test_openai_image_generation_extra_headers.py b/tests/unit/llms/openai/image_generation/test_openai_image_generation_extra_headers.py index 55ef74abd7b..11df07f5fea 100644 --- a/tests/unit/llms/openai/image_generation/test_openai_image_generation_extra_headers.py +++ b/tests/unit/llms/openai/image_generation/test_openai_image_generation_extra_headers.py @@ -12,6 +12,14 @@ import pytest from litellm.llms.openai.openai import OpenAIChatCompletion +from litellm.types.utils import ImageResponse + + +def _raw_image_response(mock_image_data): + raw_response = MagicMock() + raw_response.parse.return_value = mock_image_data + raw_response.headers = {"x-request-id": "req-image"} + return raw_response @pytest.fixture @@ -41,7 +49,7 @@ class TestImageGenerationExtraHeaders: } mock_openai_client = MagicMock() - mock_openai_client.images.generate.return_value = mock_image_data + mock_openai_client.images.with_raw_response.generate.return_value = _raw_image_response(mock_image_data) mock_openai_client.api_key = "test-key" mock_openai_client._base_url._uri_reference = "https://api.openai.com" @@ -58,7 +66,7 @@ class TestImageGenerationExtraHeaders: client=mock_openai_client, ) - _, kwargs = mock_openai_client.images.generate.call_args + _, kwargs = mock_openai_client.images.with_raw_response.generate.call_args assert kwargs.get("extra_headers") == test_headers def test_sync_image_generation_without_headers( @@ -72,7 +80,7 @@ class TestImageGenerationExtraHeaders: } mock_openai_client = MagicMock() - mock_openai_client.images.generate.return_value = mock_image_data + mock_openai_client.images.with_raw_response.generate.return_value = _raw_image_response(mock_image_data) mock_openai_client.api_key = "test-key" mock_openai_client._base_url._uri_reference = "https://api.openai.com" @@ -86,7 +94,7 @@ class TestImageGenerationExtraHeaders: client=mock_openai_client, ) - _, kwargs = mock_openai_client.images.generate.call_args + _, kwargs = mock_openai_client.images.with_raw_response.generate.call_args assert "extra_headers" not in kwargs @pytest.mark.asyncio @@ -101,7 +109,9 @@ class TestImageGenerationExtraHeaders: } mock_openai_client = MagicMock() - mock_openai_client.images.generate = AsyncMock(return_value=mock_image_data) + mock_openai_client.images.with_raw_response.generate = AsyncMock( + return_value=_raw_image_response(mock_image_data) + ) mock_openai_client.api_key = "test-key" test_headers = {"cf-aig-authorization": "Bearer custom-token"} @@ -109,7 +119,7 @@ class TestImageGenerationExtraHeaders: await openai_chat_completions.aimage_generation( prompt="A white cat", data={"model": "dall-e-3", "prompt": "A white cat"}, - model_response=MagicMock(), + model_response=ImageResponse(), timeout=60.0, logging_obj=mock_logging_obj, api_key="test-key", @@ -117,7 +127,7 @@ class TestImageGenerationExtraHeaders: client=mock_openai_client, ) - _, kwargs = mock_openai_client.images.generate.call_args + _, kwargs = mock_openai_client.images.with_raw_response.generate.call_args assert kwargs.get("extra_headers") == test_headers @pytest.mark.asyncio @@ -132,20 +142,22 @@ class TestImageGenerationExtraHeaders: } mock_openai_client = MagicMock() - mock_openai_client.images.generate = AsyncMock(return_value=mock_image_data) + mock_openai_client.images.with_raw_response.generate = AsyncMock( + return_value=_raw_image_response(mock_image_data) + ) mock_openai_client.api_key = "test-key" await openai_chat_completions.aimage_generation( prompt="A white cat", data={"model": "dall-e-3", "prompt": "A white cat"}, - model_response=MagicMock(), + model_response=ImageResponse(), timeout=60.0, logging_obj=mock_logging_obj, api_key="test-key", client=mock_openai_client, ) - _, kwargs = mock_openai_client.images.generate.call_args + _, kwargs = mock_openai_client.images.with_raw_response.generate.call_args assert "extra_headers" not in kwargs @pytest.mark.parametrize("is_async", [False, True]) @@ -169,11 +181,13 @@ class TestImageGenerationExtraHeaders: test_headers = {"cf-aig-authorization": "Bearer custom-token"} if is_async: - mock_openai_client.images.generate = AsyncMock(return_value=mock_image_data) + mock_openai_client.images.with_raw_response.generate = AsyncMock( + return_value=_raw_image_response(mock_image_data) + ) await openai_chat_completions.aimage_generation( prompt="A white cat", data={"model": "dall-e-3", "prompt": "A white cat"}, - model_response=MagicMock(), + model_response=ImageResponse(), timeout=60.0, logging_obj=mock_logging_obj, api_key="test-key", @@ -181,7 +195,7 @@ class TestImageGenerationExtraHeaders: client=mock_openai_client, ) else: - mock_openai_client.images.generate.return_value = mock_image_data + mock_openai_client.images.with_raw_response.generate.return_value = _raw_image_response(mock_image_data) openai_chat_completions.image_generation( model="dall-e-3", prompt="A white cat", @@ -197,7 +211,7 @@ class TestImageGenerationExtraHeaders: "complete_input_dict" ] assert "extra_headers" not in logged_body - _, kwargs = mock_openai_client.images.generate.call_args + _, kwargs = mock_openai_client.images.with_raw_response.generate.call_args assert kwargs.get("extra_headers") == test_headers def test_sync_image_generation_forwards_headers_to_async( @@ -242,7 +256,9 @@ class TestImageGenerationEntryPointHeaders: } mock_openai_client = MagicMock() - mock_openai_client.images.generate = AsyncMock(return_value=mock_image_data) + mock_openai_client.images.with_raw_response.generate = AsyncMock( + return_value=_raw_image_response(mock_image_data) + ) mock_openai_client.api_key = "test-key" mock_openai_client._base_url._uri_reference = "https://api.openai.com" @@ -256,6 +272,6 @@ class TestImageGenerationEntryPointHeaders: api_key="test-key", ) - mock_openai_client.images.generate.assert_called_once() - _, kwargs = mock_openai_client.images.generate.call_args + mock_openai_client.images.with_raw_response.generate.assert_called_once() + _, kwargs = mock_openai_client.images.with_raw_response.generate.call_args assert kwargs.get("extra_headers") == test_headers diff --git a/tests/unit/llms/vertex_ai/text_to_speech/test_transformation.py b/tests/unit/llms/vertex_ai/text_to_speech/test_transformation.py index b5eec42b569..ee7bdebe745 100644 --- a/tests/unit/llms/vertex_ai/text_to_speech/test_transformation.py +++ b/tests/unit/llms/vertex_ai/text_to_speech/test_transformation.py @@ -526,6 +526,7 @@ class TestVertexAILyriaTextToSpeechConfig: ): mock_response = Mock(spec=httpx.Response) mock_response.status_code = 200 + mock_response.headers = {"content-type": "application/json"} mock_response.json.return_value = response_json with ( patch.object( # test-quality-ok: litellm.speech has no seam for Vertex token minting