From c65f4ea692dde23b8254f6f442f6bea3da6e712d Mon Sep 17 00:00:00 2001 From: Mick Date: Sun, 21 Jun 2026 09:44:38 +0800 Subject: [PATCH] [diffusion] fix: validate openai sampling dimensions (#28791) --- .../runtime/entrypoints/openai/utils.py | 41 +++++++++++++++---- .../test/unit/test_openai_utils.py | 36 ++++++++++++++++ 2 files changed, 69 insertions(+), 8 deletions(-) diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/utils.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/utils.py index 9736791cc..7f6020af1 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/openai/utils.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/utils.py @@ -10,7 +10,7 @@ from contextlib import contextmanager from typing import Any, Generator, List, Optional, Union import httpx -from fastapi import UploadFile +from fastapi import HTTPException, UploadFile from sglang.multimodal_gen.configs.sample.sampling_params import ( DataType, @@ -54,6 +54,23 @@ DEFAULT_FPS = 24 DEFAULT_VIDEO_SECONDS = 4 +def _bad_request(message: str) -> HTTPException: + return HTTPException(status_code=400, detail=message) + + +def _parse_size_or_raise(size: str) -> tuple[int, int]: + width, height = parse_size(size) + if width is None or height is None or width <= 0 or height <= 0: + raise _bad_request("size must be formatted as positive WIDTHxHEIGHT") + return width, height + + +def _validate_positive_int(kwargs: dict[str, Any], name: str) -> None: + value = kwargs.get(name) + if value is not None and int(value) <= 0: + raise _bad_request(f"{name} must be positive") + + def flatten_extra_params(payload: Any) -> dict[str, Any]: """Promote vLLM-Omni-style extra_params into regular request fields.""" if not isinstance(payload, dict): @@ -124,13 +141,21 @@ def build_sampling_params(request_id: str, **kwargs) -> SamplingParams: # parse "WxH" size string if provided size = kwargs.pop("size", None) if size: - w, h = parse_size(size) - if w is not None: - # treat None dimensions as unset so parsed size can fill them - if kwargs.get("width") is None: - kwargs["width"] = w - if kwargs.get("height") is None: - kwargs["height"] = h + w, h = _parse_size_or_raise(size) + # treat None dimensions as unset so parsed size can fill them + if kwargs.get("width") is None: + kwargs["width"] = w + if kwargs.get("height") is None: + kwargs["height"] = h + + for name in ( + "width", + "height", + "num_frames", + "num_inference_steps", + "num_outputs_per_prompt", + ): + _validate_positive_int(kwargs, name) # filter out None values to let SamplingParams defaults apply kwargs = {k: v for k, v in kwargs.items() if v is not None} diff --git a/python/sglang/multimodal_gen/test/unit/test_openai_utils.py b/python/sglang/multimodal_gen/test/unit/test_openai_utils.py index 9a00e06cd..69f8ad058 100644 --- a/python/sglang/multimodal_gen/test/unit/test_openai_utils.py +++ b/python/sglang/multimodal_gen/test/unit/test_openai_utils.py @@ -6,7 +6,9 @@ import io from starlette.datastructures import UploadFile as StarletteUploadFile from sglang.multimodal_gen.runtime.entrypoints.openai.utils import ( + _parse_size_or_raise, _save_upload_to_path, + _validate_positive_int, ) @@ -21,3 +23,37 @@ def test_save_upload_to_path_accepts_starlette_upload_file(tmp_path): assert saved_path == str(target_path) assert target_path.read_bytes() == b"image-bytes" + + +def test_parse_size_or_raise_accepts_positive_size(): + assert _parse_size_or_raise("512x768") == (512, 768) + + +def test_parse_size_or_raise_rejects_malformed_size(): + try: + _parse_size_or_raise("not-a-size") + except Exception as exc: + assert exc.status_code == 400 + assert "positive WIDTHxHEIGHT" in exc.detail + else: + raise AssertionError("expected bad request") + + +def test_parse_size_or_raise_rejects_non_positive_size(): + try: + _parse_size_or_raise("0x512") + except Exception as exc: + assert exc.status_code == 400 + assert "positive WIDTHxHEIGHT" in exc.detail + else: + raise AssertionError("expected bad request") + + +def test_validate_positive_int_rejects_non_positive_sampling_fields(): + try: + _validate_positive_int({"num_frames": 0}, "num_frames") + except Exception as exc: + assert exc.status_code == 400 + assert "num_frames must be positive" in exc.detail + else: + raise AssertionError("expected bad request")