fix(images): bill unmeasurable references as one megapixel and keep measurement off the event loop

This commit is contained in:
shrey kharbanda 2026-09-24 08:45:11 +00:00
parent 6c58cb4267
commit 5f576ae5e8
10 changed files with 263 additions and 276 deletions

View file

@ -1,7 +1,7 @@
import base64
import os
import struct
from collections.abc import Iterator, Mapping, Sequence
from collections.abc import Iterable, Iterator, Mapping, Sized
from dataclasses import dataclass
from io import BytesIO
from typing import IO, Final
@ -33,6 +33,7 @@ _WEBP_VP8_START_CODE: Final = b"\x9d\x01\x2a"
_WEBP_VP8L_SIGNATURE: Final = 0x2F
_BMP_CORE_DIB_SIZE: Final = 12
_BMP_KNOWN_DIB_SIZES: Final = frozenset({12, 40, 52, 56, 64, 108, 124})
_UNMEASURED_REFERENCE_PIXELS: Final = 1024 * 1024
@dataclass(frozen=True, slots=True)
@ -45,34 +46,34 @@ class ImageDimensions:
return self.width * self.height
def total_reference_pixels(images: Sequence[FileTypes]) -> int | None:
"""All-or-nothing sum of measured reference pixels; never raises."""
def total_reference_pixels(images: Iterable[FileTypes]) -> int:
"""Reference pixels sent; never raises. A reference that cannot be measured is still metered by the
provider, so it counts as one megapixel rather than disappearing from spend."""
try:
measured: Final = tuple(read_image_dimensions(image) for image in images)
except Exception: # noqa: BLE001 # a billing helper must fall back, never fail the request
return None
dimensions: Final = tuple(size for size in measured if size is not None)
if len(dimensions) != len(measured):
return _UNMEASURED_REFERENCE_PIXELS * (len(images) if isinstance(images, Sized) else 1)
unmeasured: Final = sum(size is None for size in measured)
if unmeasured:
verbose_logger.debug(
"Reference image %d could not be measured (path, non-seekable stream, or no readable PNG, JPEG, WebP, "
"GIF or BMP header); billing generated pixels only",
measured.index(None),
"%d reference image(s) could not be measured (path, non-seekable stream, or no readable PNG, JPEG, "
"WebP, GIF or BMP header); billing each as one megapixel",
unmeasured,
)
return None
return sum(size.pixels for size in dimensions)
return sum(_UNMEASURED_REFERENCE_PIXELS if size is None else size.pixels for size in measured)
def uploaded_reference_pixels(
files: RequestFiles | None,
json_body: Mapping[str, object] | None,
) -> int | None:
) -> int:
"""Pixels across the image-bearing parts a request actually sends.
Multipart requests carry image content under ``files``; JSON requests embed base64 inside
``json_body`` (e.g. FLUX ``input_image`` fields). Measuring the outgoing payload instead of the
caller's arguments keeps billing aligned with what the provider meters when a transform filters
images out or adds parts of its own (masks, single-image providers). All-or-nothing, ``0`` when
nothing image-bearing is sent.
images out or adds parts of its own (masks, single-image providers). ``0`` when nothing
image-bearing is sent.
"""
parts: Final = _file_parts(files) + tuple(_embedded_image_values(json_body))
if not parts:
@ -110,7 +111,7 @@ def read_image_dimensions(image: FileTypes) -> ImageDimensions | None:
if embedded is None:
return None
return _dimensions_from_stream(BytesIO(embedded))
stream = BytesIO(content) if isinstance(content, (bytes, bytearray, memoryview)) else content
stream: Final = BytesIO(content) if isinstance(content, (bytes, bytearray, memoryview)) else content
return _dimensions_from_stream(stream)
except Exception: # noqa: BLE001 # an odd stream or malformed file must fall back, never raise
return None
@ -127,7 +128,7 @@ def _dimensions_from_stream(stream: IO[bytes]) -> ImageDimensions | None:
return None
position: Final = stream.tell()
try:
dimensions = _header_dimensions(stream, 0)
dimensions: Final = _header_dimensions(stream, 0)
finally:
stream.seek(position)
if dimensions is None or dimensions.width <= 0 or dimensions.height <= 0:

View file

@ -125,7 +125,7 @@ class AzureFoundryFlux2ImageEditConfig(OpenAIImageEditConfig):
if isinstance(image, bytes):
image_bytes = image
elif hasattr(image, "read"):
if image.seekable():
if hasattr(image, "seekable") and image.seekable(): # pyright: ignore[reportAny] # image is a duck-typed file-like object by contract
image.seek(0)
image_bytes = image.read()
else:

View file

@ -6994,7 +6994,7 @@ class BaseLLMHTTPHandler:
headers=headers,
)
data = image_edit_provider_config.finalize_image_edit_request_data(data, api_base)
reference_pixels: Final = uploaded_reference_pixels(files, data)
reference_pixels: Final = await asyncio.to_thread(uploaded_reference_pixels, files, data)
## LOGGING
logging_obj.pre_call(
@ -7223,7 +7223,7 @@ class BaseLLMHTTPHandler:
litellm_params=dict(litellm_params),
headers=headers,
)
reference_pixels: Final = uploaded_reference_pixels(None, data)
reference_pixels: Final = await asyncio.to_thread(uploaded_reference_pixels, None, data)
## LOGGING
logging_obj.pre_call(

View file

@ -9,6 +9,7 @@ import pytest
from PIL import Image
from litellm.images.dimensions import (
_UNMEASURED_REFERENCE_PIXELS,
ImageDimensions,
read_image_dimensions,
total_reference_pixels,
@ -309,30 +310,32 @@ def test_read_image_dimensions_returns_none_for_short_and_none_tuples():
assert read_image_dimensions(("ref.png", None)) is None
def test_total_reference_pixels_returns_none_instead_of_raising():
assert total_reference_pixels([("ref.png",)]) is None
assert total_reference_pixels(cast(Any, 123)) is None
def test_total_reference_pixels_bills_unmeasurable_input_as_one_megapixel():
assert total_reference_pixels([("ref.png",)]) == _UNMEASURED_REFERENCE_PIXELS
assert total_reference_pixels(cast(Any, 123)) == _UNMEASURED_REFERENCE_PIXELS
def test_total_reference_pixels_returns_none_for_closed_and_read_only_streams():
def test_total_reference_pixels_bills_closed_and_read_only_streams_as_one_megapixel():
closed: Final = io.BytesIO(_png(64, 64))
closed.close()
assert total_reference_pixels([closed]) is None
assert total_reference_pixels([cast(Any, _ReadOnlyObject())]) is None
assert total_reference_pixels([closed]) == _UNMEASURED_REFERENCE_PIXELS
assert total_reference_pixels([cast(Any, _ReadOnlyObject())]) == _UNMEASURED_REFERENCE_PIXELS
def test_total_reference_pixels_returns_none_when_a_reference_has_invalid_dimensions():
def test_total_reference_pixels_bills_invalid_dimensions_as_one_megapixel():
negative_width_bmp: Final = _bmp(-1024, 768)
assert read_image_dimensions(negative_width_bmp) is None
assert total_reference_pixels([negative_width_bmp, _png(1024, 1024)]) is None
assert total_reference_pixels([negative_width_bmp, _png(1024, 1024)]) == _UNMEASURED_REFERENCE_PIXELS + 1024 * 1024
def test_total_reference_pixels_returns_none_when_any_reference_is_unmeasurable():
def test_total_reference_pixels_bills_each_unmeasurable_reference_as_one_megapixel():
assert total_reference_pixels([_png(64, 64), CAT_JPEG.read_bytes()]) == 64 * 64 + 512 * 512
assert total_reference_pixels([_png(64, 64), _gif(b"GIF89a", 100, 50)]) == 64 * 64 + 100 * 50
assert total_reference_pixels([_png(64, 64), io.BytesIO(b"image"), CAT_JPEG.read_bytes()]) is None
assert total_reference_pixels([_png(64, 64), io.BytesIO(b"image"), CAT_JPEG.read_bytes()]) == (
64 * 64 + _UNMEASURED_REFERENCE_PIXELS + 512 * 512
)
assert total_reference_pixels([]) == 0
@ -343,31 +346,42 @@ def test_uploaded_reference_pixels_measures_the_uploaded_set_not_the_requested_s
# a transform that keeps only the first image (MAI) bills only what it sends
assert uploaded_reference_pixels({"image[]": ("ref.png", _png(64, 64))}, {"prompt": "hi"}) == 64 * 64
# FLUX.2 embeds each reference as a base64 field in the JSON body
assert uploaded_reference_pixels(
[],
{"model": "FLUX.2-flex", "prompt": "blend", "input_image": png_b64, "input_image_2": jpeg_b64, "n": 2},
) == 64 * 64 + 512 * 512
assert (
uploaded_reference_pixels(
[],
{"model": "FLUX.2-flex", "prompt": "blend", "input_image": png_b64, "input_image_2": jpeg_b64, "n": 2},
)
== 64 * 64 + 512 * 512
)
# an extra file part a transform adds itself (a mask) is uploaded and metered too
assert uploaded_reference_pixels(
{"image": ("edit.png", _png(64, 64)), "mask": ("mask.png", _png(10, 10))}, {"prompt": "cut"}
) == 64 * 64 + 10 * 10
assert (
uploaded_reference_pixels(
{"image": ("edit.png", _png(64, 64)), "mask": ("mask.png", _png(10, 10))}, {"prompt": "cut"}
)
== 64 * 64 + 10 * 10
)
def test_uploaded_reference_pixels_skips_non_image_fields_and_nested_bodies():
png_b64: Final = base64.b64encode(_png(64, 64)).decode()
jpeg_b64: Final = base64.b64encode(CAT_JPEG.read_bytes()).decode()
assert uploaded_reference_pixels(
{"image": ("edit.png", _png(64, 64))},
{"model": "m", "prompt": "hi", "size": "1024x1024", "extra": {"nested": jpeg_b64}, "refs": [png_b64]},
) == 64 * 64 + 512 * 512 + 64 * 64
assert (
uploaded_reference_pixels(
{"image": ("edit.png", _png(64, 64))},
{"model": "m", "prompt": "hi", "size": "1024x1024", "extra": {"nested": jpeg_b64}, "refs": [png_b64]},
)
== 64 * 64 + 512 * 512 + 64 * 64
)
assert uploaded_reference_pixels(None, {"prompt": "hi", "model": "m"}) == 0
assert uploaded_reference_pixels(None, None) == 0
def test_uploaded_reference_pixels_returns_none_when_any_uploaded_part_is_unmeasurable():
def test_uploaded_reference_pixels_bills_each_unmeasurable_part_as_one_megapixel():
png_b64: Final = base64.b64encode(_png(64, 64)).decode()
truncated_png_b64: Final = base64.b64encode(_png(64, 64)[:16]).decode()
assert uploaded_reference_pixels({"image": io.BytesIO(b"not an image")}, None) is None
assert uploaded_reference_pixels(None, {"input_image": png_b64, "input_image_2": truncated_png_b64}) is None
assert uploaded_reference_pixels({"image": io.BytesIO(b"not an image")}, None) == _UNMEASURED_REFERENCE_PIXELS
assert uploaded_reference_pixels(None, {"input_image": png_b64, "input_image_2": truncated_png_b64}) == (
64 * 64 + _UNMEASURED_REFERENCE_PIXELS
)

View file

@ -269,8 +269,6 @@ class TestImageEditCustomPricing:
assert use_custom_pricing_for_model(litellm_params) is False
def test_image_edit_forwards_custom_pricing_kwargs_to_logging(self):
"""Pricing a deployment declares under litellm_params (e.g. input_cost_per_pixel)
reaches self.litellm_params so use_custom_pricing_for_model fires on image edits."""
from litellm.images.main import image_edit
captured_litellm_params = {}
@ -344,9 +342,7 @@ class TestImageEditHandlerCredentialsForwarding:
"vertex_ai_credentials": "/path/to/creds.json",
}
with patch.object(
config, "_ensure_access_token", return_value=("token", "project")
) as mock_ensure:
with patch.object(config, "_ensure_access_token", return_value=("token", "project")) as mock_ensure:
config.validate_environment(
headers={},
model="test-model",
@ -375,9 +371,7 @@ class TestImageEditHandlerCredentialsForwarding:
"vertex_ai_credentials": "/path/to/creds.json",
}
with patch.object(
config, "_ensure_access_token", return_value=("token", "project")
) as mock_ensure:
with patch.object(config, "_ensure_access_token", return_value=("token", "project")) as mock_ensure:
config.validate_environment(
headers={},
model="test-model",
@ -447,10 +441,6 @@ class TestImageEditHandlerCredentialsForwarding:
params = list(sig.parameters.keys())
assert "litellm_params" in params, (
f"{config.__class__.__name__}.validate_environment "
"missing litellm_params parameter"
)
assert "api_base" in params, (
f"{config.__class__.__name__}.validate_environment "
"missing api_base parameter"
f"{config.__class__.__name__}.validate_environment missing litellm_params parameter"
)
assert "api_base" in params, f"{config.__class__.__name__}.validate_environment missing api_base parameter"

View file

@ -0,0 +1,39 @@
import base64
import io
from typing import Final
from litellm.llms.azure_ai.image_edit.flux2_transformation import AzureFoundryFlux2ImageEditConfig
class _ReadOnlyStream:
def __init__(self, content: bytes) -> None:
self._content: Final = content
def read(self, size: int = -1) -> bytes:
return self._content
class _NonSeekableStream(io.BytesIO):
def seekable(self) -> bool:
return False
def test_convert_image_to_base64_accepts_streams_without_seekable() -> None:
encoded: Final = AzureFoundryFlux2ImageEditConfig()._convert_image_to_base64(_ReadOnlyStream(b"png-bytes"))
assert base64.b64decode(encoded) == b"png-bytes"
def test_convert_image_to_base64_reads_non_seekable_streams_without_rewinding() -> None:
encoded: Final = AzureFoundryFlux2ImageEditConfig()._convert_image_to_base64(_NonSeekableStream(b"png-bytes"))
assert base64.b64decode(encoded) == b"png-bytes"
def test_convert_image_to_base64_rewinds_seekable_streams_before_encoding() -> None:
stream: Final = io.BytesIO(b"png-bytes")
stream.read(4)
encoded: Final = AzureFoundryFlux2ImageEditConfig()._convert_image_to_base64(stream)
assert base64.b64decode(encoded) == b"png-bytes"

View file

@ -11,10 +11,16 @@ from litellm.litellm_core_utils.llm_cost_calc.utils import CostCalculatorUtils
from litellm.llms.azure.azure import AzureChatCompletion
from litellm.llms.azure.image_generation import get_azure_image_generation_config
from litellm.llms.azure.image_generation.http_utils import azure_deployment_image_generation_json_body
from litellm.llms.azure_ai.image_generation.cost_calculator import cost_calculator
from litellm.llms.azure_ai.image_generation.flux_transformation import (
AzureFoundryFluxImageGenerationConfig,
)
from litellm.types.utils import ImageObject, ImageResponse, ImageUsage
from litellm.types.utils import (
ImageObject,
ImageResponse,
ImageUsage,
ImageUsageInputTokensDetails,
)
from litellm.utils import _invalidate_model_cost_lowercase_map, get_optional_params_image_gen
@ -265,9 +271,7 @@ _ONE_MEGAPIXEL: Final = 1024 * 1024
def test_flux2_cost_bills_references_once_for_multi_image_edits() -> None:
response: Final = ImageResponse(
data=[ImageObject(b64_json="aW1n"), ImageObject(b64_json="aW1n")]
)
response: Final = ImageResponse(data=[ImageObject(b64_json="aW1n"), ImageObject(b64_json="aW1n")])
response.set_reference_pixels(_ONE_MEGAPIXEL)
cost: Final = CostCalculatorUtils.route_image_generation_cost_calculator(
@ -380,8 +384,6 @@ def test_flux2_cost_bills_pixels_when_usage_carries_no_token_rates() -> None:
def test_flux2_cost_reads_pricing_declared_in_litellm_params_kwargs() -> None:
"""A deployment declared under litellm_params (e.g. litellm.image_edit(..., input_cost_per_pixel=x))
prices the same model_info dict the router's deployment YAML folds into metadata.model_info."""
deployment_rate: Final = 1e-06
response: Final = ImageResponse(data=[ImageObject(b64_json="aW1n")])
response.set_reference_pixels(_ONE_MEGAPIXEL)
@ -402,8 +404,6 @@ def test_flux2_cost_reads_pricing_declared_in_litellm_params_kwargs() -> None:
def test_flux2_cost_overlays_litellm_params_kwargs_on_nested_model_info() -> None:
"""Declared kwargs win per key over the router-folded metadata.model_info, so a call
that overrides one price key keeps the deployment's other declared pricing."""
nested_rate: Final = 1e-06
kwargs_rate: Final = 2e-06
response: Final = ImageResponse(data=[ImageObject(b64_json="aW1n")])
@ -425,3 +425,98 @@ def test_flux2_cost_overlays_litellm_params_kwargs_on_nested_model_info() -> Non
)
assert cost == pytest.approx(kwargs_rate * _ONE_MEGAPIXEL * 2)
def _edit_response(reference_pixels: int | None = _ONE_MEGAPIXEL * 2, **kwargs: object) -> ImageResponse:
response: Final = ImageResponse(data=[ImageObject(b64_json="aW1n")], size="1024x1024", **kwargs)
if reference_pixels is not None:
response.set_reference_pixels(reference_pixels)
return response
def _edit_cost(response: ImageResponse, size: str = "1024x1024", **kwargs: object) -> float:
return CostCalculatorUtils.route_image_generation_cost_calculator(
model="FLUX.2-flex",
completion_response=response,
custom_llm_provider="azure_ai",
size=size,
call_type="image_edit",
**kwargs,
)
def test_get_model_info_surfaces_flux2_flex_pixel_rate() -> None:
model_info = litellm.get_model_info(model="FLUX.2-flex", custom_llm_provider="azure_ai")
assert model_info["input_cost_per_pixel"] == _catalog_pixel_rate()
def test_flux2_cost_adds_reference_pixels_to_generated_pixels() -> None:
cost: Final = _edit_cost(_edit_response())
assert cost == pytest.approx(_catalog_pixel_rate() * _ONE_MEGAPIXEL * 3)
def test_flux2_cost_prefers_deployment_rates_over_catalog_for_edits() -> None:
cost: Final = _edit_cost(_edit_response(), model_info={"input_cost_per_pixel": 2e-07})
assert cost == pytest.approx(2e-07 * _ONE_MEGAPIXEL * 3)
def test_flux2_cost_honors_explicit_zero_pixel_rate() -> None:
cost: Final = _edit_cost(_edit_response(), model_info={"input_cost_per_pixel": 0.0})
assert cost == pytest.approx(0.0)
def test_flux2_cost_per_image_rate_beats_per_pixel_rate() -> None:
cost: Final = _edit_cost(
_edit_response(),
model_info={"output_cost_per_image": 0.04, "input_cost_per_pixel": 1e-07},
)
assert cost == pytest.approx(0.04 + 1e-07 * _ONE_MEGAPIXEL * 2)
def test_flux2_cost_bills_token_rates_when_usage_carries_them() -> None:
response: Final = _edit_response(
usage=ImageUsage(
input_tokens=150,
input_tokens_details=ImageUsageInputTokensDetails(image_tokens=100, text_tokens=50),
output_tokens=1000,
total_tokens=1150,
)
)
cost: Final = _edit_cost(
response,
model_info={
"input_cost_per_pixel": 5e-08,
"input_cost_per_token": 1e-05,
"input_cost_per_image_token": 2e-05,
"output_cost_per_image_token": 4e-05,
"output_cost_per_token": 4e-05,
},
)
assert cost == pytest.approx(50 * 1e-05 + 100 * 2e-05 + 1000 * 4e-05)
def test_flux2_cost_parses_dashed_size_string() -> None:
cost: Final = _edit_cost(_edit_response(), size="1024-x-1024")
assert cost == pytest.approx(_catalog_pixel_rate() * _ONE_MEGAPIXEL * 3)
def test_flux2_cost_optional_params_dimensions_beat_response_size() -> None:
cost: Final = _edit_cost(
_edit_response(reference_pixels=None),
optional_params={"width": 2048, "height": 1024},
)
assert cost == pytest.approx(_catalog_pixel_rate() * 2048 * 1024)
def test_flux2_cost_rejects_non_image_response() -> None:
with pytest.raises(TypeError, match="must be of type ImageResponse"):
cost_calculator(model="FLUX.2-flex", image_response=object())

View file

@ -1,165 +0,0 @@
from typing import Final
import pytest
import litellm
from litellm.litellm_core_utils.get_model_cost_map import GetModelCostMap
from litellm.litellm_core_utils.llm_cost_calc.utils import CostCalculatorUtils
from litellm.llms.azure_ai.image_generation.cost_calculator import cost_calculator
from litellm.types.utils import ImageObject, ImageResponse, ImageUsage, ImageUsageInputTokensDetails
from litellm.utils import _invalidate_model_cost_lowercase_map
REFERENCE_PIXELS: Final = 2 * 1024 * 1024
@pytest.fixture(autouse=True)
def use_local_model_cost_map(monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(litellm, "model_cost", GetModelCostMap.load_local_model_cost_map())
litellm.get_model_info.cache_clear()
_invalidate_model_cost_lowercase_map()
yield
litellm.get_model_info.cache_clear()
_invalidate_model_cost_lowercase_map()
def _edit_response(reference_pixels: int | None = REFERENCE_PIXELS, **kwargs: object) -> ImageResponse:
response: Final = ImageResponse(data=[ImageObject(b64_json="aW1n")], size="1024x1024", **kwargs)
if reference_pixels is not None:
response.set_reference_pixels(reference_pixels)
return response
def _edit_cost(response: ImageResponse, size: str = "1024x1024", **kwargs: object) -> float:
return CostCalculatorUtils.route_image_generation_cost_calculator(
model="FLUX.2-flex",
completion_response=response,
custom_llm_provider="azure_ai",
size=size,
call_type="image_edit",
**kwargs,
)
def test_get_model_info_surfaces_flux2_flex_pixel_rate() -> None:
model_info = litellm.get_model_info(model="FLUX.2-flex", custom_llm_provider="azure_ai")
catalog_info = litellm.model_cost["azure_ai/FLUX.2-flex"]
assert model_info["input_cost_per_pixel"] == catalog_info["input_cost_per_pixel"]
def test_flux2_flex_catalog_pixel_rate_is_azure_megapixel_price() -> None:
catalog_info = litellm.model_cost["azure_ai/FLUX.2-flex"]
assert catalog_info["input_cost_per_pixel"] * 1024 * 1024 == pytest.approx(0.05), (
"Azure Retail Prices API, product 'Azure BFL Flux Models', meters 'Flex Megapixel' and "
"'Flex Ref Megapixel' are $0.05 per MP where 1 MP = 1024x1024 pixels; confirmed against "
"Azure Cost Management usage on 2026-09-22 (1024x1024 image metered as 1.0 MP)"
)
def test_edit_cost_adds_reference_pixels_to_generated_pixels() -> None:
catalog_rate: Final = litellm.model_cost["azure_ai/FLUX.2-flex"]["input_cost_per_pixel"]
cost: Final = _edit_cost(_edit_response())
assert cost == pytest.approx(catalog_rate * 1024 * 1024 + catalog_rate * REFERENCE_PIXELS)
def test_edit_cost_prefers_deployment_rates_over_catalog() -> None:
cost: Final = _edit_cost(_edit_response(), model_info={"input_cost_per_pixel": 2e-07})
assert cost == pytest.approx(2e-07 * (1024 * 1024 + REFERENCE_PIXELS))
def test_edit_cost_honors_explicit_zero_pixel_rate() -> None:
cost: Final = _edit_cost(_edit_response(), model_info={"input_cost_per_pixel": 0.0})
assert cost == pytest.approx(0.0)
def test_per_image_rate_beats_per_pixel_rate(monkeypatch: pytest.MonkeyPatch) -> None:
cost: Final = _edit_cost(
_edit_response(),
model_info={"output_cost_per_image": 0.04, "input_cost_per_pixel": 1e-07},
)
assert cost == pytest.approx(0.04 + 1e-07 * REFERENCE_PIXELS)
def test_usage_short_circuits_pixel_billing() -> None:
response: Final = _edit_response(
usage=ImageUsage(
input_tokens=150,
input_tokens_details=ImageUsageInputTokensDetails(image_tokens=100, text_tokens=50),
output_tokens=1000,
total_tokens=1150,
)
)
cost: Final = _edit_cost(
response,
model_info={
"input_cost_per_pixel": 5e-08,
"input_cost_per_token": 1e-05,
"input_cost_per_image_token": 2e-05,
"output_cost_per_image_token": 4e-05,
"output_cost_per_token": 4e-05,
},
)
assert cost == pytest.approx(50 * 1e-05 + 100 * 2e-05 + 1000 * 4e-05)
def test_edit_without_measurement_bills_generated_pixels_only() -> None:
catalog_rate: Final = litellm.model_cost["azure_ai/FLUX.2-flex"]["input_cost_per_pixel"]
cost: Final = _edit_cost(_edit_response(reference_pixels=None))
assert cost == pytest.approx(catalog_rate * 1024 * 1024)
@pytest.mark.parametrize(
("size", "expected_generated"),
[
("1024-x-1024", "computed"),
("auto", "reference-only"),
("garbage", "reference-only"),
],
ids=["dashed-size", "auto-size-unparsed", "garbage-size-unparsed"],
)
def test_size_string_parsing(size: str, expected_generated: str) -> None:
catalog_rate: Final = litellm.model_cost["azure_ai/FLUX.2-flex"]["input_cost_per_pixel"]
expected: Final = (
catalog_rate * 1024 * 1024 if expected_generated == "computed" else 0.0
) + catalog_rate * REFERENCE_PIXELS
cost: Final = _edit_cost(_edit_response(), size=size)
assert cost == pytest.approx(expected)
def test_optional_params_dimensions_beat_response_size() -> None:
cost: Final = _edit_cost(
_edit_response(reference_pixels=None),
optional_params={"width": 2048, "height": 1024},
)
assert cost == pytest.approx(litellm.model_cost["azure_ai/FLUX.2-flex"]["input_cost_per_pixel"] * 2048 * 1024)
def test_unlisted_model_bills_deployment_rates() -> None:
cost: Final = CostCalculatorUtils.route_image_generation_cost_calculator(
model="unlisted-flux-deployment",
completion_response=_edit_response(),
custom_llm_provider="azure_ai",
size="1024x1024",
call_type="image_edit",
model_info={"input_cost_per_pixel": 1e-07},
)
assert cost == pytest.approx(1e-07 * (1024 * 1024 + REFERENCE_PIXELS))
def test_non_image_response_raises() -> None:
with pytest.raises(TypeError, match="must be of type ImageResponse"):
cost_calculator(model="FLUX.2-flex", image_response=object())

View file

@ -1442,9 +1442,7 @@ def test_sync_delete_responses_sets_json_content_type():
({}, True, None, None),
],
)
def test_resolve_anthropic_messages_timeout(
monkeypatch, litellm_params_kwargs, stream, global_timeout, expected
):
def test_resolve_anthropic_messages_timeout(monkeypatch, litellm_params_kwargs, stream, global_timeout, expected):
from litellm.constants import DEFAULT_REQUEST_TIMEOUT_SECONDS
if global_timeout is None:
@ -1460,9 +1458,7 @@ def test_resolve_anthropic_messages_timeout(
)
else:
monkeypatch.setattr("litellm.request_timeout", global_timeout, raising=False)
monkeypatch.setattr(
"litellm.request_timeout_explicitly_set", True, raising=False
)
monkeypatch.setattr("litellm.request_timeout_explicitly_set", True, raising=False)
resolved = BaseLLMHTTPHandler._resolve_anthropic_messages_timeout(
litellm_params=GenericLiteLLMParams(**litellm_params_kwargs),
@ -1487,9 +1483,7 @@ async def test_async_anthropic_messages_handler_forwards_request_timeout(monkeyp
return_value=({"x-api-key": "k"}, "https://api.anthropic.com")
)
mock_config.should_filter_anthropic_beta_headers = Mock(return_value=False)
mock_config.transform_anthropic_messages_request = Mock(
return_value={"model": "claude", "messages": []}
)
mock_config.transform_anthropic_messages_request = Mock(return_value={"model": "claude", "messages": []})
mock_config.get_complete_url = Mock(return_value="https://api.anthropic.com/v1/messages")
mock_config.sign_request = Mock(return_value=({"x-api-key": "k"}, None))
mock_config.max_retry_on_anthropic_messages_http_error = 1
@ -1535,9 +1529,7 @@ async def test_async_anthropic_messages_handler_forwards_stream_timeout(monkeypa
return_value=({"x-api-key": "k"}, "https://api.anthropic.com")
)
mock_config.should_filter_anthropic_beta_headers = Mock(return_value=False)
mock_config.transform_anthropic_messages_request = Mock(
return_value={"model": "claude", "messages": []}
)
mock_config.transform_anthropic_messages_request = Mock(return_value={"model": "claude", "messages": []})
mock_config.get_complete_url = Mock(return_value="https://api.anthropic.com/v1/messages")
mock_config.sign_request = Mock(return_value=({"x-api-key": "k"}, None))
mock_config.max_retry_on_anthropic_messages_http_error = 1
@ -1947,7 +1939,13 @@ async def test_async_anthropic_messages_handler_passes_api_key_to_agentic_hooks(
)
mock_config.sign_request = Mock(return_value=({}, None))
fake_raw_response = {"id": "msg_1", "type": "message", "role": "assistant", "content": [], "stop_reason": "end_turn"}
fake_raw_response = {
"id": "msg_1",
"type": "message",
"role": "assistant",
"content": [],
"stop_reason": "end_turn",
}
mock_config.transform_anthropic_messages_response = Mock(return_value=fake_raw_response)
mock_logging_obj = Mock()
@ -1967,10 +1965,17 @@ async def test_async_anthropic_messages_handler_passes_api_key_to_agentic_hooks(
mock_httpx_response.status_code = 200
with (
patch.object(handler, "_async_post_anthropic_messages_with_http_error_retry", new=AsyncMock(return_value=mock_httpx_response)),
patch.object(
handler,
"_async_post_anthropic_messages_with_http_error_retry",
new=AsyncMock(return_value=mock_httpx_response),
),
patch.object(handler, "_call_agentic_completion_hooks", side_effect=fake_agentic_hooks),
patch("litellm.llms.custom_httpx.llm_http_handler.get_async_httpx_client"),
patch("litellm.litellm_core_utils.get_provider_specific_headers.ProviderSpecificHeaderUtils.get_provider_specific_headers", return_value=None),
patch(
"litellm.litellm_core_utils.get_provider_specific_headers.ProviderSpecificHeaderUtils.get_provider_specific_headers",
return_value=None,
),
):
result = await handler.async_anthropic_messages_handler(
model="claude-haiku",
@ -2029,7 +2034,9 @@ async def test_async_anthropic_messages_handler_passes_deployment_api_base_to_ag
super().__init__()
self.hook_kwargs: dict | None = None
async def async_should_run_agentic_loop(self, response, model, messages, tools, stream, custom_llm_provider, kwargs):
async def async_should_run_agentic_loop(
self, response, model, messages, tools, stream, custom_llm_provider, kwargs
):
self.hook_kwargs = dict(kwargs)
return False, {}
@ -2303,7 +2310,9 @@ def test_audio_transcriptions_sends_dict_data_as_json_body():
form-encodes it and silently ignores json=; JSON-body providers (e.g.
Google Speech-to-Text) need an application/json body."""
captured = {}
client = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(_capture_json_transcription_request(captured))))
client = HTTPHandler(
client=httpx.Client(transport=httpx.MockTransport(_capture_json_transcription_request(captured)))
)
response = BaseLLMHTTPHandler().audio_transcriptions(
client=client,
@ -2387,9 +2396,7 @@ def _transform_subtitle_response(payload):
def test_subtitle_synthesis_fallback_without_timings_drops_words():
response = _transform_subtitle_response(
{"text": "hello world", "words": [{"word": "hello"}, {"word": "world"}]}
)
response = _transform_subtitle_response({"text": "hello world", "words": [{"word": "hello"}, {"word": "world"}]})
assert response.text == "hello world"
assert "words" not in response
@ -2609,9 +2616,7 @@ async def test_anthropic_invalid_thinking_signature_retry_resigns_bedrock_reques
ok_response = httpx.Response(200, json={"id": "msg_1"}, request=httpx.Request("POST", request_url))
class FakeAsyncClient:
async def post(
self, url, headers, data, stream=False, logging_obj=None, timeout=None
):
async def post(self, url, headers, data, stream=False, logging_obj=None, timeout=None):
posts.append({"headers": dict(headers), "data": data})
return invalid_signature_response if len(posts) == 1 else ok_response
@ -3228,7 +3233,13 @@ def _capture_video_create_request(captured):
captured["body"] = request.content
return httpx.Response(
200,
json={"id": "video_123", "object": "video", "status": "queued", "created_at": 1712697600, "model": "sora-2"},
json={
"id": "video_123",
"object": "video",
"status": "queued",
"created_at": 1712697600,
"model": "sora-2",
},
)
return respond
@ -3251,7 +3262,9 @@ def test_video_generation_without_file_sends_multipart_form_data():
captured = {}
client = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(_capture_video_create_request(captured))))
result = BaseLLMHTTPHandler().video_generation_handler(client=client, **_video_create_call_kwargs(OpenAIVideoConfig()))
result = BaseLLMHTTPHandler().video_generation_handler(
client=client, **_video_create_call_kwargs(OpenAIVideoConfig())
)
assert captured["content_type"].startswith("multipart/form-data")
assert _multipart_text_fields(captured["content_type"], captured["body"]) == {
@ -3291,7 +3304,9 @@ def test_azure_video_generation_without_file_sends_multipart_form_data():
captured = {}
client = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(_capture_video_create_request(captured))))
result = BaseLLMHTTPHandler().video_generation_handler(client=client, **_video_create_call_kwargs(AzureVideoConfig()))
result = BaseLLMHTTPHandler().video_generation_handler(
client=client, **_video_create_call_kwargs(AzureVideoConfig())
)
assert captured["content_type"].startswith("multipart/form-data")
assert _multipart_text_fields(captured["content_type"], captured["body"]) == {
@ -3306,7 +3321,9 @@ def test_video_generation_json_provider_keeps_json_body():
captured = {}
client = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(_capture_video_create_request(captured))))
result = BaseLLMHTTPHandler().video_generation_handler(client=client, **_video_create_call_kwargs(_JSONBodyVideoConfig()))
result = BaseLLMHTTPHandler().video_generation_handler(
client=client, **_video_create_call_kwargs(_JSONBodyVideoConfig())
)
assert captured["content_type"] == "application/json"
assert json.loads(captured["body"]) == {"model": "sora-2", "prompt": "a cat surfing", "seconds": "4"}
@ -3336,6 +3353,7 @@ def test_video_generation_with_input_reference_keeps_file_multipart():
AZURE_AI_BASE = "https://myfoundry.services.ai.azure.com"
AZURE_AI_CHAT_COMPLETIONS_URL = f"{AZURE_AI_BASE}/models/chat/completions"
def _a_tool_with_an_unsupported_field() -> dict:
return {
"type": "function",
@ -3343,14 +3361,13 @@ def _a_tool_with_an_unsupported_field() -> dict:
"strict": True,
}
A_COMPLETION = {
"id": "chatcmpl-1",
"object": "chat.completion",
"created": 1,
"model": "grok-3",
"choices": [
{"index": 0, "message": {"role": "assistant", "content": "sent"}, "finish_reason": "stop"}
],
"choices": [{"index": 0, "message": {"role": "assistant", "content": "sent"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2},
}
@ -3394,9 +3411,7 @@ def _call_azure_ai(recorder: _RecordedAzureAI, **overrides):
def test_a_tool_field_the_provider_rejects_is_dropped_and_the_call_retried():
recorder = _RecordedAzureAI(
[_rejection(TOOL_LEVEL_REJECTION), httpx.Response(200, json=A_COMPLETION)]
)
recorder = _RecordedAzureAI([_rejection(TOOL_LEVEL_REJECTION), httpx.Response(200, json=A_COMPLETION)])
response = _call_azure_ai(recorder)
@ -3407,9 +3422,7 @@ def test_a_tool_field_the_provider_rejects_is_dropped_and_the_call_retried():
def test_the_retry_changes_only_the_field_the_provider_named():
recorder = _RecordedAzureAI(
[_rejection(TOOL_LEVEL_REJECTION), httpx.Response(200, json=A_COMPLETION)]
)
recorder = _RecordedAzureAI([_rejection(TOOL_LEVEL_REJECTION), httpx.Response(200, json=A_COMPLETION)])
_call_azure_ai(recorder)
@ -3448,9 +3461,7 @@ def test_an_extra_input_outside_a_tool_is_not_retried_unless_dropping_params_was
def test_an_extra_input_outside_a_tool_is_retried_when_dropping_params_was_asked_for():
recorder = _RecordedAzureAI(
[_rejection(UNRELATED_REJECTION), httpx.Response(200, json=A_COMPLETION)]
)
recorder = _RecordedAzureAI([_rejection(UNRELATED_REJECTION), httpx.Response(200, json=A_COMPLETION)])
response = _call_azure_ai(recorder, drop_params=True)
@ -3464,9 +3475,7 @@ async def test_a_tool_field_the_provider_rejects_is_dropped_and_retried_on_the_a
):
import respx
recorder = _RecordedAzureAI(
[_rejection(TOOL_LEVEL_REJECTION), httpx.Response(200, json=A_COMPLETION)]
)
recorder = _RecordedAzureAI([_rejection(TOOL_LEVEL_REJECTION), httpx.Response(200, json=A_COMPLETION)])
with respx.mock(assert_all_called=True) as router:
router.post(AZURE_AI_CHAT_COMPLETIONS_URL).mock(side_effect=recorder)
@ -3671,7 +3680,9 @@ def _start_async_completion(config, logging_obj=None):
custom_llm_provider="openai",
model_response=ModelResponse(),
encoding=None,
logging_obj=logging_obj if logging_obj is not None else Mock(dynamic_success_callbacks=None, model_call_details={}),
logging_obj=logging_obj
if logging_obj is not None
else Mock(dynamic_success_callbacks=None, model_call_details={}),
optional_params={},
timeout=10.0,
litellm_params={},
@ -4042,7 +4053,7 @@ def test_image_edit_handler_does_not_bill_references_the_transform_dropped():
assert response.reference_pixels == 0
def test_image_edit_handler_leaves_reference_pixels_unset_when_unmeasurable():
def test_image_edit_handler_bills_unmeasurable_reference_as_one_megapixel():
client = HTTPHandler()
client.client = httpx.Client(transport=_fixed_json_transport())
@ -4059,7 +4070,7 @@ def test_image_edit_handler_leaves_reference_pixels_unset_when_unmeasurable():
client=client,
)
assert response.reference_pixels is None
assert response.reference_pixels == 1024 * 1024
class _ScriptedClientWebSocket(_FakeClientWebSocket):
@ -4174,7 +4185,9 @@ async def test_async_realtime_bridges_a_transcription_session_through_the_provid
logging_obj.dispatch_failure_handlers = AsyncMock()
handler = BaseLLMHTTPHandler()
with patch.object(handler, "_open_realtime_backend_ws", AsyncMock(side_effect=AssertionError("dialed a websocket"))) as dial:
with patch.object(
handler, "_open_realtime_backend_ws", AsyncMock(side_effect=AssertionError("dialed a websocket"))
) as dial:
await handler.async_realtime(
model="chirp_3",
websocket=client_ws,