fix(images): measure the uploaded reference set and resolve deployment pricing kwargs

- measure every file part and every embedded image field the request
  carries instead of the caller's image arg, so a transform that drops,
  filters, or adds parts and a JSON body that smuggles extra_body
  references all bill what the provider actually receives
- read seekable streams from offset 0 and restore the cursor, matching
  how httpx serializes the part
- measure input_image* base64 fields on azure image generation the
  same as edits
- merge top-level litellm_params pricing kwargs over nested model_info
  in _deployment_model_info so custom pricing works on non-router calls
- reject bool and non-positive width/height in the size fallback and
  clamp negative reference_pixels to zero
- prices_tokens now checks rates for non-zero values; get_model_info
  normalizes absent token rates to 0, which made is not None always
  true and zero-billed per-pixel models that return usage
This commit is contained in:
shrey kharbanda 2026-09-24 06:22:28 +00:00
parent 793371270b
commit d0730b51d7
11 changed files with 649 additions and 65 deletions

View file

@ -28,6 +28,7 @@ from litellm.litellm_core_utils.llm_cost_calc.usage_object_transformation import
from litellm.litellm_core_utils.llm_cost_calc.utils import (
BilledTokenRates,
CostCalculatorUtils,
_DEPLOYMENT_PRICING_KEYS, # pyright: ignore[reportPrivateUsage] # declared-pricing key set shared with sibling calculators
_generic_cost_per_character,
_get_regional_uplift_multiplier,
_get_service_tier_cost_key,
@ -2020,7 +2021,12 @@ def _deployment_model_info(
litellm_params: Final = getattr(litellm_logging_obj, "litellm_params", None)
if litellm_params is None:
return None
return next(
declared: Final = {
key: value
for key in _DEPLOYMENT_PRICING_KEYS
if (value := litellm_params.get(key)) is not None
}
nested: Final = next(
(
model_info
for metadata_key in ("metadata", "litellm_metadata")
@ -2028,6 +2034,8 @@ def _deployment_model_info(
),
None,
)
merged: Final = {**(nested or {}), **declared}
return cast(ModelInfo, merged) if merged else None
def _ocr_model_info(

View file

@ -1,15 +1,23 @@
import base64
import os
import struct
from collections.abc import Sequence
from collections.abc import Iterator, Mapping, Sequence
from dataclasses import dataclass
from io import BytesIO
from typing import IO, Final, cast
from httpx._types import RequestFiles # pyright: ignore[reportPrivateImportUsage] # same source the base transform classes use
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.token_counter import get_image_type
from litellm.types.llms.openai import FileTypes
from litellm.types.utils import ImageResponse
_HEADER_READ_SIZE: Final = 32
_IMAGE_SIGNATURE_BYTES: Final = 12
_DATA_URI_PREFIX: Final = "data:"
_EMBEDDED_IMAGE_MAX_DEPTH: Final = 4
_IMAGE_SIGNATURES: Final = frozenset({"png", "jpeg", "webp", "gif"})
_JPEG_SOF_MARKERS: Final = frozenset(range(0xC0, 0xD0)) - {0xC4, 0xC8, 0xCC}
_JPEG_MAX_SEGMENTS: Final = 1024
@ -52,29 +60,121 @@ def total_reference_pixels(images: Sequence[FileTypes]) -> int | None:
return sum(size.pixels for size in dimensions)
def uploaded_reference_pixels(
files: RequestFiles | None,
json_body: Mapping[str, object] | None,
) -> int | None:
"""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.
"""
parts: Final = _file_parts(files) + tuple(_embedded_image_values(json_body or {}))
if not parts:
return 0
return total_reference_pixels(parts)
def _file_parts(files: RequestFiles | None) -> tuple[FileTypes, ...]:
if files is None:
return ()
if isinstance(files, Mapping):
return tuple(files.values())
return tuple(files)
def _file_content(part: FileTypes) -> object:
"""The payload bytes of an httpx file part: (field, (name, content, type, headers?)) unwraps twice."""
unwrapped: Final = part[1] if isinstance(part, tuple) else part
return unwrapped[1] if isinstance(unwrapped, tuple) else unwrapped
def read_image_dimensions(image: FileTypes) -> ImageDimensions | None:
"""Measure one image file's pixel dimensions; ``None`` when unmeasurable, never raises.
Seekable content is measured from offset 0 (the bytes an upload sends) with the stream position
restored afterward. Strings are read as embedded base64 or data-URI image content; a filesystem
path does not decode to an image signature and still returns ``None``.
"""
try:
content: Final = cast(
"IO[bytes] | bytes | str | os.PathLike[str] | None",
image[1] if isinstance(image, tuple) else image,
_file_content(image),
)
if content is None or isinstance(content, (str, os.PathLike)):
if content is None:
return None
if isinstance(content, (str, os.PathLike)):
embedded: Final = _decoded_image_bytes(os.fspath(content))
if embedded is None:
return None
return _dimensions_from_stream(BytesIO(embedded))
stream = BytesIO(content) if isinstance(content, (bytes, bytearray, memoryview)) else content
if not stream.seekable():
return None
position = stream.tell()
try:
dimensions = _header_dimensions(stream, position)
finally:
stream.seek(position)
if dimensions is None or dimensions.width <= 0 or dimensions.height <= 0:
return None
return dimensions
return _dimensions_from_stream(stream)
except Exception: # noqa: BLE001 # an odd stream or malformed file must fall back, never raise
return None
def with_reference_pixels(response: ImageResponse, reference_pixels: int | None) -> ImageResponse:
if reference_pixels is not None:
response.set_reference_pixels(reference_pixels)
return response
def _dimensions_from_stream(stream: IO[bytes]) -> ImageDimensions | None:
if not stream.seekable():
return None
position: Final = stream.tell()
try:
dimensions = _header_dimensions(stream, 0)
finally:
stream.seek(position)
if dimensions is None or dimensions.width <= 0 or dimensions.height <= 0:
return None
return dimensions
def _decoded_image_bytes(text: str) -> bytes | None:
"""The image bytes a string carries, when it is base64 or data-URI encoded image content."""
payload: Final = text.partition(",")[2] if text.startswith(_DATA_URI_PREFIX) else text
if len(payload) < _IMAGE_SIGNATURE_BYTES:
return None
try:
head: Final = base64.b64decode(payload[:64])
if not _is_image_signature(head):
return None
return base64.b64decode(payload)
except ValueError: # binascii.Error subclasses ValueError; undecodable fields are not images
return None
def _embedded_image_values(value: object, depth: int = 0) -> Iterator[bytes]:
if depth > _EMBEDDED_IMAGE_MAX_DEPTH:
return
if isinstance(value, str):
decoded: Final = _decoded_image_bytes(value)
if decoded is not None:
yield decoded
elif isinstance(value, (bytes, bytearray, memoryview)):
raw: Final = bytes(value)
if _is_image_signature(raw[:_HEADER_READ_SIZE]):
yield raw
elif isinstance(value, Mapping):
for item in value.values():
yield from _embedded_image_values(item, depth + 1)
elif isinstance(value, (list, tuple)):
for item in value:
yield from _embedded_image_values(item, depth + 1)
def _is_image_signature(head: bytes) -> bool:
return len(head) >= _IMAGE_SIGNATURE_BYTES and (
get_image_type(head) in _IMAGE_SIGNATURES or head[:2] == b"BM"
)
def _header_dimensions(stream: IO[bytes], position: int) -> ImageDimensions | None:
stream.seek(position)
head: Final = cast("bytes | None", stream.read(_HEADER_READ_SIZE))

View file

@ -854,8 +854,9 @@ def deployment_pricing(model_info: ModelInfo | None) -> ModelInfo | None:
def prices_tokens(model_info: ModelInfo) -> bool:
"""Whether the price table carries any token rate, so a token-priced calculator can bill from usage."""
return any(model_info.get(key) is not None for key in _IMAGE_TOKEN_RATE_KEYS)
"""Whether the price table carries a billable token rate, so a token-priced calculator can bill
from usage. get_model_info fills absent rates with 0, which cannot bill anything."""
return any(bool(model_info.get(key)) for key in _IMAGE_TOKEN_RATE_KEYS)
def flat_image_cost(model_info: ModelInfo | None, image_response: ImageResponse) -> float:

View file

@ -16,6 +16,7 @@ from openai import (
import litellm
from litellm.constants import AZURE_OPERATION_POLLING_TIMEOUT, DEFAULT_MAX_RETRIES
from litellm.images.dimensions import uploaded_reference_pixels, with_reference_pixels
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.litellm_core_utils.url_utils import SSRFError, assert_same_origin
@ -1185,18 +1186,22 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
headers=headers,
deployment_name=model,
)
reference_pixels: Final = uploaded_reference_pixels(None, data)
provider_config: Final = get_azure_image_generation_config(data.get("model", "dall-e-2"))
if provider_config is not None:
return provider_config.transform_image_generation_response(
model=data.get("model", "dall-e-2"),
raw_response=httpx_response,
model_response=model_response or ImageResponse(),
logging_obj=logging_obj,
request_data=data,
optional_params=data,
litellm_params=data,
encoding=litellm.encoding,
return with_reference_pixels(
provider_config.transform_image_generation_response(
model=data.get("model", "dall-e-2"),
raw_response=httpx_response,
model_response=model_response or ImageResponse(),
logging_obj=logging_obj,
request_data=data,
optional_params=data,
litellm_params=data,
encoding=litellm.encoding,
),
reference_pixels,
)
else:
@ -1210,10 +1215,13 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
additional_args={"complete_input_dict": data},
original_response=stringified_response,
)
return convert_to_model_response_object(
response_object=stringified_response,
model_response_object=model_response,
response_type="image_generation",
return with_reference_pixels(
convert_to_model_response_object(
response_object=stringified_response,
model_response_object=model_response,
response_type="image_generation",
),
reference_pixels,
)
except Exception as e:
## LOGGING
@ -1265,6 +1273,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
max_retries: Final = data.pop("max_retries", 2)
if not isinstance(max_retries, int):
raise AzureOpenAIError(status_code=422, message="max retries must be an int")
reference_pixels: Final = uploaded_reference_pixels(None, data)
auth_params: Final[dict[str, object]] = {**(litellm_params or {})} # mutable-ok: SDK init takes a dict
if azure_ad_token is not None:
@ -1324,15 +1333,18 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
)
provider_config: Final = get_azure_image_generation_config(data.get("model", "dall-e-2"))
if isinstance(provider_config, AzureFoundryMAIImageGenerationConfig):
return provider_config.transform_image_generation_response(
model=data.get("model", "dall-e-2"),
raw_response=httpx_response,
model_response=model_response or ImageResponse(),
logging_obj=logging_obj,
request_data=data,
optional_params=data,
litellm_params=data,
encoding=litellm.encoding,
return with_reference_pixels(
provider_config.transform_image_generation_response(
model=data.get("model", "dall-e-2"),
raw_response=httpx_response,
model_response=model_response or ImageResponse(),
logging_obj=logging_obj,
request_data=data,
optional_params=data,
litellm_params=data,
encoding=litellm.encoding,
),
reference_pixels,
)
response: Final = httpx_response.json()
@ -1345,10 +1357,13 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
original_response=response,
)
# return response
return convert_to_model_response_object(
response_object=response,
model_response_object=model_response,
response_type="image_generation",
return with_reference_pixels(
convert_to_model_response_object(
response_object=response,
model_response_object=model_response,
response_type="image_generation",
),
reference_pixels,
)
except AzureOpenAIError as e:
raise e

View file

@ -20,7 +20,7 @@ def _rate(table: ModelInfo, key: str) -> float | None:
def _generated_pixels(optional_params: Mapping[str, object] | None, size: str | None) -> int:
width: Final = optional_params.get("width") if optional_params else None
height: Final = optional_params.get("height") if optional_params else None
if isinstance(width, int) and isinstance(height, int):
if type(width) is int and type(height) is int and width > 0 and height > 0:
return width * height
match: Final = _SIZE_PATTERN.fullmatch(size or "")
return int(match[1]) * int(match[2]) if match else 0
@ -58,7 +58,7 @@ def cost_calculator(
num_images: Final = n if n is not None else len(image_response.data or ())
generated_cost: Final = _generated_cost(resolved, num_images, _generated_pixels(optional_params, size))
per_pixel: Final = _rate(resolved, "input_cost_per_pixel") or 0.0
reference_cost: Final = per_pixel * (image_response.reference_pixels or 0)
reference_cost: Final = per_pixel * max(image_response.reference_pixels or 0, 0)
return generated_cost + reference_cost

View file

@ -33,7 +33,7 @@ from litellm._logging import _redact_string, verbose_logger
from litellm.anthropic_beta_headers_manager import update_headers_with_filtered_beta
from litellm.constants import MAX_FILE_LIST_LIMIT, REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES
from litellm.files.types import FileContentStreamingResult
from litellm.images.dimensions import total_reference_pixels
from litellm.images.dimensions import uploaded_reference_pixels, with_reference_pixels
from litellm.litellm_core_utils.agentic_followup_kwargs import build_agentic_followup_kwargs
from litellm.litellm_core_utils.agentic_loop_settings import (
DEFAULT_MAX_AGENTIC_LOOPS,
@ -383,12 +383,6 @@ 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:
response.set_reference_pixels(reference_pixels)
return response
class BaseLLMHTTPHandler:
async def _make_common_async_call(
self,
@ -6888,7 +6882,6 @@ class BaseLLMHTTPHandler:
litellm_params=dict(litellm_params),
)
reference_pixels: Final = total_reference_pixels(image if isinstance(image, list) else [image])
data, files = image_edit_provider_config.transform_image_edit_request(
model=model,
image=image,
@ -6898,6 +6891,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)
## LOGGING
logging_obj.pre_call(
@ -6936,7 +6930,7 @@ class BaseLLMHTTPHandler:
provider_config=image_edit_provider_config,
)
return _with_reference_pixels(
return with_reference_pixels(
image_edit_provider_config.transform_image_edit_response(
model=model,
raw_response=response,
@ -6991,7 +6985,6 @@ class BaseLLMHTTPHandler:
litellm_params=dict(litellm_params),
)
reference_pixels: Final = total_reference_pixels(image if isinstance(image, list) else [image])
data, files = await image_edit_provider_config.async_transform_image_edit_request(
model=model,
image=image,
@ -7001,6 +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)
## LOGGING
logging_obj.pre_call(
@ -7039,7 +7033,7 @@ class BaseLLMHTTPHandler:
provider_config=image_edit_provider_config,
)
return _with_reference_pixels(
return with_reference_pixels(
image_edit_provider_config.transform_image_edit_response(
model=model,
raw_response=response,
@ -7121,6 +7115,7 @@ class BaseLLMHTTPHandler:
litellm_params=dict(litellm_params),
headers=headers,
)
reference_pixels: Final = uploaded_reference_pixels(None, data)
## LOGGING
logging_obj.pre_call(
@ -7170,7 +7165,7 @@ class BaseLLMHTTPHandler:
encoding=None,
)
return model_response
return with_reference_pixels(model_response, reference_pixels)
async def async_image_generation_handler(
self,
@ -7228,6 +7223,7 @@ class BaseLLMHTTPHandler:
litellm_params=dict(litellm_params),
headers=headers,
)
reference_pixels: Final = uploaded_reference_pixels(None, data)
## LOGGING
logging_obj.pre_call(
@ -7277,7 +7273,7 @@ class BaseLLMHTTPHandler:
encoding=None,
)
return model_response
return with_reference_pixels(model_response, reference_pixels)
###### VIDEO GENERATION HANDLER ######
def video_generation_handler(

View file

@ -1,3 +1,4 @@
import base64
import io
import struct
import zlib
@ -7,7 +8,12 @@ from typing import Any, Final, cast
import pytest
from PIL import Image
from litellm.images.dimensions import ImageDimensions, read_image_dimensions, total_reference_pixels
from litellm.images.dimensions import (
ImageDimensions,
read_image_dimensions,
total_reference_pixels,
uploaded_reference_pixels,
)
CAT_JPEG: Final = Path(__file__).parents[2] / "e2e" / "llm_translation" / "fixtures" / "cat.jpg"
JPEG_SOI: Final = b"\xff\xd8"
@ -165,6 +171,11 @@ def _dims(width: int, height: int) -> ImageDimensions:
_dims(640, 480),
id="bmp-coreheader",
),
pytest.param(
b"BM" + bytes(12) + struct.pack("<I", 108) + struct.pack("<ii", 640, 480) + bytes(6),
_dims(640, 480),
id="bmp-v4header",
),
pytest.param(b"II*\x00" + bytes(28), None, id="unknown-format"),
),
)
@ -201,12 +212,43 @@ def test_read_image_dimensions_reads_png_jpeg_webp_headers_and_keeps_position(in
assert layouts == (b"VP8 ", b"VP8L", b"VP8X")
for image, expected in cases:
stream: Final = io.BytesIO(bytes(initial_position) + image)
stream: Final = io.BytesIO(image)
stream.seek(initial_position)
# uploads rewind seekable parts to offset 0, so a mid-file cursor still measures the whole image
assert read_image_dimensions(stream) == _dims(*expected), image[:16]
assert stream.tell() == initial_position, image[:16]
def test_read_image_dimensions_returns_none_when_stream_content_is_not_an_image():
stream: Final = io.BytesIO(bytes(5) + _png(1024, 1024))
stream.seek(5)
assert read_image_dimensions(stream) is None
assert stream.tell() == 5
def test_read_image_dimensions_reads_base64_and_data_uri_strings():
png_b64: Final = base64.b64encode(_png(512, 256)).decode()
assert read_image_dimensions(png_b64) == _dims(512, 256)
assert read_image_dimensions(f"data:image/png;base64,{png_b64}") == _dims(512, 256)
assert read_image_dimensions(("ref.png", png_b64)) == _dims(512, 256)
@pytest.mark.parametrize(
"text",
(
"/tmp/nonexistent.png",
"",
"the quick brown fox",
base64.b64encode(b"not an image body").decode(),
"data:text/plain;base64," + base64.b64encode(b"still not an image").decode(),
),
)
def test_read_image_dimensions_returns_none_for_non_image_strings(text: str):
assert read_image_dimensions(text) is None
def test_read_image_dimensions_returns_none_on_truncated_jpeg_and_keeps_position():
jpeg: Final = _pillow_bytes(1024, 768, "RGB", "JPEG")
sof_offset: Final = jpeg.index(JPEG_SOF0)
@ -292,3 +334,40 @@ def test_total_reference_pixels_returns_none_when_any_reference_is_unmeasurable(
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([]) == 0
def test_uploaded_reference_pixels_measures_the_uploaded_set_not_the_requested_set():
png_b64: Final = base64.b64encode(_png(64, 64)).decode()
jpeg_b64: Final = base64.b64encode(CAT_JPEG.read_bytes()).decode()
# 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
# 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
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(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():
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

View file

@ -14,7 +14,7 @@ from litellm.llms.azure.image_generation.http_utils import azure_deployment_imag
from litellm.llms.azure_ai.image_generation.flux_transformation import (
AzureFoundryFluxImageGenerationConfig,
)
from litellm.types.utils import ImageObject, ImageResponse
from litellm.types.utils import ImageObject, ImageResponse, ImageUsage
from litellm.utils import _invalidate_model_cost_lowercase_map, get_optional_params_image_gen
@ -255,3 +255,173 @@ def test_flux2_response_preserves_mapped_dimensions():
encoding=None,
)
assert response.size == "2048x1024"
def _catalog_pixel_rate() -> float:
return litellm.model_cost["azure_ai/FLUX.2-flex"]["input_cost_per_pixel"]
_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.set_reference_pixels(_ONE_MEGAPIXEL)
cost: Final = CostCalculatorUtils.route_image_generation_cost_calculator(
model="FLUX.2-flex",
completion_response=response,
custom_llm_provider="azure_ai",
size="1024x1024",
call_type="image_edit",
n=2,
)
assert cost == pytest.approx(_catalog_pixel_rate() * (_ONE_MEGAPIXEL * 2 + _ONE_MEGAPIXEL))
@pytest.mark.parametrize(
"dimensions",
(
{"width": True, "height": 1024},
{"width": -2048, "height": 1024},
{"width": 2048, "height": 0},
{"width": 2048.0, "height": 1024},
),
)
def test_flux2_cost_rejects_bool_and_non_positive_dimensions_for_the_size_string(
dimensions: Mapping[str, int | float | bool],
) -> None:
response: Final = ImageResponse(data=[ImageObject(b64_json="aW1n")])
cost: Final = CostCalculatorUtils.route_image_generation_cost_calculator(
model="FLUX.2-flex",
completion_response=response,
custom_llm_provider="azure_ai",
optional_params=dimensions,
size="2048x1024",
call_type="image_edit",
n=1,
)
assert cost == pytest.approx(_catalog_pixel_rate() * 2048 * 1024)
@pytest.mark.parametrize("size", ("auto", "big", "1024", "1024x"))
def test_flux2_cost_skips_unparseable_size_strings(size: str) -> None:
response: Final = ImageResponse(data=[ImageObject(b64_json="aW1n")])
response.set_reference_pixels(_ONE_MEGAPIXEL)
cost: Final = CostCalculatorUtils.route_image_generation_cost_calculator(
model="FLUX.2-flex",
completion_response=response,
custom_llm_provider="azure_ai",
size=size,
call_type="image_edit",
)
assert cost == pytest.approx(_catalog_pixel_rate() * _ONE_MEGAPIXEL)
def test_flux2_cost_honors_an_explicit_zero_output_cost_per_image() -> None:
response: Final = ImageResponse(data=[ImageObject(b64_json="aW1n")])
response.set_reference_pixels(_ONE_MEGAPIXEL)
cost: Final = CostCalculatorUtils.route_image_generation_cost_calculator(
model="FLUX.2-flex",
completion_response=response,
custom_llm_provider="azure_ai",
size="1024x1024",
call_type="image_edit",
model_info={"output_cost_per_image": 0.0, "input_cost_per_pixel": _catalog_pixel_rate()},
)
assert cost == pytest.approx(_catalog_pixel_rate() * _ONE_MEGAPIXEL)
def test_flux2_cost_ignores_negative_reference_pixels() -> None:
response: Final = ImageResponse(data=[ImageObject(b64_json="aW1n")])
response.set_reference_pixels(-2_097_152)
cost: Final = CostCalculatorUtils.route_image_generation_cost_calculator(
model="FLUX.2-flex",
completion_response=response,
custom_llm_provider="azure_ai",
size="1024x1024",
call_type="image_edit",
)
assert cost == pytest.approx(_catalog_pixel_rate() * _ONE_MEGAPIXEL)
def test_flux2_cost_bills_pixels_when_usage_carries_no_token_rates() -> None:
response: Final = ImageResponse(
data=[ImageObject(b64_json="aW1n")],
usage=ImageUsage(
input_tokens=100,
input_tokens_details={"image_tokens": 50, "text_tokens": 50},
output_tokens=50,
total_tokens=150,
),
)
response.set_reference_pixels(_ONE_MEGAPIXEL)
cost: Final = CostCalculatorUtils.route_image_generation_cost_calculator(
model="FLUX.2-flex",
completion_response=response,
custom_llm_provider="azure_ai",
size="1024x1024",
call_type="image_edit",
)
assert cost == pytest.approx(_catalog_pixel_rate() * _ONE_MEGAPIXEL * 2)
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)
logging_obj: Final = MagicMock()
logging_obj.litellm_params = {"input_cost_per_pixel": deployment_rate}
cost: Final = litellm.completion_cost(
completion_response=response,
model="azure_ai/FLUX.2-flex",
custom_llm_provider="azure_ai",
custom_pricing=True,
litellm_logging_obj=logging_obj,
size="1024x1024",
call_type="image_edit",
)
assert cost == pytest.approx(deployment_rate * _ONE_MEGAPIXEL * 2)
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")])
response.set_reference_pixels(_ONE_MEGAPIXEL)
logging_obj: Final = MagicMock()
logging_obj.litellm_params = {
"metadata": {"model_info": {"input_cost_per_pixel": nested_rate}},
"input_cost_per_pixel": kwargs_rate,
}
cost: Final = litellm.completion_cost(
completion_response=response,
model="azure_ai/FLUX.2-flex",
custom_llm_provider="azure_ai",
custom_pricing=True,
litellm_logging_obj=logging_obj,
size="1024x1024",
call_type="image_edit",
)
assert cost == pytest.approx(kwargs_rate * _ONE_MEGAPIXEL * 2)

View file

@ -3851,6 +3851,57 @@ def _echo_json_transport(captured):
return httpx.MockTransport(handle)
def _fixed_json_transport():
return httpx.MockTransport(lambda request: httpx.Response(200, json={"transformed_by": "sync"}))
class _ImageEditForwardingConfig(_ImageEditRecordingConfig):
"""Uploads the caller's images inside the JSON body, like the FLUX.2 transform."""
def _forwarded_images(self, image):
images = image if isinstance(image, list) else [image]
return [
base64.b64encode(item if isinstance(item, bytes) else item[1]).decode()
for item in images
if item is not None
]
def transform_image_edit_request(
self, model, prompt, image, image_edit_optional_request_params, litellm_params, headers
):
self.transform_calls.append("sync")
return {"transformed_by": "sync", "image": self._forwarded_images(image)}, []
async def async_transform_image_edit_request(
self, model, prompt, image, image_edit_optional_request_params, litellm_params, headers
):
self.transform_calls.append("async")
return {"transformed_by": "async", "image": self._forwarded_images(image)}, []
class _ImageEditMultipartConfig(_ImageEditRecordingConfig):
"""Uploads the caller's images as multipart file parts, like the MAI transform."""
def use_multipart_form_data(self):
return True
def _file_parts_for(self, image):
images = image if isinstance(image, list) else [image]
return [("image", ("image.png", item, "image/png")) for item in images if item is not None]
def transform_image_edit_request(
self, model, prompt, image, image_edit_optional_request_params, litellm_params, headers
):
self.transform_calls.append("sync")
return {"transformed_by": "sync"}, self._file_parts_for(image)
async def async_transform_image_edit_request(
self, model, prompt, image, image_edit_optional_request_params, litellm_params, headers
):
self.transform_calls.append("async")
return {"transformed_by": "async"}, self._file_parts_for(image)
async def test_async_image_edit_handler_awaits_the_async_transform():
config = _ImageEditRecordingConfig()
captured = {}
@ -3918,7 +3969,7 @@ def test_image_edit_handler_stamps_measured_reference_pixels():
model="edit-model",
image=[_tiny_png(4, 2), _tiny_png(1, 1)],
prompt="add a hat",
image_edit_provider_config=_ImageEditRecordingConfig(),
image_edit_provider_config=_ImageEditForwardingConfig(),
image_edit_optional_request_params={},
custom_llm_provider="openai",
litellm_params=GenericLiteLLMParams(),
@ -3939,7 +3990,7 @@ async def test_async_image_edit_handler_stamps_measured_reference_pixels():
model="edit-model",
image=_tiny_png(4, 2),
prompt="add a hat",
image_edit_provider_config=_ImageEditRecordingConfig(),
image_edit_provider_config=_ImageEditForwardingConfig(),
image_edit_optional_request_params={},
custom_llm_provider="openai",
litellm_params=GenericLiteLLMParams(),
@ -3951,15 +4002,55 @@ async def test_async_image_edit_handler_stamps_measured_reference_pixels():
assert response.reference_pixels == 8
def test_image_edit_handler_bills_every_multipart_reference_it_uploads():
client = HTTPHandler()
client.client = httpx.Client(transport=_fixed_json_transport())
response = BaseLLMHTTPHandler().image_edit_handler(
model="edit-model",
image=[_tiny_png(4, 2), _tiny_png(1, 1)],
prompt="add a hat",
image_edit_provider_config=_ImageEditMultipartConfig(),
image_edit_optional_request_params={},
custom_llm_provider="openai",
litellm_params=GenericLiteLLMParams(),
logging_obj=Mock(),
timeout=10.0,
client=client,
)
assert response.reference_pixels == 4 * 2 + 1 * 1
def test_image_edit_handler_does_not_bill_references_the_transform_dropped():
client = HTTPHandler()
client.client = httpx.Client(transport=_fixed_json_transport())
response = BaseLLMHTTPHandler().image_edit_handler(
model="edit-model",
image=[_tiny_png(4, 2), _tiny_png(1, 1)],
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,
client=client,
)
assert response.reference_pixels == 0
def test_image_edit_handler_leaves_reference_pixels_unset_when_unmeasurable():
client = HTTPHandler()
client.client = httpx.Client(transport=_echo_json_transport({}))
client.client = httpx.Client(transport=_fixed_json_transport())
response = BaseLLMHTTPHandler().image_edit_handler(
model="edit-model",
image=b"not-an-image",
prompt="add a hat",
image_edit_provider_config=_ImageEditRecordingConfig(),
image_edit_provider_config=_ImageEditMultipartConfig(),
image_edit_optional_request_params={},
custom_llm_provider="openai",
litellm_params=GenericLiteLLMParams(),

View file

@ -1,4 +1,6 @@
import base64
import json
import struct
import traceback
from typing import Callable, Optional
from unittest.mock import AsyncMock, MagicMock, Mock, patch
@ -378,6 +380,43 @@ def test_azure_image_generation_v1_api_version_uses_v1_route(api_version):
assert azure_deployment_image_generation_json_body(url, data) == data
_PNG_1024X1024: bytes = (
b"\x89PNG\r\n\x1a\n"
+ struct.pack(">I", 13)
+ b"IHDR"
+ struct.pack(">IIBBBBB", 1024, 1024, 8, 2, 0, 0, 0)
)
def test_azure_ai_image_generation_meters_extra_body_reference_images():
"""extra_body fields flattened into the JSON body are metered: base64 reference images a
caller adds through extra_body reach the provider, so they are billed like edit references."""
png_base64: str = base64.b64encode(_PNG_1024X1024).decode()
mock_http_response = MagicMock()
mock_http_response.status_code = 200
mock_http_response.json.return_value = {"data": [{"b64_json": "aW1n"}]}
with patch.object(HTTPHandler, "post", return_value=mock_http_response):
response = AzureChatCompletion().image_generation(
prompt="Blend the references",
timeout=60.0,
optional_params={
"n": 1,
"size": "1024x1024",
"extra_body": {"input_image": png_base64, "input_image_2": png_base64},
},
logging_obj=MagicMock(),
headers={},
model="FLUX.2-flex",
api_key="test-api-key",
api_base="https://example.services.ai.azure.com",
api_version="preview",
litellm_params={},
)
assert response.reference_pixels == 2 * 1024 * 1024
def test_azure_image_generation_dated_api_version_uses_deployment_route():
url = AzureChatCompletion().create_azure_base_url(
azure_client_params={

View file

@ -1,5 +1,7 @@
import base64
import json
import struct
import zlib
from collections.abc import Mapping
from typing import Final
@ -175,7 +177,90 @@ def test_flux2_image_edit_preserves_controls_and_pixel_cost(dimensions: Mapping[
**dimensions,
)
assert response._hidden_params["response_cost"] == pytest.approx(5e-08 * 2048 * 1024 * 2)
catalog_rate: Final = litellm.get_model_info(model="azure_ai/FLUX.2-flex", custom_llm_provider="azure_ai")[
"input_cost_per_pixel"
]
# the reference b"image" decodes to non-image content and is not metered; generated pixels only
assert response._hidden_params["response_cost"] == pytest.approx(catalog_rate * 2048 * 1024 * 2)
def _png_bytes(width: int, height: int) -> bytes:
"""Smallest well-formed PNG carrying real IHDR dimensions."""
def chunk(tag: bytes, payload: bytes) -> bytes:
return (
struct.pack(">I", len(payload))
+ tag
+ payload
+ struct.pack(">I", zlib.crc32(tag + payload))
)
return (
b"\x89PNG\r\n\x1a\n"
+ chunk(b"IHDR", struct.pack(">IIBBBBB", width, height, 8, 0, 0, 0, 0))
+ chunk(b"IDAT", zlib.compress(b"\x00"))
+ chunk(b"IEND", b"")
)
@pytest.mark.usefixtures("local_model_cost_map")
def test_flux2_image_edit_bills_every_uploaded_reference():
"""Each input_image field the transform builds and posts is metered at input_cost_per_pixel."""
catalog_rate: Final = litellm.get_model_info(model="azure_ai/FLUX.2-flex", custom_llm_provider="azure_ai")[
"input_cost_per_pixel"
]
client: Final = HTTPHandler(
client=httpx.Client(
transport=httpx.MockTransport(
lambda request: httpx.Response(200, json={"data": [{"b64_json": "aW1n"}]})
)
)
)
response: Final = litellm.image_edit(
model="azure_ai/FLUX.2-flex",
image=[_png_bytes(1024, 1024), _png_bytes(1024, 1024)],
prompt="Blend the references",
api_key="test-key",
api_base="https://example.services.ai.azure.com",
client=client,
size="1024x1024",
)
assert response.reference_pixels == 2 * 1024 * 1024
assert response._hidden_params["response_cost"] == pytest.approx(catalog_rate * 3 * 1024 * 1024)
@pytest.mark.usefixtures("local_model_cost_map")
def test_mai_image_edit_bills_only_the_reference_it_uploads():
"""MAI-Image uploads only the first caller image, so billing must follow the uploaded set,
not the requested set: two passed images bill one uploaded reference."""
deployment_rate: Final = 5e-08
client: Final = HTTPHandler(
client=httpx.Client(
transport=httpx.MockTransport(
lambda request: httpx.Response(200, json={"data": [{"b64_json": "aW1n"}]})
)
)
)
response: Final = litellm.image_edit(
model="azure_ai/MAI-Image-2.5",
image=[_png_bytes(1024, 1024), _png_bytes(1024, 1024)],
prompt="Add a hat",
api_key="test-key",
api_base="https://example.services.ai.azure.com",
client=client,
input_cost_per_pixel=deployment_rate,
)
catalog_price_per_image: Final = litellm.get_model_info(
model="azure_ai/MAI-Image-2.5", custom_llm_provider="azure_ai"
)["output_cost_per_image"]
assert response.reference_pixels == 1024 * 1024
assert response._hidden_params["response_cost"] == pytest.approx(
catalog_price_per_image + deployment_rate * 1024 * 1024
)
def test_flux2_image_edit_accepts_and_drops_openai_only_parameters():