[diffusion] fix: validate openai sampling dimensions (#28791)
This commit is contained in:
@@ -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}
|
||||
|
||||
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user