Return 400 instead of 500 for unfetchable or unparseable multimodal inputs (#31417)

Co-authored-by: hnyls2002 <lsyincs@gmail.com>
Co-authored-by: Liangsheng Yin <hnyls2002@gmail.com>
This commit is contained in:
Richard Gong
2026-07-28 05:23:42 -07:00
committed by GitHub
co-authored by hnyls2002 Liangsheng Yin
parent 161fffedfc
commit 32c30c0f96
4 changed files with 138 additions and 8 deletions
@@ -69,6 +69,7 @@ from sglang.srt.server_args import (
set_global_server_args_for_scheduler, set_global_server_args_for_scheduler,
) )
from sglang.srt.utils import ( from sglang.srt.utils import (
CLIENT_MEDIA_EXCEPTIONS,
add_prometheus_middleware, add_prometheus_middleware,
configure_logger, configure_logger,
load_audio, load_audio,
@@ -624,6 +625,9 @@ class MMEncoder:
elif modality == Modality.AUDIO: elif modality == Modality.AUDIO:
return load_audio(data, self.model_audio_sr) return load_audio(data, self.model_audio_sr)
except CLIENT_MEDIA_EXCEPTIONS as e:
# Not ValueError: the DP envelope classifies by `.code`, which only MMError carries.
raise BadRequestError(f"Error while loading data {data}: {e}") from e
except Exception as e: except Exception as e:
raise RuntimeError(f"Error while loading data {data}: {e}") raise RuntimeError(f"Error while loading data {data}: {e}")
@@ -22,6 +22,7 @@ from sglang.srt.managers.schedule_batch import (
from sglang.srt.multimodal.processors.executor import MultimodalProcessorExecutor from sglang.srt.multimodal.processors.executor import MultimodalProcessorExecutor
from sglang.srt.runtime_context import get_server_args from sglang.srt.runtime_context import get_server_args
from sglang.srt.utils import ( from sglang.srt.utils import (
CLIENT_MEDIA_EXCEPTIONS,
envs, envs,
is_cpu, is_cpu,
is_npu, is_npu,
@@ -652,8 +653,7 @@ class BaseMultimodalProcessor(ABC):
elif modality == Modality.AUDIO: elif modality == Modality.AUDIO:
return load_audio(data, audio_sample_rate) return load_audio(data, audio_sample_rate)
except ValueError as e: except CLIENT_MEDIA_EXCEPTIONS as e:
# Bad input (e.g. invalid base64) -> 400, not 500.
data_str = str(data) data_str = str(data)
if len(data_str) > 100: if len(data_str) > 100:
data_str = data_str[:100] + "..." data_str = data_str[:100] + "..."
+19 -1
View File
@@ -86,7 +86,7 @@ import torch
import torch.distributed as dist import torch.distributed as dist
import triton import triton
from packaging import version as pkg_version from packaging import version as pkg_version
from PIL import Image from PIL import Image, UnidentifiedImageError
from starlette.routing import Mount from starlette.routing import Mount
from torch import nn from torch import nn
from torch.library import Library from torch.library import Library
@@ -1514,6 +1514,15 @@ def get_mm_http_session() -> requests.Session:
return session return session
# Raised by the loaders below when client-supplied media cannot be fetched or
# decoded. ValueError is in the set because invalid base64 raises binascii.Error.
CLIENT_MEDIA_EXCEPTIONS = (
ValueError,
UnidentifiedImageError,
requests.exceptions.RequestException,
)
def load_audio( def load_audio(
audio_file: str, sr: Optional[int] = None, mono: bool = True audio_file: str, sr: Optional[int] = None, mono: bool = True
) -> np.ndarray: ) -> np.ndarray:
@@ -1582,10 +1591,13 @@ def load_audio(
import torch import torch
import torchaudio import torchaudio
try:
if isinstance(source, bytes): if isinstance(source, bytes):
audio, original_sr = sf.read(BytesIO(source)) audio, original_sr = sf.read(BytesIO(source))
else: else:
audio, original_sr = sf.read(source) audio, original_sr = sf.read(source)
except sf.LibsndfileError as e:
raise ValueError(f"Could not decode audio: {e}") from e
if mono and len(audio.shape) > 1: if mono and len(audio.shape) > 1:
audio = np.mean(audio, axis=1) audio = np.mean(audio, axis=1)
@@ -1785,7 +1797,13 @@ def load_video(video_file: Union[str, bytes, VideoData], use_gpu: bool = True):
raise ValueError(f"Unsupported video input type: {type(video_file)}") raise ValueError(f"Unsupported video input type: {type(video_file)}")
device = "cuda" if use_gpu else "cpu" device = "cuda" if use_gpu else "cpu"
try:
return VideoDecoderWrapper(source, device=device) return VideoDecoderWrapper(source, device=device)
except (ImportError, MemoryError):
raise # missing backend / OOM is not a bad payload
except Exception as e:
# Broad on purpose: torchcodec raises RuntimeError, decord its own type.
raise ValueError(f"Could not decode video: {e}") from e
def sample_video_frames(video, *, desired_fps: int, max_frames: int) -> list[int]: def sample_video_frames(video, *, desired_fps: int, max_frames: int) -> list[int]:
@@ -0,0 +1,108 @@
"""Unit tests for bad-input classification in ``BaseMultimodalProcessor._load_single_item``.
Media the client supplied but that cannot be fetched or decoded must raise
``ValueError``; anything else must stay a ``RuntimeError``.
"""
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=10, suite="base-a-test-cpu")
import binascii
import unittest
from unittest.mock import MagicMock, patch
import requests
from sglang.srt.managers.schedule_batch import Modality
from sglang.srt.multimodal.processors.base_processor import BaseMultimodalProcessor
from sglang.srt.utils.common import CLIENT_MEDIA_EXCEPTIONS
from sglang.test.test_utils import CustomTestCase
MODALITIES = (Modality.IMAGE, Modality.AUDIO, Modality.VIDEO)
class _StubProcessor(BaseMultimodalProcessor):
# gpu_image_decode=False keeps the decode on PIL so the test needs no GPU. The
# abstract methods are never called: only the _load_single_item classmethod is.
gpu_image_decode = False
def _session_raising(exc):
session = MagicMock()
session.get.side_effect = exc
return session
class TestBadInputIsClientError(CustomTestCase):
def _assert_client_error(self, data, modality):
with self.assertRaises(ValueError):
_StubProcessor._load_single_item(data, modality)
def test_unfetchable_url_every_modality(self):
# All three loaders fetch through get_mm_http_session(); HTTPError,
# ConnectionError and Timeout all subclass RequestException.
for exc in (
requests.exceptions.HTTPError("404 from media host"),
requests.exceptions.ConnectionError("dns failure"),
requests.exceptions.Timeout("read timed out"),
):
for modality in MODALITIES:
with self.subTest(exc=type(exc).__name__, modality=modality):
with patch(
"sglang.srt.utils.common.get_mm_http_session",
return_value=_session_raising(exc),
):
self._assert_client_error("https://media.host/clip", modality)
def test_invalid_base64(self):
self._assert_client_error("!!!not-base64!!!", Modality.IMAGE)
def test_undecodable_image_bytes(self):
# PIL raises UnidentifiedImageError, an OSError -- not a ValueError.
self._assert_client_error(b"definitely not an image", Modality.IMAGE)
def test_undecodable_audio_bytes(self):
# soundfile raises LibsndfileError, a RuntimeError -- not a ValueError.
self._assert_client_error(b"definitely not audio", Modality.AUDIO)
def test_undecodable_video_bytes(self):
# Decoder is patched so no codec backend needs to be installed.
with patch(
"sglang.srt.utils.common.VideoDecoderWrapper",
side_effect=RuntimeError("invalid data found when processing input"),
):
self._assert_client_error(b"definitely not a video", Modality.VIDEO)
class TestServerFaultStaysServerError(CustomTestCase):
"""``load_video`` catches the decoder broadly; these are the exclusions."""
def _assert_server_error(self, side_effect):
with patch(
"sglang.srt.utils.common.VideoDecoderWrapper", side_effect=side_effect
):
with self.assertRaises(RuntimeError):
_StubProcessor._load_single_item(b"payload", Modality.VIDEO)
def test_missing_decoder_backend(self):
self._assert_server_error(ImportError("no decoder backend installed"))
def test_decoder_oom(self):
self._assert_server_error(MemoryError("out of memory"))
class TestClientMediaExceptions(CustomTestCase):
def test_tuple_covers_the_documented_families(self):
for exc_type in (
requests.exceptions.HTTPError,
requests.exceptions.ConnectionError,
requests.exceptions.Timeout,
binascii.Error, # invalid base64
):
with self.subTest(exc_type=exc_type.__name__):
self.assertTrue(issubclass(exc_type, CLIENT_MEDIA_EXCEPTIONS))
if __name__ == "__main__":
unittest.main()