mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(core): stop get_image_dimensions crashing on bare base64 and unfetchable URLs
get_image_dimensions documents 'URL or base64 encoded string' input, but the base64 branch unconditionally split on ',' and unpacked two values, so bare base64 (no data: prefix) raised ValueError, and a URL whose download failed fell into the same branch and crashed the same way — instead of reaching the documented 'sensible default dimensions' fallback that only the unknown-image-format path could reach. The base64 branch now accepts both data-URL and bare base64 forms and treats a decode error as 'use defaults'; an unfetchable URL or undecodable payload leaves img_data unset and skips image-format detection, falling through to DEFAULT_IMAGE_WIDTH/HEIGHT.
This commit is contained in:
parent
973329e986
commit
e38b01c973
2 changed files with 59 additions and 5 deletions
|
|
@ -213,12 +213,19 @@ def get_image_dimensions(
|
|||
img_data = body
|
||||
except Exception:
|
||||
pass
|
||||
if img_data is None:
|
||||
# Not a URL or fetch failed — assume base64
|
||||
_header, encoded = data.split(",", 1)
|
||||
img_data = base64.b64decode(encoded)
|
||||
if img_data is None and not data.startswith(("http://", "https://")):
|
||||
# Not a URL — assume base64. Accept both data-URL form
|
||||
# ('data:<mime>;base64,<payload>') and bare base64; on a decode
|
||||
# error, leave img_data unset so the default dimensions are used.
|
||||
encoded = data.partition(",")[2] if data.startswith("data:") else data
|
||||
try:
|
||||
img_data = base64.b64decode(encoded)
|
||||
except ValueError:
|
||||
verbose_logger.debug("Failed to decode base64 image data; using default dimensions")
|
||||
|
||||
img_type: Final = get_image_type(img_data)
|
||||
# A URL that could not be fetched (or base64 that could not be decoded)
|
||||
# leaves img_data unset — fall through to the default dimensions below.
|
||||
img_type: Final = get_image_type(img_data) if img_data is not None else None
|
||||
|
||||
if img_type == "png":
|
||||
w, h = struct.unpack(">LL", img_data[16:24])
|
||||
|
|
|
|||
|
|
@ -526,6 +526,53 @@ def test_img_url_token_counter(img_url, monkeypatch):
|
|||
assert height is not None
|
||||
|
||||
|
||||
def test_get_image_dimensions_bare_base64():
|
||||
"""
|
||||
Bare base64 (no 'data:' URL prefix) is documented input for
|
||||
get_image_dimensions and must not crash on the ','-split of a data URL.
|
||||
"""
|
||||
import base64
|
||||
import struct
|
||||
|
||||
from litellm.litellm_core_utils.token_counter import get_image_dimensions
|
||||
|
||||
# Minimal PNG header carrying 100x200 dimensions.
|
||||
png = base64.b64encode(
|
||||
b"\x89PNG\r\n\x1a\n" + b"\x00" * 8 + struct.pack(">LL", 100, 200) + b"\x00" * 16
|
||||
)
|
||||
assert get_image_dimensions(data=png.decode()) == (100, 200)
|
||||
|
||||
|
||||
def test_get_image_dimensions_unfetchable_url_returns_defaults(monkeypatch):
|
||||
"""A URL that cannot be fetched falls back to the default dimensions instead of raising."""
|
||||
|
||||
def _raise(client, url, **kwargs):
|
||||
raise RuntimeError("connection failed")
|
||||
|
||||
monkeypatch.setattr(
|
||||
"litellm.litellm_core_utils.token_counter.safe_get",
|
||||
_raise,
|
||||
)
|
||||
from litellm.constants import DEFAULT_IMAGE_HEIGHT, DEFAULT_IMAGE_WIDTH
|
||||
from litellm.litellm_core_utils.token_counter import get_image_dimensions
|
||||
|
||||
assert get_image_dimensions(data="https://invalid.invalid/x.png") == (
|
||||
DEFAULT_IMAGE_WIDTH,
|
||||
DEFAULT_IMAGE_HEIGHT,
|
||||
)
|
||||
|
||||
|
||||
def test_get_image_dimensions_undecodable_data_returns_defaults():
|
||||
"""Data that is neither a fetchable URL nor decodable base64 uses the defaults."""
|
||||
from litellm.constants import DEFAULT_IMAGE_HEIGHT, DEFAULT_IMAGE_WIDTH
|
||||
from litellm.litellm_core_utils.token_counter import get_image_dimensions
|
||||
|
||||
assert get_image_dimensions(data="not-an-image!") == (
|
||||
DEFAULT_IMAGE_WIDTH,
|
||||
DEFAULT_IMAGE_HEIGHT,
|
||||
)
|
||||
|
||||
|
||||
def test_token_encode_disallowed_special():
|
||||
encode(model="gpt-3.5-turbo", text="Hello, world! <|endoftext|>")
|
||||
token_counter(model="gpt-3.5-turbo", text="Hello, world! <|endoftext|>")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue