From 298ca046b4ab6edc50cebca866177eb165004934 Mon Sep 17 00:00:00 2001 From: shrey kharbanda Date: Thu, 24 Sep 2026 00:19:44 +0000 Subject: [PATCH] feat(images): add bounded header parsers that measure reference image pixels --- litellm/images/dimensions.py | 165 ++++++++++++ tests/test_litellm/images/test_dimensions.py | 270 +++++++++++++++++++ 2 files changed, 435 insertions(+) create mode 100644 litellm/images/dimensions.py create mode 100644 tests/test_litellm/images/test_dimensions.py diff --git a/litellm/images/dimensions.py b/litellm/images/dimensions.py new file mode 100644 index 00000000000..979e777ae81 --- /dev/null +++ b/litellm/images/dimensions.py @@ -0,0 +1,165 @@ +import os +import struct +from collections.abc import Sequence +from dataclasses import dataclass +from io import BytesIO +from typing import IO, Final, cast + +from litellm._logging import verbose_logger +from litellm.litellm_core_utils.token_counter import get_image_type +from litellm.types.llms.openai import FileTypes + +_JPEG_SOF_MARKERS: Final = frozenset(range(0xC0, 0xD0)) - {0xC4, 0xC8, 0xCC} +_MAX_JPEG_SEGMENTS: Final = 1024 +_MAX_JPEG_HEADER_OFFSET: Final = 16 * 1024 * 1024 +_JPEG_FIRST_SEGMENT_OFFSET: Final = 2 +_JPEG_SOF_PAYLOAD_SIZE: Final = 5 +_JPEG_FILL_CHUNK: Final = 64 * 1024 +_HEADER_READ_SIZE: Final = 32 +_MIN_PNG_HEADER_SIZE: Final = 24 +_MIN_GIF_HEADER_SIZE: Final = 10 +_MIN_BMP_HEADER_SIZE: Final = 26 +_MIN_WEBP_HEADER_SIZE: Final = 30 + + +@dataclass(frozen=True, slots=True) +class ImageDimensions: + width: int + height: int + + @property + def pixels(self) -> int: + return self.width * self.height + + +def _content_stream(image: FileTypes) -> IO[bytes] | None: + content: Final = image[1] if isinstance(image, tuple) else image + if isinstance(content, bytes): + return BytesIO(content) + if isinstance(content, (str, os.PathLike)): + return None + return content if content.seekable() else None + + +def _png_dimensions(head: bytes) -> ImageDimensions | None: + if len(head) < _MIN_PNG_HEADER_SIZE: + return None + width, height = cast(tuple[int, int], struct.unpack(">II", head[16:24])) + return ImageDimensions(width=width, height=height) + + +def _gif_dimensions(head: bytes) -> ImageDimensions | None: + if len(head) < _MIN_GIF_HEADER_SIZE: + return None + width: Final = int.from_bytes(head[6:8], "little") + height: Final = int.from_bytes(head[8:10], "little") + return ImageDimensions(width=width, height=height) + + +def _bmp_dimensions(head: bytes) -> ImageDimensions | None: + if len(head) < _MIN_BMP_HEADER_SIZE: + return None + if int.from_bytes(head[14:18], "little") == 12: + core_width: Final = int.from_bytes(head[18:20], "little") + core_height: Final = int.from_bytes(head[20:22], "little") + return ImageDimensions(width=core_width, height=core_height) + width: Final = int.from_bytes(head[18:22], "little", signed=True) + height: Final = int.from_bytes(head[22:26], "little", signed=True) + return ImageDimensions(width=width, height=abs(height)) + + +def _webp_dimensions(head: bytes) -> ImageDimensions | None: + if len(head) < _MIN_WEBP_HEADER_SIZE: + return None + match head[12:16]: # pyright: ignore[reportMatchNotExhaustive] # unknown fourccs fall through to the None below + case b"VP8X": + return ImageDimensions( + width=int.from_bytes(head[24:27], "little") + 1, + height=int.from_bytes(head[27:30], "little") + 1, + ) + case b"VP8 ": + width, height = cast(tuple[int, int], struct.unpack("> 14) & 0x3FFF) + 1) + return None + + +def _jpeg_sof_dimensions(stream: IO[bytes], start: int, offset: int, segments_left: int) -> ImageDimensions | None: + segment_offset = offset # rebind-ok: advanced one JPEG segment per hop + for _ in range(segments_left): + if segment_offset > _MAX_JPEG_HEADER_OFFSET: + return None + stream.seek(start + segment_offset) + marker = stream.read(4) + if len(marker) < 4 or marker[0] != 0xFF: + return None + if marker[1] == 0xFF: + stream.seek(start + segment_offset) + fill = stream.read(_JPEG_FILL_CHUNK) + run = len(fill) - len(fill.lstrip(b"\xff")) + segment_offset += run - 1 + continue + if marker[1] in _JPEG_SOF_MARKERS: + 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)) + return ImageDimensions(width=width, height=height) + segment_length = int.from_bytes(marker[2:4], "big") + if segment_length < 2: + return None + segment_offset += 2 + segment_length + return None + + +def _header_dimensions(stream: IO[bytes], position: int) -> ImageDimensions | None: + stream.seek(position) + head: Final = stream.read(_HEADER_READ_SIZE) + match get_image_type(head): # pyright: ignore[reportMatchNotExhaustive] # unmatched types fall through to the None below + case "png": + return _png_dimensions(head) + case "webp": + return _webp_dimensions(head) + case "jpeg": + return _jpeg_sof_dimensions(stream, position, _JPEG_FIRST_SEGMENT_OFFSET, _MAX_JPEG_SEGMENTS) + case "gif": + return _gif_dimensions(head) + if head[:2] == b"BM": + return _bmp_dimensions(head) + return None + + +def _measure_stream(stream: IO[bytes]) -> ImageDimensions | None: + position: Final = stream.tell() + try: + dimensions: Final = _header_dimensions(stream, position) + finally: + stream.seek(position) + if dimensions is None or dimensions.width <= 0 or dimensions.height <= 0: + return None + return dimensions + + +def read_image_dimensions(image: FileTypes) -> ImageDimensions | None: + try: + stream: Final = _content_stream(image) + if stream is None: + return None + return _measure_stream(stream) + except (OSError, ValueError, AttributeError, struct.error): + return None + + +def total_reference_pixels(images: Sequence[FileTypes]) -> int | None: + measured: Final = tuple(read_image_dimensions(image) for image in images) + dimensions: Final = tuple(size for size in measured if size is not None) + if len(dimensions) != len(measured): + 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), + ) + return None + return sum(size.pixels for size in dimensions) diff --git a/tests/test_litellm/images/test_dimensions.py b/tests/test_litellm/images/test_dimensions.py new file mode 100644 index 00000000000..a5866c6c45d --- /dev/null +++ b/tests/test_litellm/images/test_dimensions.py @@ -0,0 +1,270 @@ +import io +import struct +import zlib +from pathlib import Path +from typing import Any, Final, cast + +import pytest +from PIL import Image + +from litellm.images.dimensions import ImageDimensions, read_image_dimensions, total_reference_pixels + +CAT_JPEG: Final = Path(__file__).parents[2] / "e2e" / "llm_translation" / "fixtures" / "cat.jpg" +JPEG_SOI: Final = b"\xff\xd8" +JPEG_APP1: Final = b"\xff\xe1" +JPEG_APP2: Final = b"\xff\xe2" +JPEG_SOF0: Final = b"\xff\xc0" +JPEG_MARKER_AND_LENGTH_SIZE: Final = 4 + + +def _png_chunk(tag: bytes, payload: bytes) -> bytes: + return struct.pack(">I", len(payload)) + tag + payload + struct.pack(">I", zlib.crc32(tag + payload)) + + +def _png(width: int, height: int) -> bytes: + header: Final = struct.pack(">IIBBBBB", width, height, 8, 0, 0, 0, 0) + rows: Final = b"".join(b"\x00" + bytes(width) for _ in range(height)) + return ( + b"\x89PNG\r\n\x1a\n" + + _png_chunk(b"IHDR", header) + + _png_chunk(b"IDAT", zlib.compress(rows)) + + _png_chunk(b"IEND", b"") + ) + + +def _pillow_bytes(width: int, height: int, mode: str, image_format: str, **save_options: object) -> bytes: + buffer: Final = io.BytesIO() + Image.new(mode, (width, height)).save(buffer, format=image_format, **save_options) + return buffer.getvalue() + + +def _jpeg_app_segment(marker: bytes, payload_size: int) -> bytes: + return marker + struct.pack(">H", payload_size + 2) + bytes(payload_size) + + +def _jpeg_with_leading_segments(jpeg: bytes, *segments: bytes) -> bytes: + return JPEG_SOI + b"".join(segments) + jpeg[len(JPEG_SOI) :] + + +class _NonSeekableStream(io.BytesIO): + def seekable(self) -> bool: + return False + + +class _FailingReadStream(io.BytesIO): + def read(self, size: int = -1) -> bytes: + raise OSError("simulated disk read failure") + + +def _jpeg_sof_segment(marker: int, width: int, height: int) -> bytes: + return bytes((0xFF, marker)) + struct.pack(">H", 17) + struct.pack(">BHH", 8, height, width) + + +def _jpeg_segment(marker: int, payload: bytes) -> bytes: + return bytes((0xFF, marker)) + struct.pack(">H", len(payload) + 2) + payload + + +def _webp_chunk(fourcc: bytes, body: bytes) -> bytes: + chunk: Final = fourcc + struct.pack(" bytes: + return version + struct.pack(" bytes: + return b"BM" + bytes(12) + struct.pack(" ImageDimensions: + return ImageDimensions(width=width, height=height) + + +@pytest.mark.parametrize( + ("header", "expected"), + ( + pytest.param(PADDED_JPEG, _dims(320, 200), id="jpeg-sof-after-0xff-padding"), + pytest.param(SHORT_LENGTH_JPEG, None, id="jpeg-segment-length-below-2"), + pytest.param(TRUNCATED_SOF_JPEG, None, id="jpeg-truncated-sof-payload"), + pytest.param(SEGMENTS_63_JPEG, _dims(320, 200), id="jpeg-sof-at-segment-64-boundary"), + pytest.param(SEGMENTS_65_JPEG, _dims(320, 200), id="jpeg-65-segments-within-budget"), + pytest.param(SEGMENTS_300_JPEG, _dims(320, 200), id="jpeg-300-app2-segments"), + pytest.param(SEGMENTS_1030_JPEG, None, id="jpeg-more-than-1024-segments"), + pytest.param(FILL_BYTES_JPEG, _dims(320, 200), id="jpeg-1100-fill-bytes-before-sof"), + pytest.param(FILL_BEFORE_LARGE_SEGMENT_JPEG, _dims(640, 480), id="jpeg-fill-before-large-segment"), + pytest.param(FILL_200K_JPEG, _dims(320, 200), id="jpeg-200k-fill-bytes-before-sof"), + pytest.param(FILL_65535_JPEG, _dims(320, 200), id="jpeg-65535-fill-bytes-before-sof"), + pytest.param(FILL_65536_JPEG, _dims(320, 200), id="jpeg-65536-fill-bytes-before-sof"), + pytest.param(OVERSIZED_HEADER_JPEG, None, id="jpeg-sof-beyond-16mib"), + pytest.param(PROGRESSIVE_JPEG, _dims(111, 55), id="jpeg-sof2-progressive"), + pytest.param(DHT_JPEG, _dims(400, 300), id="jpeg-dht-skipped-before-sof0"), + pytest.param(_webp_chunk(b"VP8 ", VP8_BODY), _dims(500, 250), id="webp-vp8-lossy"), + pytest.param(_webp_chunk(b"VP8L", VP8L_BODY), _dims(300, 200), id="webp-vp8l-lossless"), + pytest.param(_webp_chunk(b"VP8X", VP8X_BODY), _dims(700, 350), id="webp-vp8x-extended"), + pytest.param(_webp_chunk(b"VP8Z", bytes(10)), None, id="webp-unknown-fourcc"), + pytest.param(_webp_chunk(b"VP8X", VP8X_BODY)[:29], None, id="webp-shorter-than-30-bytes"), + pytest.param(b"\x89PNG\r\n\x1a\n" + bytes(10), None, id="png-shorter-than-24-bytes"), + pytest.param(_gif(b"GIF89a", 1024, 768), _dims(1024, 768), id="gif89a"), + pytest.param(_gif(b"GIF87a", 320, 200), _dims(320, 200), id="gif87a"), + pytest.param(b"GIF89a" + bytes(2), None, id="gif-shorter-than-10-bytes"), + pytest.param(_bmp(1024, 768), _dims(1024, 768), id="bmp-positive-height"), + pytest.param(_bmp(1024, -768), _dims(1024, 768), id="bmp-negative-height"), + pytest.param(_bmp(-1024, 768), None, id="bmp-negative-width"), + pytest.param(_bmp(0, 768), None, id="bmp-zero-width"), + pytest.param(b"BM" + bytes(18), None, id="bmp-shorter-than-26-bytes"), + pytest.param( + b"BM" + bytes(12) + struct.pack(" 70_000 + cases: Final = ( + (_png(1024, 1024), (1024, 1024)), + (CAT_JPEG.read_bytes(), (512, 512)), + (_pillow_bytes(640, 480, "RGB", "WEBP"), (640, 480)), + (_pillow_bytes(640, 480, "RGB", "WEBP", lossless=True), (640, 480)), + (_pillow_bytes(640, 480, "RGBA", "WEBP"), (640, 480)), + (large_app1_jpeg, (1024, 768)), + (far_sof_jpeg, (1024, 768)), + ) + layouts: Final = tuple(image[12:16] for image, _ in cases[2:5]) + assert layouts == (b"VP8 ", b"VP8L", b"VP8X") + + for image, expected in cases: + stream: Final = io.BytesIO(bytes(initial_position) + image) + stream.seek(initial_position) + assert read_image_dimensions(stream) == _dims(*expected), image[:16] + assert stream.tell() == initial_position, image[:16] + + +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) + prefixes: Final = (32, 160, sof_offset, sof_offset + JPEG_MARKER_AND_LENGTH_SIZE) + + for prefix in prefixes: + stream: Final = io.BytesIO(jpeg[:prefix]) + assert read_image_dimensions(stream) is None, prefix + assert stream.tell() == 0, prefix + + +def test_read_image_dimensions_returns_none_for_unreadable_bytes_paths_tuples_and_non_seekable_streams( + tmp_path: Path, +): + png: Final = _png(64, 64) + path: Final = tmp_path / "ref.png" + path.write_bytes(png) + non_seekable: Final = _NonSeekableStream(png) + + assert read_image_dimensions(io.BytesIO(b"image")) is None + assert read_image_dimensions(b"II*\x00" + bytes(26)) is None + assert read_image_dimensions(path) is None + assert read_image_dimensions(("ref.png", path, "image/png", {})) is None + assert read_image_dimensions(non_seekable) is None + assert non_seekable.tell() == 0 + assert read_image_dimensions(("ref.png", png)) == _dims(64, 64) + assert read_image_dimensions(("ref.png", io.BytesIO(png), "image/png")) == _dims(64, 64) + + +@pytest.mark.parametrize("sof_marker", (0xC0, 0xC1, 0xC2, 0xC3, 0xC5, 0xC6, 0xC7, 0xC9, 0xCA, 0xCB, 0xCD, 0xCE, 0xCF)) +def test_read_image_dimensions_reads_every_jpeg_sof_marker(sof_marker: int): + assert read_image_dimensions(JPEG_SOI + _jpeg_sof_segment(sof_marker, 320, 200)) == _dims(320, 200) + + +@pytest.mark.parametrize("skipped_marker", (JPEG_JPG_MARKER, JPEG_DAC_MARKER)) +def test_read_image_dimensions_skips_non_sof_c_range_markers(skipped_marker: int): + jpeg: Final = JPEG_SOI + _jpeg_segment(skipped_marker, bytes(8)) + _jpeg_sof_segment(JPEG_SOF0_MARKER, 400, 300) + + assert read_image_dimensions(jpeg) == _dims(400, 300) + + +class _ReadOnlyObject: + def read(self, size: int = -1) -> bytes: + return b"" + + +def test_total_reference_pixels_returns_none_for_closed_and_read_only_streams(): + 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 + + +def test_total_reference_pixels_returns_none_when_a_reference_has_invalid_dimensions(): + 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 + + +def test_total_reference_pixels_returns_none_when_any_reference_is_unmeasurable(): + 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([]) == 0