mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
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:
parent
c9e8a04139
commit
1edc4ba580
18 changed files with 714 additions and 79 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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 ##########################
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue