refactor(images): expose reference pixels through a typed ImageResponse accessor

This commit is contained in:
shrey kharbanda 2026-09-24 00:23:13 +00:00
parent f62437f13c
commit eedf4d8648
6 changed files with 47 additions and 27 deletions

View file

@ -105,7 +105,7 @@ def _jpeg_sof_dimensions(stream: IO[bytes], start: int, offset: int, segments_le
sof = stream.read(_JPEG_SOF_PAYLOAD_SIZE)
if len(sof) < _JPEG_SOF_PAYLOAD_SIZE:
return None
_precision, height, width = cast(tuple[int, int, int], struct.unpack(">BHH", sof))
height, width = cast(tuple[int, int], struct.unpack(">HH", sof[1:5]))
return ImageDimensions(width=width, height=height)
segment_length = int.from_bytes(marker[2:4], "big")
if segment_length < 2:

View file

@ -1,3 +1,4 @@
import re
from collections.abc import Mapping
from typing import Any, Final
@ -14,23 +15,16 @@ def _rate(table: ModelInfo, key: str) -> float | None:
return _get_cost_per_unit(table, key, default_value=None)
_SIZE_PATTERN: Final = re.compile(r"(\d+)(?:x|-x-)(\d+)")
def _size_pixels(size: str | None) -> int:
if size is None:
return 0
for separator in ("x", "-x-"):
if separator in size:
parts = size.split(separator)
if len(parts) != 2:
continue
try:
width = int(parts[0])
height = int(parts[1])
except ValueError:
continue
if width > 0 and height > 0:
return width * height
continue
return 0
match: Final = _SIZE_PATTERN.fullmatch(size)
if match is None or int(match[1]) <= 0 or int(match[2]) <= 0:
return 0
return int(match[1]) * int(match[2])
def _generated_pixels(
@ -73,8 +67,7 @@ def cost_calculator(
if usage_cost is not None:
return usage_cost
data: Final = image_response.data
num_images: Final = n if n is not None else (len(data) if isinstance(data, list) else 0)
num_images: Final = n if n is not None else len(image_response.data or ())
generated_meters: Final = (
("output_cost_per_image", num_images),
("input_cost_per_pixel", _generated_pixels(optional_params, size, image_response) * num_images),
@ -83,8 +76,7 @@ def cost_calculator(
(rate * units for key, units in generated_meters if (rate := _rate(resolved, key)) is not None),
0.0,
)
reference_pixels: Final = getattr(image_response, "_reference_pixels", None)
reference: Final = (_rate(resolved, "input_cost_per_reference_pixel") or 0.0) * (
reference_pixels if isinstance(reference_pixels, int) and not isinstance(reference_pixels, bool) else 0
image_response.reference_pixels or 0
)
return generated + reference

View file

@ -385,8 +385,7 @@ def _collect_ws_project_quota_callbacks() -> tuple[ProjectQuotaCallback, ...]:
def _with_reference_pixels(response: ImageResponse, reference_pixels: int | None) -> ImageResponse:
if reference_pixels is not None:
# mutable-ok: handler stamps this billing field
response._reference_pixels = reference_pixels # pyright: ignore[reportPrivateUsage] # billing stamp
response.set_reference_pixels(reference_pixels)
return response

View file

@ -2586,6 +2586,13 @@ class ImageResponse(OpenAIImageResponse, BaseLiteLLMOpenAIResponseObject):
_hidden_params: dict = {}
_reference_pixels: int | None = None
@property
def reference_pixels(self) -> int | None:
return self._reference_pixels
def set_reference_pixels(self, pixels: int) -> None:
self._reference_pixels = pixels
usage: ImageUsage | None = None
"""
Users might use litellm with older python versions, we don't want this to break for them.

View file

@ -24,16 +24,17 @@ def use_local_model_cost_map(monkeypatch: pytest.MonkeyPatch):
def _edit_response(reference_pixels: int | None = REFERENCE_PIXELS, **kwargs: object) -> ImageResponse:
response: Final = ImageResponse(data=[ImageObject(b64_json="aW1n")], size="1024x1024", **kwargs)
response._reference_pixels = reference_pixels
if reference_pixels is not None:
response.set_reference_pixels(reference_pixels)
return response
def _edit_cost(response: ImageResponse, **kwargs: object) -> float:
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="1024x1024",
size=size,
call_type="image_edit",
**kwargs,
)
@ -136,6 +137,26 @@ def test_edit_without_measurement_bills_generated_pixels_only() -> 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),

View file

@ -3927,7 +3927,8 @@ def test_image_edit_handler_stamps_measured_reference_pixels():
client=client,
)
assert response._reference_pixels == 4 * 2 + 1 * 1
assert response.reference_pixels == 4 * 2 + 1 * 1
assert "reference_pixels" not in response.model_dump()
async def test_async_image_edit_handler_stamps_measured_reference_pixels():
@ -3947,7 +3948,7 @@ async def test_async_image_edit_handler_stamps_measured_reference_pixels():
client=client,
)
assert response._reference_pixels == 8
assert response.reference_pixels == 8
def test_image_edit_handler_leaves_reference_pixels_unset_when_unmeasurable():
@ -3967,7 +3968,7 @@ def test_image_edit_handler_leaves_reference_pixels_unset_when_unmeasurable():
client=client,
)
assert response._reference_pixels is None
assert response.reference_pixels is None
class _ScriptedClientWebSocket(_FakeClientWebSocket):