vllm-project--vllm-omni
565 行
21 KiB
Python
565 行
21 KiB
Python
import logging
|
|
from inspect import Signature, signature
|
|
from types import SimpleNamespace
|
|
from unittest.mock import AsyncMock, MagicMock, patch
|
|
|
|
import pytest
|
|
import torch
|
|
from fastapi import FastAPI, Request
|
|
from fastapi.testclient import TestClient
|
|
from vllm.v1.engine.exceptions import EngineGenerateError
|
|
|
|
from vllm_omni.entrypoints.omni_base import OmniEngineDeadError
|
|
from vllm_omni.entrypoints.openai import api_server as api_server_module
|
|
from vllm_omni.entrypoints.openai.protocol.audio import (
|
|
CreateAudio,
|
|
OpenAICreateAudioGenerateRequest,
|
|
)
|
|
from vllm_omni.entrypoints.openai.serving_audio_generate import (
|
|
OmniOpenAIServingAudioGenerate,
|
|
)
|
|
from vllm_omni.inputs.data import OmniDiffusionSamplingParams
|
|
from vllm_omni.outputs import OmniRequestOutput
|
|
|
|
pytestmark = [pytest.mark.core_model, pytest.mark.cpu]
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
# Helper: create a mock audio output for endpoint tests
|
|
def create_mock_audio_output(
|
|
request_id: str = "audiogen-mock-123",
|
|
sample_rate: int = 44100,
|
|
num_samples: int = 44100,
|
|
audio_key: str = "audio",
|
|
) -> OmniRequestOutput:
|
|
"""Return an OmniRequestOutput mimicking diffusion audio model output."""
|
|
|
|
audio_tensor = torch.sin(torch.linspace(0, 440 * 2 * torch.pi, num_samples))
|
|
|
|
return OmniRequestOutput.from_diffusion(
|
|
request_id=request_id,
|
|
images=[],
|
|
prompt=None,
|
|
metrics={},
|
|
multimodal_output={audio_key: audio_tensor, "sr": sample_rate},
|
|
)
|
|
|
|
|
|
def _make_engine_client(*, audio_key: str = "audio", sample_rate: int = 44100):
|
|
"""Build a mock engine client producing audio output."""
|
|
mock_engine_client = MagicMock()
|
|
mock_engine_client.errored = False
|
|
mock_engine_client.model_type = "StableAudioPipeline"
|
|
mock_engine_client.default_sampling_params_list = [{}]
|
|
|
|
async def mock_generate_fn(*args, **kwargs):
|
|
yield create_mock_audio_output(
|
|
request_id=kwargs.get("request_id", "audiogen-mock"),
|
|
sample_rate=sample_rate,
|
|
audio_key=audio_key,
|
|
)
|
|
|
|
mock_engine_client.generate = MagicMock(side_effect=mock_generate_fn)
|
|
return mock_engine_client
|
|
|
|
|
|
def _make_server(engine_client=None):
|
|
"""Build an OmniOpenAIServingAudioGenerate with mocks."""
|
|
if engine_client is None:
|
|
engine_client = _make_engine_client()
|
|
|
|
mock_models = MagicMock()
|
|
mock_models.is_base_model.return_value = True
|
|
|
|
return OmniOpenAIServingAudioGenerate(
|
|
engine_client=engine_client,
|
|
models=mock_models,
|
|
request_logger=MagicMock(),
|
|
)
|
|
|
|
|
|
@pytest.fixture
|
|
def test_app():
|
|
server = _make_server()
|
|
|
|
original_fn = server.create_audio_generate
|
|
sig = signature(original_fn)
|
|
new_params = [p for name, p in sig.parameters.items() if name != "raw_request"]
|
|
new_sig = Signature(parameters=new_params, return_annotation=sig.return_annotation)
|
|
|
|
async def patched_create_audio_generate(*args, **kwargs):
|
|
return await original_fn(*args, **kwargs)
|
|
|
|
patched_create_audio_generate.__signature__ = new_sig
|
|
server.create_audio_generate = patched_create_audio_generate
|
|
|
|
app = FastAPI()
|
|
app.add_api_route(
|
|
"/v1/audio/generate",
|
|
server.create_audio_generate,
|
|
methods=["POST"],
|
|
response_model=None,
|
|
)
|
|
return app
|
|
|
|
|
|
@pytest.fixture
|
|
def client(test_app):
|
|
return TestClient(test_app)
|
|
|
|
|
|
def _make_api_server_test_app(handler, *, request_id: str = "audio-gen-req-1"):
|
|
app = FastAPI()
|
|
app.state.openai_serving_audio_generate = handler
|
|
app.state.engine_client = SimpleNamespace(
|
|
engine=SimpleNamespace(is_alive=lambda: False),
|
|
errored=True,
|
|
)
|
|
app.state.server = SimpleNamespace()
|
|
|
|
@app.middleware("http")
|
|
async def add_request_metadata(request: Request, call_next):
|
|
request.state.request_metadata = SimpleNamespace(request_id=request_id)
|
|
return await call_next(request)
|
|
|
|
app.add_api_route(
|
|
"/v1/audio/generate",
|
|
api_server_module.create_audio_generate,
|
|
methods=["POST"],
|
|
response_model=None,
|
|
)
|
|
return app
|
|
|
|
|
|
# Request Validation (Pydantic model)
|
|
class TestRequestValidation:
|
|
"""Validate OpenAICreateAudioGenerateRequest pydantic constraints."""
|
|
|
|
def test_valid_minimal_request(self):
|
|
req = OpenAICreateAudioGenerateRequest(input="A calm piano melody")
|
|
assert req.input == "A calm piano melody"
|
|
assert req.response_format == "wav"
|
|
assert req.speed == 1.0
|
|
|
|
def test_fields_are_wired_correctly(self):
|
|
req = OpenAICreateAudioGenerateRequest(
|
|
input="rain sounds",
|
|
model="stable-audio",
|
|
response_format="flac",
|
|
speed=1.5,
|
|
audio_length=10.0,
|
|
audio_start=2.0,
|
|
negative_prompt="noise",
|
|
guidance_scale=7.5,
|
|
num_inference_steps=100,
|
|
seed=42,
|
|
)
|
|
assert req.input == "rain sounds"
|
|
assert req.model == "stable-audio"
|
|
assert req.response_format == "flac"
|
|
assert req.speed == 1.5
|
|
assert req.audio_length == 10.0
|
|
assert req.audio_start == 2.0
|
|
assert req.negative_prompt == "noise"
|
|
assert req.guidance_scale == 7.5
|
|
assert req.num_inference_steps == 100
|
|
assert req.seed == 42
|
|
|
|
def test_invalid_response_format(self):
|
|
with pytest.raises(Exception):
|
|
OpenAICreateAudioGenerateRequest(input="test", response_format="invalid_format")
|
|
|
|
def test_speed_lower_bound(self):
|
|
with pytest.raises(Exception):
|
|
OpenAICreateAudioGenerateRequest(input="test", speed=0.1)
|
|
|
|
def test_speed_upper_bound(self):
|
|
with pytest.raises(Exception):
|
|
OpenAICreateAudioGenerateRequest(input="test", speed=5.0)
|
|
|
|
def test_speed_at_boundaries(self):
|
|
req_low = OpenAICreateAudioGenerateRequest(input="test", speed=0.25)
|
|
assert req_low.speed == 0.25
|
|
req_high = OpenAICreateAudioGenerateRequest(input="test", speed=4.0)
|
|
assert req_high.speed == 4.0
|
|
|
|
def test_stream_format_sse_rejected(self):
|
|
with pytest.raises(Exception):
|
|
OpenAICreateAudioGenerateRequest(input="test", stream_format="sse")
|
|
|
|
def test_stream_format_audio_accepted(self):
|
|
req = OpenAICreateAudioGenerateRequest(input="test", stream_format="audio")
|
|
assert req.stream_format == "audio"
|
|
|
|
|
|
# Constructor & Class Methods
|
|
class TestConstructor:
|
|
def test_default_init(self):
|
|
server = _make_server()
|
|
assert server.diffusion_mode is False
|
|
|
|
def test_for_diffusion_factory(self):
|
|
engine_client = _make_engine_client()
|
|
mock_models = MagicMock()
|
|
mock_models.is_base_model.return_value = True
|
|
|
|
server = OmniOpenAIServingAudioGenerate.for_diffusion(
|
|
engine_client=engine_client,
|
|
models=mock_models,
|
|
request_logger=MagicMock(),
|
|
)
|
|
assert server.diffusion_mode is True
|
|
|
|
def test_is_stable_audio_model_true(self):
|
|
server = _make_server()
|
|
assert server._is_stable_audio_model() is True
|
|
|
|
def test_is_stable_audio_model_false(self):
|
|
engine = _make_engine_client()
|
|
engine.model_type = "SomeOtherModel"
|
|
server = _make_server(engine_client=engine)
|
|
assert server._is_stable_audio_model() is False
|
|
|
|
|
|
# Parameter Wiring — verify request params reach the engine
|
|
class TestParameterWiring:
|
|
"""Ensure request parameters are correctly forwarded to the engine."""
|
|
|
|
@pytest.fixture
|
|
def server_and_engine(self):
|
|
engine = _make_engine_client()
|
|
server = _make_server(engine_client=engine)
|
|
return server, engine
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_prompt_wiring(self, server_and_engine):
|
|
server, engine = server_and_engine
|
|
req = OpenAICreateAudioGenerateRequest(input="birds chirping")
|
|
await server.create_audio_generate(req)
|
|
|
|
engine.generate.assert_called_once()
|
|
call_kwargs = engine.generate.call_args[1]
|
|
assert call_kwargs["prompt"]["prompt"] == "birds chirping"
|
|
assert call_kwargs["output_modalities"] == ["audio"]
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_negative_prompt_wiring(self, server_and_engine):
|
|
server, engine = server_and_engine
|
|
req = OpenAICreateAudioGenerateRequest(input="a calm ocean", negative_prompt="noise distortion")
|
|
await server.create_audio_generate(req)
|
|
|
|
call_kwargs = engine.generate.call_args[1]
|
|
assert call_kwargs["prompt"]["negative_prompt"] == "noise distortion"
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_negative_prompt_absent(self, server_and_engine):
|
|
server, engine = server_and_engine
|
|
req = OpenAICreateAudioGenerateRequest(input="a calm ocean")
|
|
await server.create_audio_generate(req)
|
|
|
|
call_kwargs = engine.generate.call_args[1]
|
|
assert "negative_prompt" not in call_kwargs["prompt"]
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_guidance_scale_wiring(self, server_and_engine):
|
|
server, engine = server_and_engine
|
|
req = OpenAICreateAudioGenerateRequest(input="test", guidance_scale=12.0)
|
|
await server.create_audio_generate(req)
|
|
|
|
call_kwargs = engine.generate.call_args[1]
|
|
sp = call_kwargs["sampling_params_list"][0]
|
|
assert isinstance(sp, OmniDiffusionSamplingParams)
|
|
assert sp.guidance_scale == 12.0
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_num_inference_steps_wiring(self, server_and_engine):
|
|
server, engine = server_and_engine
|
|
req = OpenAICreateAudioGenerateRequest(input="test", num_inference_steps=200)
|
|
await server.create_audio_generate(req)
|
|
|
|
sp = engine.generate.call_args[1]["sampling_params_list"][0]
|
|
assert sp.num_inference_steps == 200
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_seed_creates_generator(self, server_and_engine):
|
|
server, engine = server_and_engine
|
|
req = OpenAICreateAudioGenerateRequest(input="test", seed=42)
|
|
|
|
with patch("vllm_omni.entrypoints.openai.serving_audio_generate.torch") as mock_torch:
|
|
mock_gen = MagicMock()
|
|
mock_gen.manual_seed.return_value = mock_gen
|
|
mock_torch.Generator.return_value = mock_gen
|
|
|
|
await server.create_audio_generate(req)
|
|
|
|
mock_torch.Generator.assert_called_once()
|
|
mock_gen.manual_seed.assert_called_once_with(42)
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_seed_none_skips_generator(self, server_and_engine):
|
|
server, engine = server_and_engine
|
|
req = OpenAICreateAudioGenerateRequest(input="test")
|
|
|
|
await server.create_audio_generate(req)
|
|
|
|
sp = engine.generate.call_args[1]["sampling_params_list"][0]
|
|
assert sp.generator is None
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_audio_length_wiring(self, server_and_engine):
|
|
server, engine = server_and_engine
|
|
req = OpenAICreateAudioGenerateRequest(input="test", audio_length=10.0, audio_start=2.0)
|
|
await server.create_audio_generate(req)
|
|
|
|
sp = engine.generate.call_args[1]["sampling_params_list"][0]
|
|
assert sp.extra_args["audio_start_in_s"] == 2.0
|
|
assert sp.extra_args["audio_end_in_s"] == 12.0 # start + length
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_audio_length_default_start(self, server_and_engine):
|
|
server, engine = server_and_engine
|
|
req = OpenAICreateAudioGenerateRequest(input="test", audio_length=5.0)
|
|
await server.create_audio_generate(req)
|
|
|
|
sp = engine.generate.call_args[1]["sampling_params_list"][0]
|
|
assert sp.extra_args["audio_start_in_s"] == 0.0
|
|
assert sp.extra_args["audio_end_in_s"] == 5.0
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_no_audio_length_skips_extra_args(self, server_and_engine):
|
|
server, engine = server_and_engine
|
|
req = OpenAICreateAudioGenerateRequest(input="test")
|
|
await server.create_audio_generate(req)
|
|
|
|
sp = engine.generate.call_args[1]["sampling_params_list"][0]
|
|
assert sp.extra_args == {}
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_defaults_not_set_when_omitted(self, server_and_engine):
|
|
"""Guidance scale and num_inference_steps keep dataclass defaults when not in request."""
|
|
server, engine = server_and_engine
|
|
req = OpenAICreateAudioGenerateRequest(input="test")
|
|
await server.create_audio_generate(req)
|
|
|
|
sp = engine.generate.call_args[1]["sampling_params_list"][0]
|
|
defaults = OmniDiffusionSamplingParams()
|
|
assert sp.guidance_scale == defaults.guidance_scale
|
|
assert sp.num_inference_steps == defaults.num_inference_steps
|
|
|
|
|
|
# Audio Response Format
|
|
class TestAudioResponseFormat:
|
|
def test_wav_response(self, client):
|
|
payload = {"input": "a gentle rain", "response_format": "wav"}
|
|
response = client.post("/v1/audio/generate", json=payload)
|
|
assert response.status_code == 200
|
|
assert response.headers["content-type"] == "audio/wav"
|
|
assert len(response.content) > 0
|
|
|
|
def test_mp3_response(self, client):
|
|
payload = {"input": "a gentle rain", "response_format": "mp3"}
|
|
response = client.post("/v1/audio/generate", json=payload)
|
|
assert response.status_code == 200
|
|
assert response.headers["content-type"] == "audio/mpeg"
|
|
assert len(response.content) > 0
|
|
|
|
def test_flac_response(self, client):
|
|
payload = {"input": "a gentle rain", "response_format": "flac"}
|
|
response = client.post("/v1/audio/generate", json=payload)
|
|
assert response.status_code == 200
|
|
assert response.headers["content-type"] == "audio/flac"
|
|
assert len(response.content) > 0
|
|
|
|
def test_invalid_format_rejected(self, client):
|
|
payload = {"input": "test", "response_format": "banana"}
|
|
response = client.post("/v1/audio/generate", json=payload)
|
|
assert response.status_code == 422
|
|
|
|
@patch("vllm_omni.entrypoints.openai.serving_audio_generate.OmniOpenAIServingAudioGenerate.create_audio")
|
|
def test_speed_parameter_forwarded(self, mock_create_audio, test_app):
|
|
mock_audio_response = MagicMock()
|
|
mock_audio_response.audio_data = b"dummy_audio"
|
|
mock_audio_response.media_type = "audio/wav"
|
|
mock_create_audio.return_value = mock_audio_response
|
|
|
|
c = TestClient(test_app)
|
|
payload = {"input": "test", "response_format": "wav", "speed": 2.5}
|
|
c.post("/v1/audio/generate", json=payload)
|
|
|
|
mock_create_audio.assert_called_once()
|
|
audio_obj = mock_create_audio.call_args[0][0]
|
|
assert isinstance(audio_obj, CreateAudio)
|
|
assert audio_obj.speed == 2.5
|
|
|
|
@patch("vllm_omni.entrypoints.openai.serving_audio_generate.OmniOpenAIServingAudioGenerate.create_audio")
|
|
def test_sample_rate_from_output(self, mock_create_audio, test_app):
|
|
mock_audio_response = MagicMock()
|
|
mock_audio_response.audio_data = b"dummy"
|
|
mock_audio_response.media_type = "audio/wav"
|
|
mock_create_audio.return_value = mock_audio_response
|
|
|
|
c = TestClient(test_app)
|
|
payload = {"input": "test"}
|
|
c.post("/v1/audio/generate", json=payload)
|
|
|
|
audio_obj = mock_create_audio.call_args[0][0]
|
|
assert audio_obj.sample_rate == 44100 # Stable Audio default
|
|
|
|
|
|
# Error Handling
|
|
class TestErrorHandling:
|
|
@pytest.mark.asyncio
|
|
async def test_no_output_returns_error(self):
|
|
engine = _make_engine_client()
|
|
|
|
async def empty_gen(*args, **kwargs):
|
|
return
|
|
yield # unreachable – makes this an async generator
|
|
|
|
engine.generate = MagicMock(side_effect=empty_gen)
|
|
server = _make_server(engine_client=engine)
|
|
req = OpenAICreateAudioGenerateRequest(input="test")
|
|
resp = await server.create_audio_generate(req)
|
|
|
|
# create_error_response returns an ErrorResponse with .error.message
|
|
assert "No output generated" in resp.error.message
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_no_audio_in_output_returns_error(self):
|
|
engine = _make_engine_client()
|
|
|
|
async def gen_without_audio(*args, **kwargs):
|
|
yield OmniRequestOutput.from_diffusion(
|
|
request_id="test",
|
|
images=[],
|
|
prompt=None,
|
|
metrics={},
|
|
multimodal_output={}, # no audio key
|
|
)
|
|
|
|
engine.generate = MagicMock(side_effect=gen_without_audio)
|
|
server = _make_server(engine_client=engine)
|
|
req = OpenAICreateAudioGenerateRequest(input="test")
|
|
resp = await server.create_audio_generate(req)
|
|
|
|
assert "did not produce audio" in resp.error.message
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_engine_errored_raises(self):
|
|
engine = _make_engine_client()
|
|
engine.errored = True
|
|
engine.dead_error = RuntimeError("engine is dead")
|
|
server = _make_server(engine_client=engine)
|
|
|
|
req = OpenAICreateAudioGenerateRequest(input="test")
|
|
with pytest.raises(RuntimeError, match="engine is dead"):
|
|
await server.create_audio_generate(req)
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_model_outputs_key_fallback(self):
|
|
"""Audio data under 'model_outputs' key should be accepted."""
|
|
engine = _make_engine_client(audio_key="model_outputs")
|
|
server = _make_server(engine_client=engine)
|
|
req = OpenAICreateAudioGenerateRequest(input="test")
|
|
resp = await server.create_audio_generate(req)
|
|
|
|
# Should succeed and return a Response with audio bytes
|
|
assert hasattr(resp, "body")
|
|
assert len(resp.body) > 0
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_value_error_returns_error_response(self):
|
|
engine = _make_engine_client()
|
|
|
|
async def gen_value_error(*args, **kwargs):
|
|
raise ValueError("bad value")
|
|
yield # unreachable
|
|
|
|
engine.generate = MagicMock(side_effect=gen_value_error)
|
|
server = _make_server(engine_client=engine)
|
|
req = OpenAICreateAudioGenerateRequest(input="test")
|
|
resp = await server.create_audio_generate(req)
|
|
|
|
assert "bad value" in resp.error.message
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_generic_exception_returns_error_response(self):
|
|
engine = _make_engine_client()
|
|
|
|
async def gen_runtime_error(*args, **kwargs):
|
|
raise RuntimeError("something went wrong")
|
|
yield # unreachable
|
|
|
|
engine.generate = MagicMock(side_effect=gen_runtime_error)
|
|
server = _make_server(engine_client=engine)
|
|
req = OpenAICreateAudioGenerateRequest(input="test")
|
|
resp = await server.create_audio_generate(req)
|
|
|
|
assert "Audio generation failed" in resp.error.message
|
|
|
|
@pytest.mark.parametrize(
|
|
("exc", "expected_message", "expected_stage_id"),
|
|
[
|
|
(OmniEngineDeadError("engine dead", error_stage_id=2), "engine dead", 2),
|
|
(EngineGenerateError("engine generate failed"), "engine generate failed", None),
|
|
],
|
|
)
|
|
def test_api_server_engine_error_response_includes_request_and_stage_id(
|
|
self,
|
|
exc,
|
|
expected_message,
|
|
expected_stage_id,
|
|
):
|
|
handler = MagicMock()
|
|
handler.create_audio_generate = AsyncMock(side_effect=exc)
|
|
app = _make_api_server_test_app(handler)
|
|
|
|
with patch.object(api_server_module, "terminate_if_errored") as terminate_mock:
|
|
with TestClient(app) as client:
|
|
response = client.post("/v1/audio/generate", json={"input": "Hello"})
|
|
|
|
assert response.status_code == 500
|
|
payload = response.json()
|
|
assert payload["error"]["message"] == expected_message
|
|
assert payload["error"]["code"] == 500
|
|
assert payload["error"]["request_id"] == "audio-gen-req-1"
|
|
assert payload["error"]["error_stage_id"] == expected_stage_id
|
|
terminate_mock.assert_called_once()
|
|
|
|
|
|
# End-to-End via TestClient
|
|
class TestAudioGenerateAPI:
|
|
def test_basic_success(self, client):
|
|
payload = {"input": "ambient forest sounds"}
|
|
response = client.post("/v1/audio/generate", json=payload)
|
|
assert response.status_code == 200
|
|
assert len(response.content) > 0
|
|
|
|
def test_with_all_params(self, client):
|
|
payload = {
|
|
"input": "gentle piano",
|
|
"response_format": "wav",
|
|
"speed": 1.0,
|
|
"audio_length": 5.0,
|
|
"audio_start": 0.0,
|
|
"negative_prompt": "noise",
|
|
"guidance_scale": 7.0,
|
|
"num_inference_steps": 50,
|
|
"seed": 123,
|
|
}
|
|
response = client.post("/v1/audio/generate", json=payload)
|
|
assert response.status_code == 200
|
|
assert response.headers["content-type"] == "audio/wav"
|
|
|
|
def test_missing_input_rejected(self, client):
|
|
payload = {}
|
|
response = client.post("/v1/audio/generate", json=payload)
|
|
assert response.status_code == 422
|
|
|
|
def test_extra_unknown_fields_ignored(self, client):
|
|
payload = {"input": "test", "unknown_field": "value"}
|
|
response = client.post("/v1/audio/generate", json=payload)
|
|
# Pydantic v2 ignores extra fields by default
|
|
assert response.status_code == 200
|