Fix Diffusion Request Validation to allow missing input artifacts if the input only contains text (#16610)

This commit is contained in:
Kangyan-Zhou
2026-01-31 23:38:40 -08:00
committed by GitHub
parent 4d28cda007
commit 9c168fcac7
2 changed files with 33 additions and 19 deletions
@@ -71,6 +71,7 @@ class VideoResponse(BaseModel):
class VideoGenerationsRequest(BaseModel): class VideoGenerationsRequest(BaseModel):
prompt: str prompt: str
input_reference: Optional[str] = None input_reference: Optional[str] = None
reference_url: Optional[str] = None
model: Optional[str] = None model: Optional[str] = None
seconds: Optional[int] = 4 seconds: Optional[int] = 4
size: Optional[str] = "" size: Optional[str] = ""
@@ -178,20 +178,26 @@ async def create_video(
content_type = request.headers.get("content-type", "").lower() content_type = request.headers.get("content-type", "").lower()
request_id = generate_request_id() request_id = generate_request_id()
server_args = get_global_server_args()
task_type = server_args.pipeline_config.task_type
if "multipart/form-data" in content_type: if "multipart/form-data" in content_type:
if not prompt: if not prompt:
raise HTTPException(status_code=400, detail="prompt is required") raise HTTPException(status_code=400, detail="prompt is required")
if input_reference is None and reference_url is None: # Validate image input based on model task type
image_list = merge_image_input_list(input_reference, reference_url)
if task_type.requires_image_input() and not image_list:
raise HTTPException( raise HTTPException(
status_code=400, status_code=400,
detail="input_reference file or reference_url is required", detail="input_reference or reference_url is required for image-to-video generation",
) )
image_list = merge_image_input_list(input_reference, reference_url) input_path = None
# Save first input image if image_list:
# Save first input image for image-to-video generation
image = image_list[0] image = image_list[0]
uploads_dir = os.path.join("outputs", "uploads") uploads_dir = os.path.join("outputs", "uploads")
os.makedirs(uploads_dir, exist_ok=True) os.makedirs(uploads_dir, exist_ok=True)
filename = image.filename if hasattr(image, "filename") else f"url_image" filename = image.filename if hasattr(image, "filename") else "url_image"
input_path = os.path.join(uploads_dir, f"{request_id}_{filename}") input_path = os.path.join(uploads_dir, f"{request_id}_{filename}")
try: try:
input_path = await save_image_to_path(image, input_path) input_path = await save_image_to_path(image, input_path)
@@ -246,6 +252,15 @@ async def create_video(
extra_json = payload.pop("extra_json", None) extra_json = payload.pop("extra_json", None)
if isinstance(extra_json, dict): if isinstance(extra_json, dict):
payload.update(extra_json) payload.update(extra_json)
# Validate image input based on model task type
has_image_input = payload.get("reference_url") or payload.get(
"input_reference"
)
if task_type.requires_image_input() and not has_image_input:
raise HTTPException(
status_code=400,
detail="input_reference or reference_url is required for image-to-video generation",
)
# for not multipart/form-data type # for not multipart/form-data type
if payload.get("reference_url"): if payload.get("reference_url"):
image_list = merge_image_input_list(payload.get("reference_url")) image_list = merge_image_input_list(payload.get("reference_url"))
@@ -253,9 +268,7 @@ async def create_video(
image = image_list[0] image = image_list[0]
uploads_dir = os.path.join("outputs", "uploads") uploads_dir = os.path.join("outputs", "uploads")
os.makedirs(uploads_dir, exist_ok=True) os.makedirs(uploads_dir, exist_ok=True)
filename = ( filename = image.filename if hasattr(image, "filename") else "url_image"
image.filename if hasattr(image, "filename") else f"url_image"
)
input_path = os.path.join(uploads_dir, f"{request_id}_{filename}") input_path = os.path.join(uploads_dir, f"{request_id}_{filename}")
try: try:
input_path = await save_image_to_path(image, input_path) input_path = await save_image_to_path(image, input_path)
@@ -277,7 +290,7 @@ async def create_video(
# Build Req for scheduler # Build Req for scheduler
batch = prepare_request( batch = prepare_request(
server_args=get_global_server_args(), server_args=server_args,
sampling_params=sampling_params, sampling_params=sampling_params,
) )
# Add diffusers_kwargs if provided # Add diffusers_kwargs if provided