fix(logging): pass provider response headers to callbacks on every endpoint (#42824)

* fix(logging): pass provider response headers to callbacks on every endpoint

Custom callbacks only received kwargs["response_headers"] for chat
completions. Responses, image generation and edit, speech, and
transcription calls either never recorded the provider's headers or
recorded them in one place and not the other.

Every handler now records the provider's httpx headers on the response's
hidden params as "headers" (raw) and "additional_headers" (processed,
with LiteLLM's own entries winning on a clash), and the logging object
derives model_call_details["response_headers"] from those hidden params
before cost calculation on the non-stream and both streaming success
paths, keeping a handler-set value authoritative. Binary speech responses
expose their hidden params to the standard logging payload, and the sync
OpenAI transcription request always fetches the raw response.

* test(images): point the legacy image and speech fakes at the raw response surface

Image generation now goes through the SDK's raw response so the provider headers can be read, and the speech binary response now carries hidden params. The unit fakes in the image generation, xinference, proxy provider, image edit, Vertex speech, and otel suites still pinned the old call surface and the old "no hidden params" assertion, so they read an uncalled mock or a fake response without headers.

* test(images): drop the rewritten mock comments and the generated edit PNGs

* test(images): move the llm-span test's image fake to the raw response surface

---------

Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-09-24 13:01:12 -07:00 • committed by GitHub
parent c9e8a04139
commit 1edc4ba580
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
18 changed files with 714 additions and 79 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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