diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/image_api.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/image_api.py index 26f4b97bd..2f6e58e7e 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/openai/image_api.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/image_api.py @@ -1,5 +1,6 @@ # Copied and adapted from: https://github.com/hao-ai-lab/FastVideo +import asyncio import base64 import contextlib import json @@ -63,6 +64,13 @@ def _get_extra_field(request, field_name): return value +def _get_request_field_or_extra(request, field_name): + value = getattr(request, field_name, None) + if value is not None: + return value + return _get_extra_field(request, field_name) + + def _parse_extra_container(value: Any) -> dict[str, Any]: if isinstance(value, str): try: @@ -83,6 +91,72 @@ def _read_b64_for_paths(paths: list[str]) -> list[str]: return result +async def _upload_and_cleanup_images(paths: list[str]) -> list[str | None]: + return await asyncio.gather( + *(cloud_storage.upload_and_cleanup(path) for path in paths) + ) + + +def _fallback_image_urls( + request_id: str, num_outputs: int, is_persistent: bool +) -> list[str] | None: + if not is_persistent: + return None + if num_outputs <= 1: + return [f"/v1/images/{request_id}/content"] + return [ + f"/v1/images/{request_id}/content?variant={idx}" for idx in range(num_outputs) + ] + + +def _select_image_variant_path(item: dict, variant: str | None) -> str | None: + file_paths = item.get("file_paths") + if file_paths: + variant_idx = _image_variant_index(variant) + if variant_idx is None: + return None + if variant_idx < 0 or variant_idx >= len(file_paths): + return None + return file_paths[variant_idx] + + if variant not in (None, "0", 0): + return None + return item.get("file_path") + + +def _image_variant_index(variant: str | None) -> int | None: + try: + return 0 if variant is None else int(variant) + except (TypeError, ValueError): + return None + + +def _select_image_variant_cloud_url(item: dict, variant: str | None) -> str | None: + variant_idx = _image_variant_index(variant) + if variant_idx is None: + return None + + urls = item.get("urls") + if urls and 0 <= variant_idx < len(urls): + return urls[variant_idx] + if variant_idx == 0: + return item.get("url") + return None + + +def _raise_if_image_variant_not_found(item: dict, variant: str | None) -> None: + file_paths = item.get("file_paths") + if not file_paths: + return + + variant_idx = _image_variant_index(variant) + if variant_idx is None or variant_idx < 0 or variant_idx >= len(file_paths): + raise HTTPException( + status_code=404, + detail=f"Image variant {variant} not found", + ) + + def _build_image_response_kwargs( save_file_path_list: list[str], resp_format: str, @@ -92,7 +166,9 @@ def _build_image_response_kwargs( *, b64_list: list[str] | None = None, cloud_url: str | None = None, + cloud_urls: list[str | None] | None = None, fallback_url: str | None = None, + fallback_urls: list[str] | None = None, is_persistent: bool = True, ) -> dict: """Build ImageResponse data list. @@ -115,25 +191,34 @@ def _build_image_response_kwargs( ] ret = {"data": data} elif resp_format == "url": - url = cloud_url or fallback_url - if not url: + if cloud_urls is None and cloud_url is not None: + cloud_urls = [cloud_url] + if fallback_urls is None and fallback_url is not None: + fallback_urls = [fallback_url] + + data = [] + for idx, path in enumerate(save_file_path_list): + url = None + if cloud_urls is not None and idx < len(cloud_urls): + url = cloud_urls[idx] + if not url and fallback_urls is not None and idx < len(fallback_urls): + url = fallback_urls[idx] + if not url: + break + data.append( + ImageResponseData( + url=url, + revised_prompt=prompt, + file_path=os.path.abspath(path) if is_persistent else None, + ) + ) + + if len(data) != len(save_file_path_list): raise HTTPException( status_code=400, detail="response_format='url' requires cloud storage to be configured.", ) - ret = { - "data": [ - ImageResponseData( - url=url, - revised_prompt=prompt, - file_path=( - os.path.abspath(save_file_path_list[0]) - if is_persistent - else None - ), - ) - ], - } + ret = {"data": data} else: raise HTTPException( status_code=400, detail=f"response_format={resp_format} is not supported" @@ -201,21 +286,11 @@ async def generations( perf_dump_path=request.perf_dump_path, use_pe=_get_extra_field(request, "use_pe"), preset=_get_extra_field(request, "preset"), - progressive_mode=( - request.progressive_mode - if request.progressive_mode is not None - else _get_extra_field(request, "progressive_mode") - ), - progressive_levels=( - request.progressive_levels - if request.progressive_levels is not None - else _get_extra_field(request, "progressive_levels") - ), - progressive_delta=( - request.progressive_delta - if request.progressive_delta is not None - else _get_extra_field(request, "progressive_delta") + progressive_mode=_get_request_field_or_extra(request, "progressive_mode"), + progressive_levels=_get_request_field_or_extra( + request, "progressive_levels" ), + progressive_delta=_get_request_field_or_extra(request, "progressive_delta"), ) trace_headers = extract_trace_headers(raw_request.headers) batch = prepare_request( @@ -246,16 +321,29 @@ async def generations( else None ) - cloud_url = await cloud_storage.upload_and_cleanup(save_file_path) - is_persistent = server_args.output_path is not None + cloud_urls = await _upload_and_cleanup_images(save_file_path_list) + cloud_url = cloud_urls[0] if cloud_urls else None + fallback_urls = _fallback_image_urls( + request_id, len(save_file_path_list), is_persistent + ) await IMAGE_STORE.upsert( request_id, { "id": request_id, "created_at": int(time.time()), "file_path": None if cloud_url or not is_persistent else save_file_path, + "file_paths": ( + None + if not is_persistent + else [ + None if url else path + for path, url in zip(save_file_path_list, cloud_urls) + ] + ), "url": cloud_url, + "urls": cloud_urls, + "num_outputs": len(save_file_path_list), }, ) @@ -266,8 +354,8 @@ async def generations( request_id, result, b64_list=b64_list, - cloud_url=cloud_url, - fallback_url=f"/v1/images/{request_id}/content" if is_persistent else None, + cloud_urls=cloud_urls, + fallback_urls=fallback_urls, is_persistent=is_persistent, ) @@ -381,19 +469,32 @@ async def edits( else None ) - cloud_url = await cloud_storage.upload_and_cleanup(save_file_path) - is_persistent = server_args.output_path is not None is_input_persistent = server_args.input_save_path is not None + cloud_urls = await _upload_and_cleanup_images(save_file_path_list) + cloud_url = cloud_urls[0] if cloud_urls else None + fallback_urls = _fallback_image_urls( + request_id, len(save_file_path_list), is_persistent + ) await IMAGE_STORE.upsert( request_id, { "id": request_id, "created_at": int(time.time()), "file_path": None if cloud_url or not is_persistent else save_file_path, + "file_paths": ( + None + if not is_persistent + else [ + None if url else path + for path, url in zip(save_file_path_list, cloud_urls) + ] + ), "url": cloud_url, + "urls": cloud_urls, "input_image_paths": input_paths if is_input_persistent else None, "num_input_images": len(input_paths), + "num_outputs": len(save_file_path_list), }, ) @@ -404,8 +505,8 @@ async def edits( request_id, result, b64_list=b64_list, - cloud_url=cloud_url, - fallback_url=f"/v1/images/{request_id}/content" if is_persistent else None, + cloud_urls=cloud_urls, + fallback_urls=fallback_urls, is_persistent=is_persistent, ) @@ -420,13 +521,18 @@ async def download_image_content( if not item: raise HTTPException(status_code=404, detail="Image not found") - if item.get("url"): + _raise_if_image_variant_not_found(item, variant) + file_path = _select_image_variant_path(item, variant) + if not file_path: + cloud_url = _select_image_variant_cloud_url(item, variant) + else: + cloud_url = None + if not file_path and cloud_url: raise HTTPException( status_code=400, - detail=f"Image has been uploaded to cloud storage. Please use the cloud URL: {item.get('url')}", + detail=f"Image has been uploaded to cloud storage. Please use the cloud URL: {cloud_url}", ) - file_path = item.get("file_path") if not file_path: raise HTTPException( status_code=404, diff --git a/python/sglang/multimodal_gen/test/unit/test_openai_image_api.py b/python/sglang/multimodal_gen/test/unit/test_openai_image_api.py new file mode 100644 index 000000000..abba9174f --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_openai_image_api.py @@ -0,0 +1,123 @@ +import os + +from fastapi import HTTPException + +from sglang.multimodal_gen.runtime.entrypoints.openai.image_api import ( + _build_image_response_kwargs, + _fallback_image_urls, + _raise_if_image_variant_not_found, + _select_image_variant_cloud_url, + _select_image_variant_path, +) +from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch + + +def test_url_response_returns_one_item_per_output_path(): + paths = ["first.png", "second.png"] + + response = _build_image_response_kwargs( + paths, + "url", + "a lantern", + "req-123", + OutputBatch(), + cloud_urls=["https://cdn.example/first.png", "https://cdn.example/second.png"], + fallback_urls=_fallback_image_urls("req-123", len(paths), True), + is_persistent=True, + ) + + assert [item.url for item in response["data"]] == [ + "https://cdn.example/first.png", + "https://cdn.example/second.png", + ] + assert [item.file_path for item in response["data"]] == [ + os.path.abspath("first.png"), + os.path.abspath("second.png"), + ] + + +def test_url_response_uses_variant_fallback_urls_for_multiple_persistent_outputs(): + paths = ["first.png", "second.png"] + + response = _build_image_response_kwargs( + paths, + "url", + "a lantern", + "req-123", + OutputBatch(), + fallback_urls=_fallback_image_urls("req-123", len(paths), True), + is_persistent=True, + ) + + assert [item.url for item in response["data"]] == [ + "/v1/images/req-123/content?variant=0", + "/v1/images/req-123/content?variant=1", + ] + + +def test_select_image_variant_path_reads_indexed_file_paths(): + item = { + "file_path": "first.png", + "file_paths": ["first.png", "second.png"], + } + + assert _select_image_variant_path(item, None) == "first.png" + assert _select_image_variant_path(item, "1") == "second.png" + + +def test_raise_if_image_variant_not_found_handles_out_of_range_variant(): + item = { + "file_path": "first.png", + "file_paths": ["first.png", "second.png"], + } + + try: + _raise_if_image_variant_not_found(item, "5") + except HTTPException as exc: + assert exc.status_code == 404 + assert exc.detail == "Image variant 5 not found" + else: + raise AssertionError("Expected HTTPException") + + +def test_select_image_variant_path_returns_none_for_cloud_only_variant(): + item = { + "file_path": None, + "file_paths": [None, "second.png"], + "urls": ["https://cdn.example/first.png"], + } + + assert _select_image_variant_path(item, "0") is None + assert _select_image_variant_path(item, "1") == "second.png" + + +def test_select_image_variant_cloud_url_keeps_variant_alignment(): + item = { + "url": "https://cdn.example/first.png", + "urls": [ + "https://cdn.example/first.png", + None, + "https://cdn.example/third.png", + ], + } + + assert ( + _select_image_variant_cloud_url(item, None) == "https://cdn.example/first.png" + ) + assert _select_image_variant_cloud_url(item, "1") is None + assert _select_image_variant_cloud_url(item, "2") == "https://cdn.example/third.png" + + +def test_select_image_variant_cloud_url_falls_back_to_single_url(): + item = {"url": "https://cdn.example/only.png"} + + assert _select_image_variant_cloud_url(item, None) == "https://cdn.example/only.png" + assert _select_image_variant_cloud_url(item, "0") == "https://cdn.example/only.png" + assert _select_image_variant_cloud_url(item, "1") is None + + +def test_select_image_variant_path_falls_back_to_single_file_path(): + item = {"file_path": "only.png"} + + assert _select_image_variant_path(item, None) == "only.png" + assert _select_image_variant_path(item, "0") == "only.png"