feat(images): add bounded header parsers that measure reference image pixels

This commit is contained in:
shrey kharbanda 2026-09-24 00:19:44 +00:00
parent 6c8afb221f
commit 298ca046b4
2 changed files with 435 additions and 0 deletions

View file

@ -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("<HH", head[26:30]))
return ImageDimensions(width=width & 0x3FFF, height=height & 0x3FFF)
case b"VP8L":
bits: Final = int.from_bytes(head[21:25], "little")
return ImageDimensions(width=(bits & 0x3FFF) + 1, height=((bits >> 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)

View file

@ -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("<I", len(body)) + body
return b"RIFF" + struct.pack("<I", len(chunk) + 4) + b"WEBP" + chunk
JPEG_APP0: Final = 0xE0
JPEG_APP2_MARKER: Final = 0xE2
JPEG_DHT: Final = 0xC4
JPEG_JPG_MARKER: Final = 0xC8
JPEG_DAC_MARKER: Final = 0xCC
JPEG_SOF0_MARKER: Final = 0xC0
JPEG_SOF2_MARKER: Final = 0xC2
PADDED_JPEG: Final = JPEG_SOI + b"\xff\xff" + _jpeg_sof_segment(JPEG_SOF0_MARKER, 320, 200)
SHORT_LENGTH_JPEG: Final = JPEG_SOI + bytes((0xFF, JPEG_APP0, 0x00, 0x01))
TRUNCATED_SOF_JPEG: Final = JPEG_SOI + bytes((0xFF, JPEG_SOF0_MARKER, 0x00, 0x11, 0x08))
SEGMENTS_63_JPEG: Final = (
JPEG_SOI + _jpeg_segment(JPEG_APP0, bytes(4)) * 63 + _jpeg_sof_segment(JPEG_SOF0_MARKER, 320, 200)
)
SEGMENTS_65_JPEG: Final = (
JPEG_SOI + _jpeg_segment(JPEG_APP0, bytes(4)) * 65 + _jpeg_sof_segment(JPEG_SOF0_MARKER, 320, 200)
)
SEGMENTS_300_JPEG: Final = (
JPEG_SOI + _jpeg_segment(JPEG_APP2_MARKER, bytes(4)) * 300 + _jpeg_sof_segment(JPEG_SOF0_MARKER, 320, 200)
)
SEGMENTS_1030_JPEG: Final = (
JPEG_SOI + _jpeg_segment(JPEG_APP0, bytes(4)) * 1030 + _jpeg_sof_segment(JPEG_SOF0_MARKER, 320, 200)
)
FILL_BYTES_JPEG: Final = JPEG_SOI + b"\xff" * 1100 + _jpeg_sof_segment(JPEG_SOF0_MARKER, 320, 200)
FILL_BEFORE_LARGE_SEGMENT_JPEG: Final = (
JPEG_SOI + b"\xff" + _jpeg_segment(JPEG_APP2_MARKER, bytes(0xFEFE)) + _jpeg_sof_segment(JPEG_SOF0_MARKER, 640, 480)
)
FILL_200K_JPEG: Final = JPEG_SOI + b"\xff" * 200_000 + _jpeg_sof_segment(JPEG_SOF0_MARKER, 320, 200)
FILL_65535_JPEG: Final = JPEG_SOI + b"\xff" * 65535 + _jpeg_sof_segment(JPEG_SOF0_MARKER, 320, 200)
FILL_65536_JPEG: Final = JPEG_SOI + b"\xff" * 65536 + _jpeg_sof_segment(JPEG_SOF0_MARKER, 320, 200)
OVERSIZED_HEADER_JPEG: Final = (
JPEG_SOI + _jpeg_segment(JPEG_APP0, bytes(65533)) * 260 + _jpeg_sof_segment(JPEG_SOF0_MARKER, 320, 200)
)
PROGRESSIVE_JPEG: Final = JPEG_SOI + _jpeg_sof_segment(JPEG_SOF2_MARKER, 111, 55)
DHT_JPEG: Final = JPEG_SOI + _jpeg_segment(JPEG_DHT, bytes(8)) + _jpeg_sof_segment(JPEG_SOF0_MARKER, 400, 300)
VP8_BODY: Final = bytes(3) + b"\x9d\x01\x2a" + struct.pack("<HH", 0xC000 | 500, 0xC000 | 250) + bytes(4)
VP8L_BITS: Final = 299 | (199 << 14)
VP8L_BODY: Final = b"\x2f" + VP8L_BITS.to_bytes(4, "little") + bytes(5)
VP8X_BODY: Final = bytes(4) + (700 - 1).to_bytes(3, "little") + (350 - 1).to_bytes(3, "little")
def _gif(version: bytes, width: int, height: int) -> bytes:
return version + struct.pack("<HH", width, height) + bytes(22)
def _bmp(width: int, height: int) -> bytes:
return b"BM" + bytes(12) + struct.pack("<I", 40) + struct.pack("<ii", width, height) + bytes(6)
def _dims(width: int, height: int) -> 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("<I", 12) + struct.pack("<HH", 640, 480) + bytes(4),
_dims(640, 480),
id="bmp-coreheader",
),
pytest.param(b"II*\x00" + bytes(28), None, id="unknown-format"),
),
)
def test_read_image_dimensions_parses_hand_built_headers(header: bytes, expected: ImageDimensions | None):
assert read_image_dimensions(header) == expected
def test_read_image_dimensions_returns_none_and_restores_position_on_read_error():
stream: Final = _FailingReadStream(bytes(64))
stream.seek(5)
assert read_image_dimensions(stream) is None
assert stream.tell() == 5
@pytest.mark.parametrize("initial_position", (0, 3))
def test_read_image_dimensions_reads_png_jpeg_webp_headers_and_keeps_position(initial_position: int):
pillow_jpeg: Final = _pillow_bytes(1024, 768, "RGB", "JPEG")
large_app1_jpeg: Final = _jpeg_with_leading_segments(pillow_jpeg, _jpeg_app_segment(JPEG_APP1, 60_000))
far_sof_jpeg: Final = _jpeg_with_leading_segments(
pillow_jpeg, _jpeg_app_segment(JPEG_APP2, 40_000), _jpeg_app_segment(JPEG_APP2, 40_000)
)
assert far_sof_jpeg.index(JPEG_SOF0) > 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