[diffusion] fix: validate openai sampling dimensions (#28791)

This commit is contained in:
Mick
2026-06-21 09:44:38 +08:00
committed by GitHub
parent 6a16573a7f
commit c65f4ea692
2 changed files with 69 additions and 8 deletions
@@ -10,7 +10,7 @@ from contextlib import contextmanager
from typing import Any, Generator, List, Optional, Union from typing import Any, Generator, List, Optional, Union
import httpx import httpx
from fastapi import UploadFile from fastapi import HTTPException, UploadFile
from sglang.multimodal_gen.configs.sample.sampling_params import ( from sglang.multimodal_gen.configs.sample.sampling_params import (
DataType, DataType,
@@ -54,6 +54,23 @@ DEFAULT_FPS = 24
DEFAULT_VIDEO_SECONDS = 4 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]: def flatten_extra_params(payload: Any) -> dict[str, Any]:
"""Promote vLLM-Omni-style extra_params into regular request fields.""" """Promote vLLM-Omni-style extra_params into regular request fields."""
if not isinstance(payload, dict): if not isinstance(payload, dict):
@@ -124,14 +141,22 @@ def build_sampling_params(request_id: str, **kwargs) -> SamplingParams:
# parse "WxH" size string if provided # parse "WxH" size string if provided
size = kwargs.pop("size", None) size = kwargs.pop("size", None)
if size: if size:
w, h = parse_size(size) w, h = _parse_size_or_raise(size)
if w is not None:
# treat None dimensions as unset so parsed size can fill them # treat None dimensions as unset so parsed size can fill them
if kwargs.get("width") is None: if kwargs.get("width") is None:
kwargs["width"] = w kwargs["width"] = w
if kwargs.get("height") is None: if kwargs.get("height") is None:
kwargs["height"] = h 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 # filter out None values to let SamplingParams defaults apply
kwargs = {k: v for k, v in kwargs.items() if v is not None} kwargs = {k: v for k, v in kwargs.items() if v is not None}
kwargs.setdefault("save_output", True) kwargs.setdefault("save_output", True)
@@ -6,7 +6,9 @@ import io
from starlette.datastructures import UploadFile as StarletteUploadFile from starlette.datastructures import UploadFile as StarletteUploadFile
from sglang.multimodal_gen.runtime.entrypoints.openai.utils import ( from sglang.multimodal_gen.runtime.entrypoints.openai.utils import (
_parse_size_or_raise,
_save_upload_to_path, _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 saved_path == str(target_path)
assert target_path.read_bytes() == b"image-bytes" 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")