mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
refactor(images): expose reference pixels through a typed ImageResponse accessor
This commit is contained in:
parent
f62437f13c
commit
eedf4d8648
6 changed files with 47 additions and 27 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue