vllm-project--vllm-omni
2139 行
74 KiB
Python
2139 行
74 KiB
Python
# SPDX-License-Identifier: Apache-2.0
|
|
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
|
"""
|
|
Tests for async image generation API endpoints.
|
|
|
|
This module contains unit tests and integration tests (with mocking) for the
|
|
OpenAI-compatible async text-to-image generation API endpoints in api_server.py.
|
|
"""
|
|
|
|
import base64
|
|
import io
|
|
import json
|
|
from argparse import Namespace
|
|
from http import HTTPStatus
|
|
from types import SimpleNamespace
|
|
|
|
import pytest
|
|
from fastapi import FastAPI, HTTPException
|
|
from fastapi.testclient import TestClient
|
|
from PIL import Image
|
|
from pytest_mock import MockerFixture
|
|
from vllm import SamplingParams
|
|
from vllm.entrypoints.openai.models.protocol import BaseModelPath
|
|
from vllm.sampling_params import RequestOutputKind
|
|
|
|
from vllm_omni.entrypoints.async_omni import AsyncOmni
|
|
from vllm_omni.entrypoints.openai.api_server import _check_max_generated_image_size, _DiffusionServingModels, router
|
|
from vllm_omni.entrypoints.openai.image_api_utils import (
|
|
encode_image_base64,
|
|
parse_size,
|
|
)
|
|
from vllm_omni.entrypoints.openai.serving_chat import OmniOpenAIServingChat
|
|
from vllm_omni.errors import GuardrailViolationError
|
|
from vllm_omni.inputs.data import OmniDiffusionSamplingParams
|
|
|
|
pytestmark = [pytest.mark.core_model, pytest.mark.cpu]
|
|
|
|
# Unit Tests
|
|
|
|
|
|
def test_parse_size_valid():
|
|
"""Test size parsing with valid inputs"""
|
|
assert parse_size("1024x1024") == (1024, 1024)
|
|
assert parse_size("512x768") == (512, 768)
|
|
assert parse_size("256x256") == (256, 256)
|
|
assert parse_size("1792x1024") == (1792, 1024)
|
|
assert parse_size("1024x1792") == (1024, 1792)
|
|
|
|
|
|
def test_parse_size_invalid():
|
|
"""Test size parsing with invalid inputs"""
|
|
with pytest.raises(ValueError, match="Invalid size format"):
|
|
parse_size("invalid")
|
|
|
|
with pytest.raises(ValueError, match="Invalid size format"):
|
|
parse_size("1024")
|
|
|
|
with pytest.raises(ValueError, match="Invalid size format"):
|
|
parse_size("1024x")
|
|
|
|
with pytest.raises(ValueError, match="Invalid size format"):
|
|
parse_size("x1024")
|
|
|
|
|
|
def test_parse_size_negative():
|
|
"""Test size parsing with negative or zero dimensions"""
|
|
with pytest.raises(ValueError, match="positive integers"):
|
|
parse_size("0x1024")
|
|
|
|
with pytest.raises(ValueError, match="positive integers"):
|
|
parse_size("1024x0")
|
|
|
|
with pytest.raises(ValueError):
|
|
parse_size("-1024x1024")
|
|
|
|
|
|
def test_parse_size_edge_cases():
|
|
"""Test size parsing with edge cases like empty strings and non-integers"""
|
|
# Empty string
|
|
with pytest.raises(ValueError, match="non-empty string"):
|
|
parse_size("")
|
|
|
|
# Non-integer dimensions
|
|
with pytest.raises(ValueError, match="must be integers"):
|
|
parse_size("abc x def")
|
|
|
|
with pytest.raises(ValueError, match="must be integers"):
|
|
parse_size("1024.5x768.5")
|
|
|
|
# Missing separator (user might forget 'x')
|
|
with pytest.raises(ValueError, match="separator"):
|
|
parse_size("1024 1024")
|
|
|
|
|
|
def test_encode_image_base64():
|
|
"""Test image encoding to base64"""
|
|
# Create a simple test image
|
|
img = Image.new("RGB", (64, 64), color="red")
|
|
b64_str = encode_image_base64(img)
|
|
|
|
# Should be valid base64
|
|
assert isinstance(b64_str, str)
|
|
assert len(b64_str) > 0
|
|
|
|
# Should decode back to PNG
|
|
decoded = base64.b64decode(b64_str)
|
|
decoded_img = Image.open(io.BytesIO(decoded))
|
|
|
|
# Verify properties
|
|
assert decoded_img.size == (64, 64)
|
|
assert decoded_img.format == "PNG"
|
|
|
|
|
|
# Integration Tests (with mocking)
|
|
|
|
|
|
class MockGenerationResult:
|
|
"""Mock result object compatible with current diffusion output shape."""
|
|
|
|
def __init__(self, images):
|
|
self.images = images
|
|
self.request_output = SimpleNamespace(images=images)
|
|
self.stage_durations = {}
|
|
self.peak_memory_mb = 0.0
|
|
|
|
|
|
class MockStageResult:
|
|
"""Mock multi-stage output for streaming image edit tests."""
|
|
|
|
def __init__(self, *, stage_id, final_output_type, text="", texts=None, images=None):
|
|
self.stage_id = stage_id
|
|
self.final_output_type = final_output_type
|
|
self.images = images or []
|
|
if texts is not None:
|
|
outputs = [SimpleNamespace(text=item, index=index) for index, item in enumerate(texts)]
|
|
elif text:
|
|
outputs = [SimpleNamespace(text=text, index=0)]
|
|
else:
|
|
outputs = []
|
|
self.request_output = SimpleNamespace(
|
|
outputs=outputs,
|
|
images=self.images,
|
|
)
|
|
self.stage_durations = {}
|
|
self.peak_memory_mb = 0.0
|
|
|
|
|
|
class FakeAsyncOmni:
|
|
"""Fake AsyncOmni that yields a single diffusion output."""
|
|
|
|
def __init__(self, images=None):
|
|
self.stage_configs = [
|
|
SimpleNamespace(stage_type="llm", is_comprehension=True),
|
|
SimpleNamespace(stage_type="diffusion", is_comprehension=False),
|
|
]
|
|
self.default_sampling_params_list = [SamplingParams(temperature=0.1), OmniDiffusionSamplingParams()]
|
|
self.captured_sampling_params_list = None
|
|
self.captured_prompt = None
|
|
self._images = images or [Image.new("RGB", (64, 64), color="green")]
|
|
|
|
async def generate(self, prompt, request_id, sampling_params=None, sampling_params_list=None, **kwargs):
|
|
if sampling_params_list is not None:
|
|
self.captured_sampling_params_list = sampling_params_list
|
|
else:
|
|
self.captured_sampling_params_list = [sampling_params]
|
|
self.captured_prompt = prompt
|
|
images = [img.copy() for img in self._images]
|
|
yield MockGenerationResult(images)
|
|
|
|
def __class_getitem__(cls, item):
|
|
return cls
|
|
|
|
|
|
@pytest.fixture
|
|
def mock_async_diffusion(mocker: MockerFixture):
|
|
"""Mock diffusion engine that matches the current async-generator API."""
|
|
|
|
class MockAsyncDiffusion:
|
|
def __init__(self) -> None:
|
|
self.is_running = True
|
|
self.check_health = mocker.AsyncMock()
|
|
self.captured_sampling_params_list = None
|
|
self.captured_prompt = None
|
|
self.generate_calls = 0
|
|
|
|
async def generate(self, **kwargs):
|
|
self.generate_calls += 1
|
|
n = kwargs["sampling_params_list"][0].num_outputs_per_prompt
|
|
self.captured_sampling_params_list = kwargs["sampling_params_list"]
|
|
self.captured_prompt = kwargs["prompt"]
|
|
images = [Image.new("RGB", (64, 64), color="blue") for _ in range(n)]
|
|
yield MockGenerationResult(images)
|
|
|
|
return MockAsyncDiffusion()
|
|
|
|
|
|
@pytest.fixture
|
|
def test_client(mock_async_diffusion):
|
|
"""Create test client with mocked async diffusion engine"""
|
|
from fastapi import FastAPI
|
|
|
|
from vllm_omni.entrypoints.openai.api_server import router
|
|
|
|
app = FastAPI()
|
|
app.include_router(router)
|
|
|
|
# Set up app state with diffusion engine
|
|
app.state.engine_client = mock_async_diffusion
|
|
app.state.diffusion_engine = mock_async_diffusion # Also set for health endpoint
|
|
app.state.stage_configs = [SimpleNamespace(stage_type="diffusion")]
|
|
from vllm.entrypoints.openai.models.protocol import BaseModelPath
|
|
|
|
from vllm_omni.entrypoints.openai.api_server import _DiffusionServingModels
|
|
|
|
app.state.openai_serving_models = _DiffusionServingModels(
|
|
[BaseModelPath(name="Qwen/Qwen-Image", model_path="Qwen/Qwen-Image")]
|
|
)
|
|
app.state.args = Namespace(
|
|
default_sampling_params='{"0": {"num_inference_steps":4, "guidance_scale":7.5, "generator_device":"cpu"}}',
|
|
max_generated_image_size=1024 * 1792,
|
|
)
|
|
|
|
return TestClient(app)
|
|
|
|
|
|
@pytest.fixture
|
|
def async_omni_test_client():
|
|
"""Create test client with mocked AsyncOmni engine."""
|
|
from fastapi import FastAPI
|
|
|
|
from vllm_omni.entrypoints.async_omni import AsyncOmni
|
|
from vllm_omni.entrypoints.openai.api_server import router
|
|
from vllm_omni.entrypoints.openai.serving_chat import OmniOpenAIServingChat
|
|
|
|
class FakeAsyncOmniClass(AsyncOmni):
|
|
def __init__(self):
|
|
stage_configs = [
|
|
SimpleNamespace(stage_type="llm", is_comprehension=True),
|
|
SimpleNamespace(stage_type="diffusion", is_comprehension=False),
|
|
]
|
|
default_sampling_params_list = [
|
|
SamplingParams(temperature=0.1),
|
|
OmniDiffusionSamplingParams(
|
|
num_inference_steps=4,
|
|
guidance_scale=7.5,
|
|
generator_device="cpu",
|
|
),
|
|
]
|
|
self.engine = SimpleNamespace(
|
|
stage_configs=stage_configs,
|
|
default_sampling_params_list=default_sampling_params_list,
|
|
)
|
|
self.default_sampling_params_list = default_sampling_params_list
|
|
self.captured_sampling_params_list = None
|
|
self.captured_prompt = None
|
|
self._images = [Image.new("RGB", (64, 64), color="green")]
|
|
self.od_config = SimpleNamespace(supports_multimodal_inputs=True)
|
|
|
|
async def generate(self, prompt, request_id, sampling_params=None, sampling_params_list=None, **kwargs):
|
|
if sampling_params_list is not None:
|
|
self.captured_sampling_params_list = sampling_params_list
|
|
else:
|
|
self.captured_sampling_params_list = [sampling_params]
|
|
self.captured_prompt = prompt
|
|
images = [img.copy() for img in self._images]
|
|
yield MockGenerationResult(images)
|
|
|
|
def __class_getitem__(cls, item):
|
|
return cls
|
|
|
|
def get_diffusion_od_config(self):
|
|
return self.od_config
|
|
|
|
app = FastAPI()
|
|
app.include_router(router)
|
|
|
|
engine = FakeAsyncOmniClass()
|
|
chat_handler = object.__new__(OmniOpenAIServingChat)
|
|
chat_handler.engine_client = engine
|
|
chat_handler._diffusion_engine = None
|
|
app.state.openai_serving_chat = chat_handler
|
|
app.state.engine_client = engine
|
|
app.state.stage_configs = [
|
|
SimpleNamespace(stage_type="llm"),
|
|
SimpleNamespace(stage_type="diffusion"),
|
|
]
|
|
app.state.args = Namespace(
|
|
default_sampling_params='{"1": {"num_inference_steps":4, "guidance_scale":7.5, "generator_device":"cpu"}}',
|
|
max_generated_image_size=1048576, # 1024*1024 to support resolution tests
|
|
)
|
|
return TestClient(app)
|
|
|
|
|
|
@pytest.fixture
|
|
def async_omni_rgba_test_client():
|
|
"""Create test client with mocked AsyncOmni engine returning RGBA output."""
|
|
from fastapi import FastAPI
|
|
|
|
from vllm_omni.entrypoints.async_omni import AsyncOmni
|
|
from vllm_omni.entrypoints.openai.api_server import router
|
|
from vllm_omni.entrypoints.openai.serving_chat import OmniOpenAIServingChat
|
|
|
|
class FakeAsyncOmniClass(AsyncOmni):
|
|
def __init__(self):
|
|
stage_configs = [
|
|
SimpleNamespace(stage_type="llm", is_comprehension=True),
|
|
SimpleNamespace(stage_type="diffusion", is_comprehension=False),
|
|
]
|
|
default_sampling_params_list = [
|
|
SamplingParams(temperature=0.1),
|
|
OmniDiffusionSamplingParams(),
|
|
]
|
|
self.engine = SimpleNamespace(
|
|
stage_configs=stage_configs,
|
|
default_sampling_params_list=default_sampling_params_list,
|
|
)
|
|
self.default_sampling_params_list = default_sampling_params_list
|
|
self.captured_sampling_params_list = None
|
|
self.captured_prompt = None
|
|
self._images = [Image.new("RGBA", (64, 64), color=(0, 255, 0, 128))]
|
|
self.od_config = SimpleNamespace(supports_multimodal_inputs=True)
|
|
|
|
async def generate(self, prompt, request_id, sampling_params=None, sampling_params_list=None, **kwargs):
|
|
if sampling_params_list is not None:
|
|
self.captured_sampling_params_list = sampling_params_list
|
|
else:
|
|
self.captured_sampling_params_list = [sampling_params]
|
|
self.captured_prompt = prompt
|
|
images = [img.copy() for img in self._images]
|
|
yield MockGenerationResult(images)
|
|
|
|
def __class_getitem__(cls, item):
|
|
return cls
|
|
|
|
def get_diffusion_od_config(self):
|
|
return self.od_config
|
|
|
|
app = FastAPI()
|
|
app.include_router(router)
|
|
|
|
engine = FakeAsyncOmniClass()
|
|
chat_handler = object.__new__(OmniOpenAIServingChat)
|
|
chat_handler.engine_client = engine
|
|
chat_handler._diffusion_engine = None
|
|
app.state.openai_serving_chat = chat_handler
|
|
app.state.engine_client = engine
|
|
app.state.stage_configs = [
|
|
SimpleNamespace(stage_type="llm"),
|
|
SimpleNamespace(stage_type="diffusion"),
|
|
]
|
|
app.state.args = Namespace(
|
|
default_sampling_params='{"1": {"num_inference_steps":4, "guidance_scale":7.5, "generator_device":"cpu"}}',
|
|
max_generated_image_size=1048576,
|
|
)
|
|
return TestClient(app)
|
|
|
|
|
|
@pytest.fixture
|
|
def async_omni_stage_configs_only_client():
|
|
"""Create test client with refactored AsyncOmni compatibility surface only."""
|
|
from fastapi import FastAPI
|
|
|
|
from vllm_omni.entrypoints.async_omni import AsyncOmni
|
|
from vllm_omni.entrypoints.openai.api_server import router
|
|
from vllm_omni.entrypoints.openai.serving_chat import OmniOpenAIServingChat
|
|
|
|
class FakeAsyncOmniClass(AsyncOmni):
|
|
def __init__(self):
|
|
stage_configs = [
|
|
SimpleNamespace(stage_type="llm", is_comprehension=True),
|
|
SimpleNamespace(stage_type="diffusion", is_comprehension=False),
|
|
]
|
|
default_sampling_params_list = [
|
|
SamplingParams(temperature=0.1),
|
|
OmniDiffusionSamplingParams(),
|
|
]
|
|
self.engine = SimpleNamespace(
|
|
stage_configs=stage_configs,
|
|
default_sampling_params_list=default_sampling_params_list,
|
|
)
|
|
self.default_sampling_params_list = default_sampling_params_list
|
|
self.captured_sampling_params_list = None
|
|
self.captured_prompt = None
|
|
self._images = [Image.new("RGB", (64, 64), color="green")]
|
|
self.od_config = SimpleNamespace(supports_multimodal_inputs=True)
|
|
|
|
async def generate(self, prompt, request_id, sampling_params=None, sampling_params_list=None, **kwargs):
|
|
if sampling_params_list is not None:
|
|
self.captured_sampling_params_list = sampling_params_list
|
|
else:
|
|
self.captured_sampling_params_list = [sampling_params]
|
|
self.captured_prompt = prompt
|
|
images = [img.copy() for img in self._images]
|
|
yield MockGenerationResult(images)
|
|
|
|
def __class_getitem__(cls, item):
|
|
return cls
|
|
|
|
def get_diffusion_od_config(self):
|
|
return self.od_config
|
|
|
|
app = FastAPI()
|
|
app.include_router(router)
|
|
|
|
engine = FakeAsyncOmniClass()
|
|
assert not hasattr(engine, "stage_list")
|
|
app.state.engine_client = engine
|
|
chat_handler = object.__new__(OmniOpenAIServingChat)
|
|
chat_handler.engine_client = engine
|
|
chat_handler._diffusion_engine = None
|
|
app.state.openai_serving_chat = chat_handler
|
|
app.state.args = Namespace(
|
|
default_sampling_params='{"1": {"num_inference_steps":4, "guidance_scale":7.5, "generator_device":"cpu"}}',
|
|
max_generated_image_size=1024 * 1792,
|
|
)
|
|
return TestClient(app)
|
|
|
|
|
|
@pytest.fixture
|
|
def streaming_image_edit_client():
|
|
"""Create a multi-stage client whose engine yields AR text before image output."""
|
|
from fastapi import FastAPI
|
|
|
|
from vllm_omni.entrypoints.async_omni import AsyncOmni
|
|
from vllm_omni.entrypoints.openai.api_server import router
|
|
from vllm_omni.entrypoints.openai.serving_chat import OmniOpenAIServingChat
|
|
|
|
class FakeAsyncOmniClass(AsyncOmni):
|
|
def __init__(self):
|
|
stage_configs = [
|
|
SimpleNamespace(stage_type="llm", is_comprehension=True),
|
|
SimpleNamespace(stage_type="diffusion", is_comprehension=False),
|
|
]
|
|
default_sampling_params_list = [
|
|
SamplingParams(temperature=0.1),
|
|
OmniDiffusionSamplingParams(),
|
|
]
|
|
self.engine = SimpleNamespace(
|
|
stage_configs=stage_configs,
|
|
default_sampling_params_list=default_sampling_params_list,
|
|
)
|
|
self.default_sampling_params_list = default_sampling_params_list
|
|
self.captured_sampling_params_list = None
|
|
self.captured_prompt = None
|
|
self.od_config = SimpleNamespace(supports_multimodal_inputs=True)
|
|
|
|
async def generate(self, prompt, request_id, sampling_params=None, sampling_params_list=None, **kwargs):
|
|
self.captured_prompt = prompt
|
|
self.captured_sampling_params_list = sampling_params_list or [sampling_params]
|
|
assert self.captured_sampling_params_list[0].output_kind == RequestOutputKind.DELTA
|
|
yield MockStageResult(stage_id=0, final_output_type="text", text="recap")
|
|
yield MockStageResult(stage_id=0, final_output_type="text", text=" done")
|
|
yield MockStageResult(
|
|
stage_id=1,
|
|
final_output_type="image",
|
|
images=[Image.new("RGB", (32, 24), color="purple")],
|
|
)
|
|
|
|
def __class_getitem__(cls, item):
|
|
return cls
|
|
|
|
def get_diffusion_od_config(self):
|
|
return self.od_config
|
|
|
|
app = FastAPI()
|
|
app.include_router(router)
|
|
|
|
engine = FakeAsyncOmniClass()
|
|
chat_handler = object.__new__(OmniOpenAIServingChat)
|
|
chat_handler.engine_client = engine
|
|
chat_handler._diffusion_engine = None
|
|
app.state.openai_serving_chat = chat_handler
|
|
app.state.engine_client = engine
|
|
app.state.stage_configs = [
|
|
SimpleNamespace(stage_type="llm"),
|
|
SimpleNamespace(stage_type="diffusion"),
|
|
]
|
|
app.state.args = Namespace(
|
|
default_sampling_params='{"1": {"num_inference_steps":4, "guidance_scale":7.5, "generator_device":"cpu"}}',
|
|
max_generated_image_size=1024 * 1792,
|
|
)
|
|
return TestClient(app)
|
|
|
|
|
|
def test_health_endpoint(test_client):
|
|
"""Test health check endpoint for diffusion mode"""
|
|
response = test_client.get("/health")
|
|
assert response.status_code == 200
|
|
data = response.json()
|
|
assert data["status"] == "healthy"
|
|
|
|
|
|
def test_health_endpoint_no_engine():
|
|
"""Test health check endpoint when no engine is initialized"""
|
|
from fastapi import FastAPI
|
|
|
|
from vllm_omni.entrypoints.openai.api_server import router
|
|
|
|
app = FastAPI()
|
|
app.include_router(router)
|
|
# Don't set any engine
|
|
|
|
client = TestClient(app)
|
|
response = client.get("/health")
|
|
assert response.status_code == 503
|
|
data = response.json()
|
|
assert data["status"] == "unhealthy"
|
|
|
|
|
|
def test_health_endpoint_dead_engine():
|
|
"""Health returns 503 when the engine raises EngineDeadError."""
|
|
from unittest.mock import AsyncMock
|
|
|
|
from fastapi import FastAPI
|
|
from vllm.v1.engine.exceptions import EngineDeadError
|
|
|
|
from vllm_omni.entrypoints.openai.api_server import router
|
|
|
|
app = FastAPI()
|
|
app.include_router(router)
|
|
|
|
dead_engine = AsyncMock()
|
|
dead_engine.check_health = AsyncMock(side_effect=EngineDeadError())
|
|
app.state.engine_client = dead_engine
|
|
|
|
client = TestClient(app)
|
|
response = client.get("/health")
|
|
assert response.status_code == 503
|
|
data = response.json()
|
|
assert data["status"] == "unhealthy"
|
|
|
|
|
|
def test_models_endpoint(test_client):
|
|
"""Test /v1/models endpoint for diffusion mode"""
|
|
response = test_client.get("/v1/models")
|
|
assert response.status_code == 200
|
|
data = response.json()
|
|
assert data["object"] == "list"
|
|
assert len(data["data"]) == 1
|
|
assert data["data"][0]["id"] == "Qwen/Qwen-Image"
|
|
assert data["data"][0]["object"] == "model"
|
|
|
|
|
|
def test_models_endpoint_no_engine():
|
|
"""Test /v1/models endpoint when no engine is initialized"""
|
|
from fastapi import FastAPI
|
|
|
|
from vllm_omni.entrypoints.openai.api_server import router
|
|
|
|
app = FastAPI()
|
|
app.include_router(router)
|
|
# Don't set any engine
|
|
|
|
client = TestClient(app)
|
|
response = client.get("/v1/models")
|
|
assert response.status_code == 200
|
|
data = response.json()
|
|
assert data["object"] == "list"
|
|
assert len(data["data"]) == 0
|
|
|
|
|
|
def test_generate_single_image(test_client):
|
|
"""Test generating a single image"""
|
|
# Single-stage path should not require openai_serving_chat.
|
|
assert not hasattr(test_client.app.state, "openai_serving_chat")
|
|
|
|
response = test_client.post(
|
|
"/v1/images/generations",
|
|
json={
|
|
"prompt": "a cat",
|
|
"n": 1,
|
|
"size": "1024x1024",
|
|
},
|
|
)
|
|
assert response.status_code == 200
|
|
data = response.json()
|
|
|
|
# Check response structure
|
|
assert "created" in data
|
|
assert isinstance(data["created"], int)
|
|
assert "data" in data
|
|
assert len(data["data"]) == 1
|
|
assert "b64_json" in data["data"][0]
|
|
|
|
# Verify image can be decoded
|
|
img_bytes = base64.b64decode(data["data"][0]["b64_json"])
|
|
img = Image.open(io.BytesIO(img_bytes))
|
|
assert img.size == (64, 64) # Our mock returns 64x64 images
|
|
assert test_client.app.state.engine_client.captured_prompt["modalities"] == ["image"]
|
|
|
|
|
|
def test_generate_images_guardrail_error_returns_400(test_client, mock_async_diffusion):
|
|
async def blocked_generate(**kwargs):
|
|
raise GuardrailViolationError("Input was blocked by Cosmos3 guardrails.")
|
|
yield MockGenerationResult([]) # pragma: no cover
|
|
|
|
mock_async_diffusion.generate = blocked_generate
|
|
response = test_client.post(
|
|
"/v1/images/generations",
|
|
json={
|
|
"prompt": "blocked prompt",
|
|
"n": 1,
|
|
"size": "1024x1024",
|
|
},
|
|
)
|
|
assert response.status_code == 400
|
|
assert response.json()["detail"] == "Input was blocked by Cosmos3 guardrails."
|
|
|
|
|
|
def test_generate_images_async_omni_sampling_params(async_omni_test_client):
|
|
"""Test AsyncOmni path uses per-stage sampling params."""
|
|
response = async_omni_test_client.post(
|
|
"/v1/images/generations",
|
|
json={
|
|
"prompt": "a cat",
|
|
"n": 2,
|
|
"size": "256x256",
|
|
"seed": 7,
|
|
},
|
|
)
|
|
assert response.status_code == 200
|
|
engine = async_omni_test_client.app.state.engine_client
|
|
captured = engine.captured_sampling_params_list
|
|
assert captured is not None
|
|
assert len(captured) == 2
|
|
assert captured[0].temperature == 0.1
|
|
assert captured[1].num_outputs_per_prompt == 2
|
|
assert captured[1].height == 256
|
|
assert captured[1].width == 256
|
|
assert captured[1].seed == 7
|
|
|
|
|
|
def test_generate_images_async_omni_stage_configs_only(async_omni_stage_configs_only_client):
|
|
"""Regression: image generation accepts refactored AsyncOmni without stage_list."""
|
|
response = async_omni_stage_configs_only_client.post(
|
|
"/v1/images/generations",
|
|
json={
|
|
"prompt": "a castle",
|
|
"n": 1,
|
|
"size": "256x256",
|
|
"seed": 11,
|
|
},
|
|
)
|
|
assert response.status_code == 200
|
|
data = response.json()
|
|
assert len(data["data"]) == 1
|
|
engine = async_omni_stage_configs_only_client.app.state.engine_client
|
|
captured = engine.captured_sampling_params_list
|
|
assert captured is not None
|
|
assert len(captured) == 2
|
|
assert captured[1].seed == 11
|
|
|
|
|
|
def test_multistage_images_async_omni_construction(async_omni_test_client):
|
|
"""Regression: multistage image generation builds the expected chat-style payload."""
|
|
response = async_omni_test_client.post(
|
|
"/v1/images/generations",
|
|
json={
|
|
"prompt": "a cat",
|
|
"n": 2,
|
|
"size": "128x256",
|
|
"seed": 7,
|
|
"num_inference_steps": 12,
|
|
"guidance_scale": 6.5,
|
|
},
|
|
)
|
|
assert response.status_code == 200
|
|
|
|
engine = async_omni_test_client.app.state.engine_client
|
|
captured_prompt = engine.captured_prompt
|
|
assert captured_prompt["prompt"] == "a cat"
|
|
assert captured_prompt["modalities"] == ["image"]
|
|
assert captured_prompt["mm_processor_kwargs"] == {
|
|
"target_h": 256,
|
|
"target_w": 128,
|
|
}
|
|
|
|
captured = engine.captured_sampling_params_list
|
|
assert captured is not None
|
|
assert len(captured) == 2
|
|
assert captured[0].temperature == 0.1
|
|
assert captured[0].seed == 7
|
|
assert captured[1].num_outputs_per_prompt == 2
|
|
assert captured[1].width == 128
|
|
assert captured[1].height == 256
|
|
assert captured[1].seed == 7
|
|
assert captured[1].num_inference_steps == 12
|
|
assert captured[1].guidance_scale == 6.5
|
|
|
|
|
|
def test_generate_images_async_omni_glm_image_sets_stage0_max_tokens():
|
|
"""GLM-Image multistage: stage-0 gets target_h/w from requested size.
|
|
|
|
max_tokens comes from the deploy YAML default (upper-bound ceiling),
|
|
NOT computed dynamically from height/width.
|
|
"""
|
|
|
|
class FakeAsyncOmniClass(AsyncOmni):
|
|
def __init__(self):
|
|
stage_configs = [
|
|
SimpleNamespace(stage_type="llm", is_comprehension=True, model_arch="GlmImageForConditionalGeneration"),
|
|
SimpleNamespace(stage_type="diffusion", is_comprehension=False, model_arch="GlmImagePipeline"),
|
|
]
|
|
# YAML default max_tokens for GLM-Image AR stage (upper bound for 2048x2048 t2i)
|
|
default_sampling_params_list = [
|
|
SamplingParams(temperature=0.1, seed=42, max_tokens=4353),
|
|
OmniDiffusionSamplingParams(height=1024, width=1024),
|
|
]
|
|
self.engine = SimpleNamespace(
|
|
stage_configs=stage_configs,
|
|
default_sampling_params_list=default_sampling_params_list,
|
|
)
|
|
self.default_sampling_params_list = default_sampling_params_list
|
|
self.captured_sampling_params_list = None
|
|
self.captured_prompt = None
|
|
self._images = [Image.new("RGB", (64, 64), color="green")]
|
|
self.od_config = SimpleNamespace(supports_multimodal_inputs=True)
|
|
|
|
async def generate(self, prompt, request_id, sampling_params=None, sampling_params_list=None, **kwargs):
|
|
self.captured_sampling_params_list = (
|
|
sampling_params_list if sampling_params_list is not None else [sampling_params]
|
|
)
|
|
self.captured_prompt = prompt
|
|
yield MockGenerationResult([img.copy() for img in self._images])
|
|
|
|
def __class_getitem__(cls, item):
|
|
return cls
|
|
|
|
def get_diffusion_od_config(self):
|
|
return self.od_config
|
|
|
|
app = FastAPI()
|
|
app.include_router(router)
|
|
engine = FakeAsyncOmniClass()
|
|
chat_handler = object.__new__(OmniOpenAIServingChat)
|
|
chat_handler.engine_client = engine
|
|
chat_handler._diffusion_engine = None
|
|
app.state.openai_serving_chat = chat_handler
|
|
app.state.engine_client = engine
|
|
app.state.stage_configs = [
|
|
SimpleNamespace(stage_type="llm", model_arch="GlmImageForConditionalGeneration"),
|
|
SimpleNamespace(stage_type="diffusion", model_arch="GlmImagePipeline"),
|
|
]
|
|
app.state.openai_serving_models = _DiffusionServingModels(
|
|
[BaseModelPath(name="THUDM/GLM-4.5V", model_path="THUDM/GLM-4.5V")]
|
|
)
|
|
app.state.args = Namespace(
|
|
default_sampling_params='{"1": {"num_inference_steps":4, "guidance_scale":7.5, "generator_device":"cpu"}}',
|
|
max_generated_image_size=1048576,
|
|
)
|
|
client = TestClient(app)
|
|
|
|
response = client.post(
|
|
"/v1/images/generations",
|
|
json={
|
|
"prompt": "a coral reef",
|
|
"n": 1,
|
|
"size": "1024x1024",
|
|
"seed": 7,
|
|
},
|
|
)
|
|
assert response.status_code == 200
|
|
|
|
captured = engine.captured_sampling_params_list
|
|
assert captured is not None
|
|
assert len(captured) == 2
|
|
# max_tokens comes from YAML default, not computed dynamically
|
|
assert captured[0].max_tokens == 4353
|
|
assert captured[0].extra_args["target_h"] == 1024
|
|
assert captured[0].extra_args["target_w"] == 1024
|
|
assert captured[1].height == 1024
|
|
assert captured[1].width == 1024
|
|
|
|
|
|
def test_image_edits_async_omni_stage_configs_only(async_omni_stage_configs_only_client):
|
|
"""Regression: image edits accepts refactored AsyncOmni without stage_list."""
|
|
img_bytes = make_test_image_bytes((16, 16))
|
|
response = async_omni_stage_configs_only_client.post(
|
|
"/v1/images/edits",
|
|
files=[("image", img_bytes)],
|
|
data={
|
|
"prompt": "edit me",
|
|
"size": "auto",
|
|
},
|
|
)
|
|
assert response.status_code == 200
|
|
engine = async_omni_stage_configs_only_client.app.state.engine_client
|
|
captured = engine.captured_sampling_params_list
|
|
assert captured is not None
|
|
assert len(captured) == 2
|
|
|
|
|
|
def _parse_sse_payloads(body: str):
|
|
payloads = []
|
|
for line in body.splitlines():
|
|
if not line.startswith("data: "):
|
|
continue
|
|
data = line[len("data: ") :]
|
|
payloads.append(data if data == "[DONE]" else json.loads(data))
|
|
return payloads
|
|
|
|
|
|
def test_image_edits_streaming_returns_ar_delta_then_image(streaming_image_edit_client):
|
|
img_bytes = make_test_image_bytes((16, 16))
|
|
response = streaming_image_edit_client.post(
|
|
"/v1/images/edits",
|
|
files=[("image", img_bytes)],
|
|
data={
|
|
"prompt": "edit me",
|
|
"size": "auto",
|
|
"stream": "true",
|
|
"output_format": "png",
|
|
},
|
|
)
|
|
|
|
assert response.status_code == 200
|
|
assert response.headers["content-type"].startswith("text/event-stream")
|
|
|
|
payloads = _parse_sse_payloads(response.text)
|
|
assert [p["type"] if isinstance(p, dict) and "type" in p else p for p in payloads] == [
|
|
"ar_delta",
|
|
"ar_delta",
|
|
"image",
|
|
"[DONE]",
|
|
]
|
|
assert payloads[0]["delta"] == "recap"
|
|
assert payloads[0]["index"] == 0
|
|
assert payloads[1]["delta"] == " done"
|
|
assert payloads[2]["output_format"] == "png"
|
|
assert payloads[2]["size"] == "16x16"
|
|
|
|
image_payload = payloads[2]["data"][0]
|
|
img = Image.open(io.BytesIO(base64.b64decode(image_payload["b64_json"])))
|
|
assert img.size == (32, 24)
|
|
|
|
|
|
def test_image_edits_streaming_ar_delta_chunks_include_index(streaming_image_edit_client):
|
|
async def generate_multi_output_delta(
|
|
prompt, request_id, sampling_params=None, sampling_params_list=None, **kwargs
|
|
):
|
|
yield MockStageResult(stage_id=0, final_output_type="text", texts=["first", "second"])
|
|
yield MockStageResult(
|
|
stage_id=1,
|
|
final_output_type="image",
|
|
images=[Image.new("RGB", (32, 24), color="purple")],
|
|
)
|
|
|
|
streaming_image_edit_client.app.state.engine_client.generate = generate_multi_output_delta
|
|
img_bytes = make_test_image_bytes((16, 16))
|
|
response = streaming_image_edit_client.post(
|
|
"/v1/images/edits",
|
|
files=[("image", img_bytes)],
|
|
data={
|
|
"prompt": "edit me",
|
|
"stream": "true",
|
|
},
|
|
)
|
|
|
|
assert response.status_code == 200
|
|
payloads = _parse_sse_payloads(response.text)
|
|
assert [(payloads[0]["index"], payloads[0]["delta"]), (payloads[1]["index"], payloads[1]["delta"])] == [
|
|
(0, "first"),
|
|
(1, "second"),
|
|
]
|
|
|
|
|
|
def test_image_edits_streaming_errors_without_final_image(streaming_image_edit_client):
|
|
async def generate_without_image(prompt, request_id, sampling_params=None, sampling_params_list=None, **kwargs):
|
|
yield MockStageResult(stage_id=0, final_output_type="text", text="recap")
|
|
|
|
streaming_image_edit_client.app.state.engine_client.generate = generate_without_image
|
|
img_bytes = make_test_image_bytes((16, 16))
|
|
response = streaming_image_edit_client.post(
|
|
"/v1/images/edits",
|
|
files=[("image", img_bytes)],
|
|
data={
|
|
"prompt": "edit me",
|
|
"stream": "true",
|
|
},
|
|
)
|
|
|
|
assert response.status_code == 200
|
|
payloads = _parse_sse_payloads(response.text)
|
|
assert payloads[0]["type"] == "ar_delta"
|
|
assert payloads[1]["object"] == "error"
|
|
assert "without a final image" in payloads[1]["error"]["message"]
|
|
assert payloads[2] == "[DONE]"
|
|
|
|
|
|
def test_image_edits_streaming_errors_on_empty_final_image(streaming_image_edit_client):
|
|
async def generate_empty_image(prompt, request_id, sampling_params=None, sampling_params_list=None, **kwargs):
|
|
yield MockStageResult(stage_id=0, final_output_type="text", text="recap")
|
|
yield MockStageResult(stage_id=1, final_output_type="image", images=[])
|
|
|
|
streaming_image_edit_client.app.state.engine_client.generate = generate_empty_image
|
|
img_bytes = make_test_image_bytes((16, 16))
|
|
response = streaming_image_edit_client.post(
|
|
"/v1/images/edits",
|
|
files=[("image", img_bytes)],
|
|
data={
|
|
"prompt": "edit me",
|
|
"stream": "true",
|
|
},
|
|
)
|
|
|
|
assert response.status_code == 200
|
|
payloads = _parse_sse_payloads(response.text)
|
|
assert payloads[0]["type"] == "ar_delta"
|
|
assert payloads[1]["object"] == "error"
|
|
assert "empty final image" in payloads[1]["error"]["message"]
|
|
assert payloads[2] == "[DONE]"
|
|
|
|
|
|
def test_image_edits_streaming_guardrail_error_uses_400(streaming_image_edit_client):
|
|
async def generate_guardrail_error(prompt, request_id, sampling_params=None, sampling_params_list=None, **kwargs):
|
|
raise GuardrailViolationError("Input was blocked by Cosmos3 guardrails.")
|
|
yield MockStageResult(stage_id=1, final_output_type="image", images=[]) # pragma: no cover
|
|
|
|
streaming_image_edit_client.app.state.engine_client.generate = generate_guardrail_error
|
|
img_bytes = make_test_image_bytes((16, 16))
|
|
response = streaming_image_edit_client.post(
|
|
"/v1/images/edits",
|
|
files=[("image", img_bytes)],
|
|
data={
|
|
"prompt": "blocked prompt",
|
|
"stream": "true",
|
|
},
|
|
)
|
|
|
|
assert response.status_code == 200
|
|
payloads = _parse_sse_payloads(response.text)
|
|
assert payloads[0]["object"] == "error"
|
|
assert payloads[0]["error"]["message"] == "Input was blocked by Cosmos3 guardrails."
|
|
assert payloads[0]["error"]["type"] == "BadRequestError"
|
|
assert payloads[0]["error"]["code"] == 400
|
|
assert payloads[1] == "[DONE]"
|
|
|
|
|
|
def test_image_edits_streaming_rejects_single_stage(test_client):
|
|
img_bytes = make_test_image_bytes((16, 16))
|
|
response = test_client.post(
|
|
"/v1/images/edits",
|
|
files=[("image", img_bytes)],
|
|
data={
|
|
"prompt": "edit me",
|
|
"stream": "true",
|
|
},
|
|
)
|
|
|
|
assert response.status_code == 400
|
|
assert "multi-stage" in response.json()["detail"]
|
|
|
|
|
|
def test_image_edits_streaming_rejects_single_stage_before_loading_url(test_client):
|
|
response = test_client.post(
|
|
"/v1/images/edits",
|
|
data={
|
|
"prompt": "edit me",
|
|
"url": "https://example.invalid/not-fetched.png",
|
|
"stream": "true",
|
|
},
|
|
)
|
|
|
|
assert response.status_code == 400
|
|
assert "multi-stage" in response.json()["detail"]
|
|
|
|
|
|
def test_generate_images_max_size_rejected(async_omni_test_client):
|
|
"""Test that a size exceeding max_generated_image_size returns 400."""
|
|
response = async_omni_test_client.post(
|
|
"/v1/images/generations",
|
|
json={
|
|
"prompt": "a cat",
|
|
"size": "2048x2048", # 4,194,304 pixels > max_generated_image_size (1,048,576)
|
|
},
|
|
)
|
|
assert response.status_code == 400
|
|
|
|
|
|
def test_generate_multiple_images(test_client):
|
|
"""Test generating multiple images"""
|
|
response = test_client.post(
|
|
"/v1/images/generations",
|
|
json={
|
|
"prompt": "a dog",
|
|
"n": 3,
|
|
"size": "512x512",
|
|
},
|
|
)
|
|
assert response.status_code == 200
|
|
data = response.json()
|
|
assert len(data["data"]) == 3
|
|
|
|
# All images should be valid
|
|
for img_data in data["data"]:
|
|
assert "b64_json" in img_data
|
|
img_bytes = base64.b64decode(img_data["b64_json"])
|
|
img = Image.open(io.BytesIO(img_bytes))
|
|
assert img.format == "PNG"
|
|
|
|
|
|
def test_with_negative_prompt(test_client):
|
|
"""Test with negative prompt"""
|
|
response = test_client.post(
|
|
"/v1/images/generations",
|
|
json={
|
|
"prompt": "beautiful landscape",
|
|
"negative_prompt": "blurry, low quality",
|
|
"size": "1024x1024",
|
|
},
|
|
)
|
|
assert response.status_code == 200
|
|
|
|
|
|
def test_with_seed(test_client):
|
|
"""Test with seed for reproducibility"""
|
|
response = test_client.post(
|
|
"/v1/images/generations",
|
|
json={
|
|
"prompt": "a tree",
|
|
"seed": 42,
|
|
"size": "1024x1024",
|
|
},
|
|
)
|
|
assert response.status_code == 200
|
|
|
|
|
|
def test_with_seed_zero(test_client):
|
|
"""Test with seed=0 for reproducibility"""
|
|
response = test_client.post(
|
|
"/v1/images/generations",
|
|
json={
|
|
"prompt": "a tree",
|
|
"seed": 0,
|
|
"size": "1024x1024",
|
|
},
|
|
)
|
|
assert response.status_code == 200
|
|
engine = test_client.app.state.engine_client
|
|
captured = engine.captured_sampling_params_list[0]
|
|
# Verify that seed=0 is correctly passed
|
|
assert captured.seed == 0, (
|
|
f"Expected seed=0, but got seed={captured.seed}. This indicates the bug where seed=0 is treated as falsy."
|
|
)
|
|
|
|
|
|
def test_with_custom_parameters(test_client):
|
|
"""Test with custom diffusion parameters"""
|
|
response = test_client.post(
|
|
"/v1/images/generations",
|
|
json={
|
|
"prompt": "a mountain",
|
|
"size": "1024x1024",
|
|
"num_inference_steps": 100,
|
|
"true_cfg_scale": 5.5,
|
|
"seed": 123,
|
|
},
|
|
)
|
|
assert response.status_code == 200
|
|
|
|
|
|
def test_flow_shift_forwarded_to_extra_args(test_client):
|
|
"""flow_shift must reach the diffusion sampling params via extra_args.
|
|
|
|
Regression: ``ImageGenerationRequest`` had no ``flow_shift`` field and the
|
|
single-stage handler never forwarded it, so a request like
|
|
``{"flow_shift": 10.0}`` was silently dropped and Cosmos3 T2I always ran at
|
|
its hardcoded per-mode default shift. The pipeline reads
|
|
``extra_args["flow_shift"]`` (via ``_get_sp_param``), so it must land there.
|
|
"""
|
|
response = test_client.post(
|
|
"/v1/images/generations",
|
|
json={
|
|
"prompt": "a robot in a lab",
|
|
"size": "960x960",
|
|
"num_inference_steps": 50,
|
|
"guidance_scale": 4.0,
|
|
"flow_shift": 10.0,
|
|
},
|
|
)
|
|
assert response.status_code == 200
|
|
captured = test_client.app.state.engine_client.captured_sampling_params_list[0]
|
|
assert captured.extra_args["flow_shift"] == 10.0
|
|
|
|
|
|
def test_flow_shift_absent_when_not_requested(test_client):
|
|
"""Omitting flow_shift must not inject an override, so the pipeline keeps
|
|
its per-mode default (e.g. Cosmos3 T2I shift=3.0)."""
|
|
response = test_client.post(
|
|
"/v1/images/generations",
|
|
json={"prompt": "a tree", "size": "1024x1024"},
|
|
)
|
|
assert response.status_code == 200
|
|
captured = test_client.app.state.engine_client.captured_sampling_params_list[0]
|
|
assert "flow_shift" not in (captured.extra_args or {})
|
|
|
|
|
|
def test_invalid_size(test_client):
|
|
"""Test with invalid size parameter - rejected by Pydantic"""
|
|
response = test_client.post(
|
|
"/v1/images/generations",
|
|
json={
|
|
"prompt": "a cat",
|
|
"size": "invalid",
|
|
},
|
|
)
|
|
# Pydantic validation errors return 422 (Unprocessable Entity)
|
|
# "invalid" has no "x" so Pydantic rejects it
|
|
assert response.status_code == 422
|
|
# Check error detail contains size validation message
|
|
detail = str(response.json()["detail"])
|
|
assert "size" in detail.lower() or "invalid" in detail.lower()
|
|
|
|
|
|
def test_invalid_size_parse_error(test_client):
|
|
"""Test with malformed size - passes Pydantic but fails parse_size()"""
|
|
response = test_client.post(
|
|
"/v1/images/generations",
|
|
json={
|
|
"prompt": "a cat",
|
|
"size": "1024x", # Has "x" so Pydantic accepts, but parse_size() rejects
|
|
},
|
|
)
|
|
# parse_size() raises ValueError → endpoint converts to 400 (Bad Request)
|
|
assert response.status_code == 400
|
|
detail = str(response.json()["detail"])
|
|
assert "size" in detail.lower() or "invalid" in detail.lower()
|
|
|
|
|
|
def test_missing_prompt(test_client):
|
|
"""Test with missing required prompt field"""
|
|
response = test_client.post(
|
|
"/v1/images/generations",
|
|
json={
|
|
"size": "1024x1024",
|
|
},
|
|
)
|
|
# Pydantic validation error
|
|
assert response.status_code == 422
|
|
|
|
|
|
def test_invalid_n_parameter(test_client):
|
|
"""Test with invalid n parameter (out of range)"""
|
|
# n < 1
|
|
response = test_client.post(
|
|
"/v1/images/generations",
|
|
json={
|
|
"prompt": "a cat",
|
|
"n": 0,
|
|
},
|
|
)
|
|
assert response.status_code == 422
|
|
|
|
# n > 10
|
|
response = test_client.post(
|
|
"/v1/images/generations",
|
|
json={
|
|
"prompt": "a cat",
|
|
"n": 11,
|
|
},
|
|
)
|
|
assert response.status_code == 422
|
|
|
|
|
|
def test_url_response_format_not_supported(test_client):
|
|
"""Test that URL format returns error"""
|
|
response = test_client.post(
|
|
"/v1/images/generations",
|
|
json={
|
|
"prompt": "a cat",
|
|
"response_format": "url",
|
|
},
|
|
)
|
|
# Pydantic validation errors return 422 (Unprocessable Entity)
|
|
assert response.status_code == 422
|
|
# Check error mentions response_format or b64_json
|
|
detail = str(response.json()["detail"])
|
|
assert "b64_json" in detail.lower() or "response" in detail.lower()
|
|
|
|
|
|
def test_model_not_loaded():
|
|
"""Test error when diffusion engine is not initialized"""
|
|
from fastapi import FastAPI
|
|
|
|
from vllm_omni.entrypoints.openai.api_server import router
|
|
|
|
app = FastAPI()
|
|
app.include_router(router)
|
|
# Don't set diffusion_engine to simulate uninitialized state
|
|
app.state.diffusion_engine = None
|
|
|
|
client = TestClient(app)
|
|
response = client.post(
|
|
"/v1/images/generations",
|
|
json={
|
|
"prompt": "a cat",
|
|
},
|
|
)
|
|
assert response.status_code == 503
|
|
assert "not initialized" in response.json()["detail"].lower()
|
|
|
|
|
|
def test_different_image_sizes(test_client):
|
|
"""Test various valid image sizes"""
|
|
sizes = ["256x256", "512x512", "1024x1024", "1792x1024", "1024x1792"]
|
|
|
|
for size in sizes:
|
|
response = test_client.post(
|
|
"/v1/images/generations",
|
|
json={
|
|
"prompt": "a test image",
|
|
"size": size,
|
|
},
|
|
)
|
|
assert response.status_code == 200, f"Failed for size {size}"
|
|
|
|
|
|
def test_parameter_validation():
|
|
"""Test Pydantic model validation"""
|
|
from vllm_omni.entrypoints.openai.protocol.images import ImageGenerationRequest
|
|
|
|
# Valid request - optional parameters default to None
|
|
req = ImageGenerationRequest(prompt="test")
|
|
assert req.prompt == "test"
|
|
assert req.n == 1
|
|
assert req.model is None
|
|
assert req.size is None # Engine will use model defaults
|
|
assert req.num_inference_steps is None # Engine will use model defaults
|
|
assert req.true_cfg_scale is None # Engine will use model defaults
|
|
|
|
# Invalid num_inference_steps (out of range)
|
|
with pytest.raises(ValueError):
|
|
ImageGenerationRequest(prompt="test", num_inference_steps=0)
|
|
|
|
with pytest.raises(ValueError):
|
|
ImageGenerationRequest(prompt="test", num_inference_steps=201)
|
|
|
|
# Invalid guidance_scale (out of range)
|
|
with pytest.raises(ValueError):
|
|
ImageGenerationRequest(prompt="test", guidance_scale=-1.0)
|
|
|
|
with pytest.raises(ValueError):
|
|
ImageGenerationRequest(prompt="test", guidance_scale=21.0)
|
|
|
|
# Invalid layers for layered models (must stay within the backend-supported range)
|
|
with pytest.raises(ValueError):
|
|
ImageGenerationRequest(prompt="test", layers=1)
|
|
|
|
with pytest.raises(ValueError):
|
|
ImageGenerationRequest(prompt="test", layers=11)
|
|
|
|
|
|
# Pass-Through Tests
|
|
|
|
|
|
def test_parameters_passed_through(test_client, mock_async_diffusion):
|
|
"""Verify all parameters passed through without modification"""
|
|
response = test_client.post(
|
|
"/v1/images/generations",
|
|
json={
|
|
"prompt": "test",
|
|
"num_inference_steps": 100,
|
|
"guidance_scale": 7.5,
|
|
"true_cfg_scale": 3.0,
|
|
"seed": 42,
|
|
},
|
|
)
|
|
assert response.status_code == 200
|
|
|
|
assert mock_async_diffusion.generate_calls == 1
|
|
captured = mock_async_diffusion.captured_sampling_params_list[0]
|
|
assert captured.num_inference_steps == 100
|
|
assert captured.guidance_scale == 7.5
|
|
assert captured.true_cfg_scale == 3.0
|
|
assert captured.seed == 42
|
|
|
|
|
|
def test_model_field_omitted_works(test_client):
|
|
"""Test that omitting model field works"""
|
|
response = test_client.post(
|
|
"/v1/images/generations",
|
|
json={
|
|
"prompt": "test",
|
|
"size": "1024x1024",
|
|
# model field omitted
|
|
},
|
|
)
|
|
assert response.status_code == 200
|
|
|
|
|
|
def test_generate_images_rejects_model_mismatch(test_client):
|
|
response = test_client.post(
|
|
"/v1/images/generations",
|
|
json={
|
|
"prompt": "test",
|
|
"model": "Qwen/Qwen-Image-2512",
|
|
"size": "1024x1024",
|
|
},
|
|
)
|
|
assert response.status_code == 400
|
|
assert "model mismatch" in response.json()["detail"].lower()
|
|
|
|
|
|
def test_image_file_response_format_multiple(test_client):
|
|
"""Test response_format=file with n>1 returns ZIP archive"""
|
|
response = test_client.post(
|
|
"/v1/images/generations",
|
|
json={
|
|
"prompt": "a dog",
|
|
"n": 3,
|
|
"response_format": "file",
|
|
},
|
|
)
|
|
|
|
assert response.status_code == 200
|
|
assert response.headers["content-type"] == "application/zip"
|
|
assert "attachment" in response.headers.get("content-disposition", "")
|
|
assert ".zip" in response.headers.get("content-disposition", "")
|
|
|
|
# Verify it's a valid ZIP with 3 PNG files
|
|
import zipfile
|
|
|
|
zip_buffer = io.BytesIO(response.content)
|
|
with zipfile.ZipFile(zip_buffer, "r") as zf:
|
|
files = zf.namelist()
|
|
assert len(files) == 3
|
|
assert all(f.endswith(".png") for f in files)
|
|
|
|
# Verify each file is a valid PNG
|
|
for filename in files:
|
|
img_bytes = zf.read(filename)
|
|
img = Image.open(io.BytesIO(img_bytes))
|
|
assert img.format == "PNG"
|
|
|
|
|
|
def test_image_file_response_format_single(test_client):
|
|
"""Test response_format=file with n=1 returns a single image file."""
|
|
response = test_client.post(
|
|
"/v1/images/generations",
|
|
json={
|
|
"prompt": "a dog",
|
|
"n": 1,
|
|
"response_format": "file",
|
|
},
|
|
)
|
|
|
|
assert response.status_code == 200
|
|
assert response.headers["content-type"] == "image/png"
|
|
assert "attachment" in response.headers.get("content-disposition", "")
|
|
assert ".png" in response.headers.get("content-disposition", "")
|
|
|
|
img = Image.open(io.BytesIO(response.content))
|
|
assert img.format == "PNG"
|
|
|
|
|
|
def make_test_image_bytes(size=(64, 64)) -> bytes:
|
|
img = Image.new(
|
|
"RGB",
|
|
size,
|
|
)
|
|
buf = io.BytesIO()
|
|
img.save(buf, format="PNG")
|
|
return buf.getvalue()
|
|
|
|
|
|
def test_image_edit_images_processing(async_omni_test_client):
|
|
img_bytes_1 = make_test_image_bytes((16, 16))
|
|
img_bytes_2 = make_test_image_bytes((32, 32))
|
|
|
|
# uploadfile with image key
|
|
response = async_omni_test_client.post(
|
|
"/v1/images/edits",
|
|
files=[
|
|
("image", img_bytes_1),
|
|
("image", img_bytes_2),
|
|
],
|
|
data={"prompt": "hello world."},
|
|
)
|
|
assert response.status_code == 200
|
|
engine = async_omni_test_client.app.state.engine_client
|
|
captured_prompt = engine.captured_prompt
|
|
processed_images = captured_prompt["multi_modal_data"]["image"]
|
|
assert len(processed_images) == 2
|
|
assert isinstance(processed_images[0], Image.Image)
|
|
assert isinstance(processed_images[1], Image.Image)
|
|
assert processed_images[0].size == (16, 16)
|
|
assert processed_images[1].size == (32, 32)
|
|
|
|
# uploadfile with image[] key
|
|
response = async_omni_test_client.post(
|
|
"/v1/images/edits",
|
|
files=[
|
|
("image[]", img_bytes_2),
|
|
("image[]", img_bytes_1),
|
|
],
|
|
data={"prompt": "hello world."},
|
|
)
|
|
|
|
assert response.status_code == 200
|
|
engine = async_omni_test_client.app.state.engine_client
|
|
captured_prompt = engine.captured_prompt
|
|
processed_images = captured_prompt["multi_modal_data"]["image"]
|
|
assert len(processed_images) == 2
|
|
assert isinstance(processed_images[0], Image.Image)
|
|
assert isinstance(processed_images[1], Image.Image)
|
|
assert processed_images[0].size == (32, 32)
|
|
assert processed_images[1].size == (16, 16)
|
|
|
|
# base64 url
|
|
buf1 = io.BytesIO()
|
|
img1 = Image.new("RGB", (16, 16))
|
|
img1.save(buf1, format="PNG")
|
|
b64_1 = "data:image/png;base64," + base64.b64encode(buf1.getvalue()).decode()
|
|
|
|
buf2 = io.BytesIO()
|
|
img2 = Image.new("RGB", (24, 24))
|
|
img2.save(buf2, format="PNG")
|
|
b64_2 = "data:image/png;base64," + base64.b64encode(buf2.getvalue()).decode()
|
|
|
|
response = async_omni_test_client.post(
|
|
"/v1/images/edits",
|
|
data={
|
|
"prompt": "hello from base64",
|
|
"url": [b64_1, b64_2],
|
|
},
|
|
)
|
|
assert response.status_code == 200
|
|
processed_images = engine.captured_prompt["multi_modal_data"]["image"]
|
|
assert len(processed_images) == 2
|
|
assert isinstance(processed_images[0], Image.Image)
|
|
assert isinstance(processed_images[1], Image.Image)
|
|
assert processed_images[0].size == (16, 16)
|
|
assert processed_images[1].size == (24, 24)
|
|
|
|
|
|
def test_image_edit_rejects_multiple_images_when_model_does_not_support_them(async_omni_test_client):
|
|
img_bytes_1 = make_test_image_bytes((16, 16))
|
|
img_bytes_2 = make_test_image_bytes((32, 32))
|
|
|
|
engine = async_omni_test_client.app.state.engine_client
|
|
engine.get_diffusion_od_config = lambda: SimpleNamespace(supports_multimodal_inputs=False)
|
|
|
|
response = async_omni_test_client.post(
|
|
"/v1/images/edits",
|
|
files=[
|
|
("image", img_bytes_1),
|
|
("image", img_bytes_2),
|
|
],
|
|
data={"prompt": "hello world."},
|
|
)
|
|
|
|
assert response.status_code == 400
|
|
assert (
|
|
response.json()["detail"] == "Received multiple input images. Only a single image is supported by this model."
|
|
)
|
|
assert engine.captured_prompt is None
|
|
|
|
|
|
def test_image_edit_rejects_model_mismatch(test_client):
|
|
img_bytes = make_test_image_bytes((16, 16))
|
|
response = test_client.post(
|
|
"/v1/images/edits",
|
|
files=[("image", img_bytes)],
|
|
data={
|
|
"prompt": "edit me",
|
|
"model": "Qwen/Qwen-Image-Edit",
|
|
},
|
|
)
|
|
assert response.status_code == 400
|
|
assert "model mismatch" in response.json()["detail"].lower()
|
|
|
|
|
|
def test_image_edit_rejects_too_many_images_for_qwen_image_edit_2511(async_omni_test_client):
|
|
engine = async_omni_test_client.app.state.engine_client
|
|
engine.get_diffusion_od_config = lambda: SimpleNamespace(
|
|
supports_multimodal_inputs=True,
|
|
max_multimodal_image_inputs=4,
|
|
)
|
|
|
|
response = async_omni_test_client.post(
|
|
"/v1/images/edits",
|
|
files=[
|
|
("image", make_test_image_bytes((16, 16))),
|
|
("image", make_test_image_bytes((16, 16))),
|
|
("image", make_test_image_bytes((16, 16))),
|
|
("image", make_test_image_bytes((16, 16))),
|
|
("image", make_test_image_bytes((16, 16))),
|
|
],
|
|
data={"prompt": "hello world."},
|
|
)
|
|
|
|
assert response.status_code == 400
|
|
assert response.json()["detail"] == "Received 5 input images. At most 4 images are supported by this model."
|
|
assert engine.captured_prompt is None
|
|
|
|
|
|
def test_image_edit_rejects_too_many_images_for_qwen_image_edit_2511_before_loading(
|
|
async_omni_test_client, monkeypatch: pytest.MonkeyPatch
|
|
):
|
|
import vllm_omni.entrypoints.openai.api_server as api_server_module
|
|
|
|
engine = async_omni_test_client.app.state.engine_client
|
|
engine.get_diffusion_od_config = lambda: SimpleNamespace(
|
|
supports_multimodal_inputs=True,
|
|
max_multimodal_image_inputs=4,
|
|
)
|
|
|
|
def _fail_load(*args, **kwargs):
|
|
raise AssertionError("_load_input_images should not run for over-limit requests")
|
|
|
|
monkeypatch.setattr(api_server_module, "_load_input_images", _fail_load)
|
|
|
|
response = async_omni_test_client.post(
|
|
"/v1/images/edits",
|
|
files=[
|
|
("image", make_test_image_bytes((16, 16))),
|
|
("image", make_test_image_bytes((16, 16))),
|
|
("image", make_test_image_bytes((16, 16))),
|
|
("image", make_test_image_bytes((16, 16))),
|
|
("image", make_test_image_bytes((16, 16))),
|
|
],
|
|
data={"prompt": "hello world."},
|
|
)
|
|
|
|
assert response.status_code == 400
|
|
assert response.json()["detail"] == "Received 5 input images. At most 4 images are supported by this model."
|
|
assert engine.captured_prompt is None
|
|
|
|
|
|
def test_image_edit_ignores_mock_like_multimodal_limit(async_omni_test_client):
|
|
engine = async_omni_test_client.app.state.engine_client
|
|
engine.get_diffusion_od_config = lambda: SimpleNamespace(
|
|
supports_multimodal_inputs=SimpleNamespace(),
|
|
max_multimodal_image_inputs=SimpleNamespace(),
|
|
)
|
|
|
|
response = async_omni_test_client.post(
|
|
"/v1/images/edits",
|
|
files=[("image", make_test_image_bytes((16, 16)))],
|
|
data={"prompt": "hello world."},
|
|
)
|
|
|
|
assert response.status_code == 200
|
|
captured_prompt = engine.captured_prompt
|
|
assert captured_prompt is not None
|
|
# Multi-stage path uses "img2img" key for single reference image
|
|
processed_images = captured_prompt["multi_modal_data"]["img2img"]
|
|
assert isinstance(processed_images, Image.Image)
|
|
assert processed_images.size == (16, 16)
|
|
|
|
|
|
def test_image_edit_parameter_pass(async_omni_test_client):
|
|
img_bytes_1 = make_test_image_bytes((16, 16))
|
|
|
|
# uploadfile with image key
|
|
response = async_omni_test_client.post(
|
|
"/v1/images/edits",
|
|
files=[("image", img_bytes_1)],
|
|
data={
|
|
"prompt": "hello world.",
|
|
"size": "16x24",
|
|
"output_format": "jpeg",
|
|
"num_inference_steps": 20,
|
|
"guidance_scale": 8.0,
|
|
"seed": 1234,
|
|
"negative_prompt": "negative",
|
|
"n": 2,
|
|
},
|
|
)
|
|
assert response.status_code == 200
|
|
engine = async_omni_test_client.app.state.engine_client
|
|
captured_prompt = engine.captured_prompt
|
|
captured_sampling_params = engine.captured_sampling_params_list[-1]
|
|
|
|
assert captured_prompt["prompt"] == "hello world."
|
|
assert captured_prompt["negative_prompt"] == "negative"
|
|
assert captured_sampling_params.num_inference_steps == 20
|
|
assert captured_sampling_params.guidance_scale == 8.0
|
|
assert captured_sampling_params.seed == 1234
|
|
assert captured_sampling_params.num_outputs_per_prompt == 2
|
|
assert captured_sampling_params.width == 16
|
|
assert captured_sampling_params.height == 24
|
|
|
|
data = response.json()
|
|
# All images should be valid
|
|
for img_data in data["data"]:
|
|
assert "b64_json" in img_data
|
|
img_bytes = base64.b64decode(img_data["b64_json"])
|
|
img = Image.open(io.BytesIO(img_bytes))
|
|
assert img.format.lower() == "jpeg"
|
|
assert data["output_format"] == "jpeg"
|
|
assert data["size"] == "16x24"
|
|
|
|
|
|
def test_image_edit_layers_and_resolution(async_omni_test_client):
|
|
"""Test layers and resolution parameters for layered models."""
|
|
img_bytes = make_test_image_bytes((16, 16))
|
|
|
|
response = async_omni_test_client.post(
|
|
"/v1/images/edits",
|
|
files=[("image", img_bytes)],
|
|
data={
|
|
"prompt": "decompose into layers",
|
|
"layers": 4,
|
|
"resolution": 1024,
|
|
},
|
|
)
|
|
assert response.status_code == 200
|
|
engine = async_omni_test_client.app.state.engine_client
|
|
captured_sampling_params = engine.captured_sampling_params_list[-1]
|
|
assert captured_sampling_params.layers == 4
|
|
assert captured_sampling_params.resolution == 1024
|
|
|
|
|
|
def test_image_edit_resolution_auto_size(async_omni_test_client):
|
|
"""Test that size='auto' with resolution lets pipeline calculate dimensions."""
|
|
img_bytes = make_test_image_bytes((16, 16))
|
|
|
|
response = async_omni_test_client.post(
|
|
"/v1/images/edits",
|
|
files=[("image", img_bytes)],
|
|
data={
|
|
"prompt": "test",
|
|
"size": "auto",
|
|
"resolution": 640,
|
|
},
|
|
)
|
|
assert response.status_code == 200
|
|
engine = async_omni_test_client.app.state.engine_client
|
|
captured_sampling_params = engine.captured_sampling_params_list[-1]
|
|
# When resolution is set with size=auto, width/height should be None
|
|
# to let pipeline calculate based on resolution
|
|
assert captured_sampling_params.width is None
|
|
assert captured_sampling_params.height is None
|
|
assert captured_sampling_params.resolution == 640
|
|
|
|
|
|
def test_image_edit_invalid_resolution(async_omni_test_client):
|
|
"""Test that invalid resolution values are rejected with 400."""
|
|
img_bytes = make_test_image_bytes((16, 16))
|
|
|
|
response = async_omni_test_client.post(
|
|
"/v1/images/edits",
|
|
files=[("image", img_bytes)],
|
|
data={
|
|
"prompt": "test",
|
|
"resolution": 512, # Invalid, only 640 or 1024 are supported
|
|
},
|
|
)
|
|
assert response.status_code == 400
|
|
detail = response.json()["detail"]
|
|
assert "Invalid resolution" in detail
|
|
assert "512" in detail
|
|
|
|
|
|
def test_image_edit_invalid_layers(async_omni_test_client):
|
|
"""Test that layered image edits reject out-of-range layers with 400."""
|
|
img_bytes = make_test_image_bytes((16, 16))
|
|
|
|
# Test layers below the supported range
|
|
response = async_omni_test_client.post(
|
|
"/v1/images/edits",
|
|
files=[("image", img_bytes)],
|
|
data={
|
|
"prompt": "test",
|
|
"layers": 1,
|
|
},
|
|
)
|
|
assert response.status_code == 400
|
|
detail = response.json()["detail"]
|
|
assert "Invalid layers" in detail
|
|
assert "layers must be between 2 and 10 inclusive" in detail
|
|
|
|
# Test layers above the supported range
|
|
response = async_omni_test_client.post(
|
|
"/v1/images/edits",
|
|
files=[("image", img_bytes)],
|
|
data={
|
|
"prompt": "test",
|
|
"layers": 11,
|
|
},
|
|
)
|
|
assert response.status_code == 400
|
|
detail = response.json()["detail"]
|
|
assert "Invalid layers" in detail
|
|
assert "layers must be between 2 and 10 inclusive" in detail
|
|
|
|
|
|
def test_image_edit_resolution_and_size_conflict(async_omni_test_client):
|
|
"""Test that providing both resolution and explicit size raises 400."""
|
|
img_bytes = make_test_image_bytes((16, 16))
|
|
|
|
response = async_omni_test_client.post(
|
|
"/v1/images/edits",
|
|
files=[("image", img_bytes)],
|
|
data={
|
|
"prompt": "test",
|
|
"resolution": 1024,
|
|
"size": "512x512", # Conflict: both resolution and explicit size
|
|
},
|
|
)
|
|
assert response.status_code == 400
|
|
detail = response.json()["detail"]
|
|
assert "Cannot specify both" in detail
|
|
assert "resolution" in detail
|
|
assert "size" in detail
|
|
|
|
|
|
def test_image_edit_parameter_default(async_omni_test_client):
|
|
img_bytes_1 = make_test_image_bytes((24, 16))
|
|
|
|
# uploadfile with image key
|
|
response = async_omni_test_client.post(
|
|
"/v1/images/edits",
|
|
files=[("image", img_bytes_1)],
|
|
data={
|
|
"prompt": "hello world.",
|
|
"size": "auto",
|
|
},
|
|
)
|
|
assert response.status_code == 200
|
|
engine = async_omni_test_client.app.state.engine_client
|
|
captured_sampling_params = engine.captured_sampling_params_list[-1]
|
|
|
|
# size="auto" on multi-stage pipelines deliberately leaves the diffusion
|
|
# stages sampling_params width/height unset so AR-driven pipelines (e.g.
|
|
# HunyuanImage-3.0) can let ar2diffusion override the final bucket from
|
|
# the AR-predicted ratio token; see
|
|
# test_image_edits_size_auto_preserves_bridge_size for the contract.
|
|
# Single-stage diffusion (test_image_edit_parameter_default_single_stage)
|
|
# still pins width/height to the input image size via api_servers
|
|
# gen_params, which is unchanged.
|
|
assert captured_sampling_params.width is None
|
|
assert captured_sampling_params.height is None
|
|
assert captured_sampling_params.num_outputs_per_prompt == 1
|
|
assert captured_sampling_params.num_inference_steps == 4
|
|
assert captured_sampling_params.guidance_scale == 7.5
|
|
assert captured_sampling_params.generator_device == "cpu"
|
|
|
|
# Test that a size exceeding max_generated_image_size returns 400
|
|
response = async_omni_test_client.post(
|
|
"/v1/images/edits",
|
|
files=[("image", img_bytes_1)],
|
|
data={
|
|
"prompt": "hello world.",
|
|
"size": "2048x2048", # 4,194,304 pixels > max_generated_image_size (1,048,576)
|
|
},
|
|
)
|
|
assert response.status_code == 400
|
|
|
|
|
|
def test_image_edit_parameter_default_single_stage(test_client):
|
|
img_bytes_1 = make_test_image_bytes((24, 16))
|
|
|
|
# uploadfile with image key
|
|
response = test_client.post(
|
|
"/v1/images/edits",
|
|
files=[("image", img_bytes_1)],
|
|
data={
|
|
"prompt": "hello world.",
|
|
},
|
|
)
|
|
assert response.status_code == 200
|
|
engine = test_client.app.state.engine_client
|
|
captured_sampling_params = engine.captured_sampling_params_list[0]
|
|
|
|
assert captured_sampling_params.width == 24
|
|
assert captured_sampling_params.height == 16
|
|
assert captured_sampling_params.num_outputs_per_prompt == 1
|
|
assert captured_sampling_params.num_inference_steps == 4
|
|
assert captured_sampling_params.guidance_scale == 7.5
|
|
assert captured_sampling_params.generator_device == "cpu"
|
|
|
|
# Size exceeding max_generated_image_size (1024*1792) returns 400
|
|
response = test_client.post(
|
|
"/v1/images/edits",
|
|
files=[("image", img_bytes_1)],
|
|
data={
|
|
"prompt": "hello world.",
|
|
"size": "2048x2048",
|
|
},
|
|
)
|
|
assert response.status_code == 400
|
|
|
|
|
|
def test_image_edit_compression_jpeg(test_client):
|
|
img_bytes_1 = make_test_image_bytes((16, 16))
|
|
# uploadfile with image key
|
|
response = test_client.post(
|
|
"/v1/images/edits",
|
|
files=[("image", img_bytes_1)],
|
|
data={"prompt": "hello world.", "output_format": "jpeg", "output_compression": 100},
|
|
)
|
|
assert response.status_code == 200
|
|
data = response.json()
|
|
img_bytes_100 = base64.b64decode(data["data"][0]["b64_json"])
|
|
img = Image.open(io.BytesIO(img_bytes_100))
|
|
assert img.format.lower() == "jpeg"
|
|
|
|
response = test_client.post(
|
|
"/v1/images/edits",
|
|
files=[("image", img_bytes_1)],
|
|
data={
|
|
"prompt": "hello world.",
|
|
"output_format": "jpeg",
|
|
"output_compression": 50,
|
|
},
|
|
)
|
|
assert response.status_code == 200
|
|
data = response.json()
|
|
img_bytes_50 = base64.b64decode(data["data"][0]["b64_json"])
|
|
|
|
response = test_client.post(
|
|
"/v1/images/edits",
|
|
files=[("image", img_bytes_1)],
|
|
data={
|
|
"prompt": "hello world.",
|
|
"output_format": "jpeg",
|
|
"output_compression": 10,
|
|
},
|
|
)
|
|
assert response.status_code == 200
|
|
data = response.json()
|
|
img_bytes_10 = base64.b64decode(data["data"][0]["b64_json"])
|
|
|
|
assert len(img_bytes_10) < len(img_bytes_50)
|
|
assert len(img_bytes_50) < len(img_bytes_100)
|
|
|
|
|
|
def test_image_edit_rgba_output_converts_to_jpeg(async_omni_rgba_test_client):
|
|
img_bytes_1 = make_test_image_bytes((16, 16))
|
|
|
|
response = async_omni_rgba_test_client.post(
|
|
"/v1/images/edits",
|
|
files=[("image", img_bytes_1)],
|
|
data={
|
|
"prompt": "hello world.",
|
|
"output_format": "jpeg",
|
|
},
|
|
)
|
|
assert response.status_code == 200
|
|
|
|
data = response.json()
|
|
img_bytes = base64.b64decode(data["data"][0]["b64_json"])
|
|
img = Image.open(io.BytesIO(img_bytes))
|
|
assert img.format.lower() == "jpeg"
|
|
assert img.mode == "RGB"
|
|
assert data["output_format"] == "jpeg"
|
|
|
|
|
|
def test_image_edit_compression_png(async_omni_test_client):
|
|
img_bytes_1 = make_test_image_bytes((16, 16))
|
|
# uploadfile with image key
|
|
response = async_omni_test_client.post(
|
|
"/v1/images/edits",
|
|
files=[("image", img_bytes_1)],
|
|
data={"prompt": "hello world.", "output_format": "PNG", "output_compression": 100},
|
|
)
|
|
assert response.status_code == 200
|
|
data = response.json()
|
|
img_bytes_100 = base64.b64decode(data["data"][0]["b64_json"])
|
|
img = Image.open(io.BytesIO(img_bytes_100))
|
|
assert img.format.lower() == "png"
|
|
|
|
response = async_omni_test_client.post(
|
|
"/v1/images/edits",
|
|
files=[("image", img_bytes_1)],
|
|
data={
|
|
"prompt": "hello world.",
|
|
"output_format": "PNG",
|
|
"output_compression": 50,
|
|
},
|
|
)
|
|
assert response.status_code == 200
|
|
data = response.json()
|
|
img_bytes_50 = base64.b64decode(data["data"][0]["b64_json"])
|
|
|
|
response = async_omni_test_client.post(
|
|
"/v1/images/edits",
|
|
files=[("image", img_bytes_1)],
|
|
data={
|
|
"prompt": "hello world.",
|
|
"output_format": "PNG",
|
|
"output_compression": 10,
|
|
},
|
|
)
|
|
assert response.status_code == 200
|
|
data = response.json()
|
|
img_bytes_10 = base64.b64decode(data["data"][0]["b64_json"])
|
|
|
|
assert len(img_bytes_10) < len(img_bytes_50)
|
|
assert len(img_bytes_50) < len(img_bytes_100)
|
|
|
|
|
|
def test_image_edit_with_seed_zero(async_omni_test_client):
|
|
"""Test that seed=0 is correctly handled in image editing.
|
|
|
|
Previously, seed=0 was incorrectly replaced by a random seed due to the
|
|
falsy value check using `or` operator. This test ensures seed=0 is
|
|
properly passed through to the sampling parameters in image editing.
|
|
"""
|
|
img_bytes_1 = make_test_image_bytes((16, 16))
|
|
|
|
response = async_omni_test_client.post(
|
|
"/v1/images/edits",
|
|
files=[("image", img_bytes_1)],
|
|
data={
|
|
"prompt": "edit this image",
|
|
"seed": 0,
|
|
},
|
|
)
|
|
assert response.status_code == 200
|
|
engine = async_omni_test_client.app.state.engine_client
|
|
captured_sampling_params = engine.captured_sampling_params_list[-1]
|
|
# Verify that seed=0 is correctly passed
|
|
assert captured_sampling_params.seed == 0, (
|
|
f"Expected seed=0, but got seed={captured_sampling_params.seed}. "
|
|
"This indicates the bug where seed=0 is treated as falsy."
|
|
)
|
|
|
|
|
|
def test_image_edit_with_seed_zero_single_stage(test_client):
|
|
"""Test that seed=0 is correctly handled in image editing (single stage).
|
|
|
|
Test seed=0 handling in image editing with single stage path.
|
|
"""
|
|
img_bytes_1 = make_test_image_bytes((16, 16))
|
|
|
|
response = test_client.post(
|
|
"/v1/images/edits",
|
|
files=[("image", img_bytes_1)],
|
|
data={
|
|
"prompt": "edit this image",
|
|
"seed": 0,
|
|
},
|
|
)
|
|
assert response.status_code == 200
|
|
engine = test_client.app.state.engine_client
|
|
captured_sampling_params = engine.captured_sampling_params_list[0]
|
|
# Verify that seed=0 is correctly passed
|
|
assert captured_sampling_params.seed == 0, (
|
|
f"Expected seed=0, but got seed={captured_sampling_params.seed}. "
|
|
"This indicates the bug where seed=0 is treated as falsy."
|
|
)
|
|
|
|
|
|
def test_normalize_image():
|
|
"""Test _normalize_image with various input types"""
|
|
import numpy as np
|
|
|
|
from vllm_omni.entrypoints.openai.api_server import _normalize_image
|
|
|
|
# Test PIL Image input
|
|
img = Image.new("RGB", (64, 64), color="red")
|
|
result = _normalize_image(img)
|
|
assert isinstance(result, Image.Image)
|
|
assert result.size == (64, 64)
|
|
|
|
# Test uint8 numpy array
|
|
arr = np.random.randint(0, 255, (64, 64, 3), dtype=np.uint8)
|
|
result = _normalize_image(arr)
|
|
assert isinstance(result, Image.Image)
|
|
assert result.size == (64, 64)
|
|
|
|
# Test float [0, 1] numpy array
|
|
arr = np.random.rand(64, 64, 3).astype(np.float32)
|
|
result = _normalize_image(arr)
|
|
assert isinstance(result, Image.Image)
|
|
assert result.size == (64, 64)
|
|
|
|
# Test float [-1, 1] numpy array
|
|
arr = np.random.rand(64, 64, 3).astype(np.float32) * 2 - 1
|
|
result = _normalize_image(arr)
|
|
assert isinstance(result, Image.Image)
|
|
assert result.size == (64, 64)
|
|
|
|
# Test batch dimensions (1, 1, H, W, C)
|
|
arr = np.random.randint(0, 255, (1, 1, 64, 64, 3), dtype=np.uint8)
|
|
result = _normalize_image(arr)
|
|
assert isinstance(result, Image.Image)
|
|
assert result.size == (64, 64)
|
|
|
|
|
|
def test_extract_images_from_result():
|
|
"""Test _extract_images_from_result with various result formats"""
|
|
import numpy as np
|
|
|
|
from vllm_omni.entrypoints.openai.api_server import _extract_images_from_result
|
|
|
|
# Test empty result
|
|
class EmptyResult:
|
|
pass
|
|
|
|
result = EmptyResult()
|
|
images = _extract_images_from_result(result)
|
|
assert images == []
|
|
|
|
# Test nested batch: [np.array(shape=(3, 64, 64, 3))]
|
|
batch = np.random.randint(0, 255, (3, 1, 64, 64, 3), dtype=np.uint8)
|
|
|
|
class BatchResult:
|
|
def __init__(self):
|
|
self.images = [batch]
|
|
|
|
result = BatchResult()
|
|
images = _extract_images_from_result(result)
|
|
assert len(images) == 3
|
|
assert all(isinstance(img, Image.Image) for img in images)
|
|
assert all(img.size == (64, 64) for img in images)
|
|
|
|
# Test dict path: result.request_output["images"]
|
|
class DictRequestOutput:
|
|
def __init__(self):
|
|
self.request_output = {"images": [np.random.randint(0, 255, (64, 64, 3), dtype=np.uint8)]}
|
|
|
|
result = DictRequestOutput()
|
|
images = _extract_images_from_result(result)
|
|
assert len(images) == 1
|
|
assert isinstance(images[0], Image.Image)
|
|
|
|
# Test attribute path: result.request_output.images
|
|
class AttrRequestOutput:
|
|
def __init__(self):
|
|
self.request_output = type(
|
|
"obj", (), {"images": [np.random.randint(0, 255, (32, 32, 3), dtype=np.uint8)]}
|
|
)()
|
|
|
|
result = AttrRequestOutput()
|
|
images = _extract_images_from_result(result)
|
|
assert len(images) == 1
|
|
assert isinstance(images[0], Image.Image)
|
|
assert images[0].size == (32, 32)
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# _check_max_generated_image_size unit tests
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
def test_width_height_within_limit_passes():
|
|
args = SimpleNamespace(max_generated_image_size=1024 * 1024)
|
|
# Exactly at limit is allowed (> not >=)
|
|
_check_max_generated_image_size(args, 1024, 1024)
|
|
# Below limit
|
|
_check_max_generated_image_size(args, 512, 512)
|
|
|
|
|
|
def test_width_height_exceeds_limit_raises_400():
|
|
limit = 1024 * 1024
|
|
args = SimpleNamespace(max_generated_image_size=limit)
|
|
with pytest.raises(HTTPException) as exc_info:
|
|
_check_max_generated_image_size(args, 1025, 1024)
|
|
assert exc_info.value.status_code == HTTPStatus.BAD_REQUEST.value
|
|
assert "1025x1024" in exc_info.value.detail
|
|
assert str(limit) in exc_info.value.detail
|
|
|
|
|
|
def test_width_height_error_message_contains_size_hint():
|
|
args = SimpleNamespace(max_generated_image_size=512 * 512)
|
|
with pytest.raises(HTTPException) as exc_info:
|
|
_check_max_generated_image_size(args, 1024, 512)
|
|
assert "--max-generated-image-size" in exc_info.value.detail
|
|
|
|
|
|
def test_resolution_within_limit_passes():
|
|
args = SimpleNamespace(max_generated_image_size=1024 * 1024)
|
|
# Exactly at limit is allowed (> not >=): 1024*1024 == limit
|
|
_check_max_generated_image_size(args, None, None, resolution=1024)
|
|
# Below limit
|
|
_check_max_generated_image_size(args, None, None, resolution=512)
|
|
|
|
|
|
def test_resolution_exceeds_limit_raises_400():
|
|
limit = 1024 * 1024
|
|
args = SimpleNamespace(max_generated_image_size=limit)
|
|
with pytest.raises(HTTPException) as exc_info:
|
|
_check_max_generated_image_size(args, None, None, resolution=1025)
|
|
assert exc_info.value.status_code == HTTPStatus.BAD_REQUEST.value
|
|
detail = exc_info.value.detail
|
|
assert "1025" in detail
|
|
assert "1025x1025" in detail
|
|
assert str(limit) in detail
|
|
|
|
|
|
def test_image_edits_size_auto_preserves_bridge_size(async_omni_stage_configs_only_client):
|
|
"""size=auto must NOT pin the diffusion stage sampling_params.height/width.
|
|
|
|
Regression: prior to the fix, edit_images resolved size=auto to the
|
|
first input image dimensions and forwarded them through gen_params +
|
|
extra_body to the diffusion stages sampling_params. AR-driven
|
|
pipelines (e.g. HunyuanImage-3.0) rely on ar2diffusions
|
|
bridge to override the final bucket via the AR-predicted ratio token,
|
|
and the DiT pre_process_func only fills sampling_params from the
|
|
bridge value when sampling_params.width is None (see
|
|
pipeline_hunyuan_image3.py:290). Non-None width from the input image
|
|
silently suppressed the AR decision, producing the wrong bucket
|
|
(e.g. 1024x1024 square instead of the AR-decided 1280x720 landscape
|
|
for multi-image fusion).
|
|
|
|
Cross-pins the multi-image fix at the API level: 2 reference images
|
|
with bot_task=think must produce 2 <img> placeholders in the captured
|
|
AR prompt (build_prompt called with num_images=2).
|
|
"""
|
|
img_a = make_test_image_bytes((32, 32))
|
|
img_b = make_test_image_bytes((128, 64))
|
|
response = async_omni_stage_configs_only_client.post(
|
|
"/v1/images/edits",
|
|
files=[("image", img_a), ("image", img_b)],
|
|
data={
|
|
"prompt": "fuse",
|
|
"size": "auto",
|
|
"bot_task": "think",
|
|
},
|
|
)
|
|
assert response.status_code == 200, response.text
|
|
|
|
engine = async_omni_stage_configs_only_client.app.state.engine_client
|
|
captured = engine.captured_sampling_params_list
|
|
assert captured is not None
|
|
assert len(captured) == 2
|
|
|
|
diffusion_params = captured[1]
|
|
assert diffusion_params.height is None, (
|
|
f"size=auto leaked into diffusion sampling_params.height={diffusion_params.height}; "
|
|
"must stay None so AR-driven pipelines can apply the bridges decision."
|
|
)
|
|
assert diffusion_params.width is None, (
|
|
f"size=auto leaked into diffusion sampling_params.width={diffusion_params.width}; "
|
|
"must stay None so AR-driven pipelines can apply the bridges decision."
|
|
)
|
|
|
|
KEY = "prompt"
|
|
IMG = "<img>"
|
|
captured_prompt = engine.captured_prompt
|
|
if isinstance(captured_prompt, dict) and isinstance(captured_prompt.get("prompt"), str):
|
|
assert captured_prompt["prompt"].count("<img>") == 2, (
|
|
f"N=2 reference images must emit 2 <img> placeholders in AR prompt; got {captured_prompt[KEY].count(IMG)} -- prompt: {captured_prompt[KEY]!r}"
|
|
)
|