Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project

from types import SimpleNamespace

Expand Down Expand Up @@ -168,7 +168,7 @@ def test_postprocess_preserves_image_contract_for_serving():
format_diffusion_outputs,
normalize_diffusion_postprocess_output,
)
from vllm_omni.entrypoints.openai.api_server import _extract_images_from_result
from vllm_omni.entrypoints.openai.images.helpers import _extract_images_from_result

sampling = OmniDiffusionSamplingParams(
num_inference_steps=1,
Expand Down
26 changes: 11 additions & 15 deletions tests/entrypoints/openai_api/test_api_server_guards.py
Original file line number Diff line number Diff line change
Expand Up @@ -55,6 +55,7 @@
from vllm.v1.engine.exceptions import EngineDeadError, EngineGenerateError

from vllm_omni.entrypoints.openai import api_server
from vllm_omni.entrypoints.serve.utils import errors as serve_errors

pytestmark = [pytest.mark.core_model, pytest.mark.cpu]

Expand Down Expand Up @@ -656,22 +657,23 @@ def test_images_generation_without_engine_preserves_service_unavailable_error()
assert exc_info.value.detail == "Multi-stage engine not initialized. Start server with a multi-stage omni model."


def test_images_generation_without_multistage_chat_handler_preserves_unavailable_error(monkeypatch) -> None:
def test_images_generation_without_multistage_chat_handler_preserves_unavailable_error() -> None:
"""Lock images generation when chat serving was not wired.

Fails if multi-stage images stop requiring ``openai_serving_chat``, or the
503 detail changes after images/chat ownership splits.
"""
app = FastAPI()
app.state.engine_client = SimpleNamespace(stage_configs=[object(), object()])
# Prefer real stage_type input over monkeypatching ``get_stage_type``:
# ``_get_engine_and_model`` lives in ``app_state`` and binds utils directly.
app.state.engine_client = SimpleNamespace(stage_configs=[{"stage_type": "llm"}, {"stage_type": "diffusion"}])
app.state.stage_configs = app.state.engine_client.stage_configs
app.state.openai_serving_models = SimpleNamespace(base_model_paths=[SimpleNamespace(name="demo-model")])
app.state.openai_serving_chat = None
app.state.args = SimpleNamespace(max_generated_image_size=None)

raw_request = _request_for(app, method="POST", path="/v1/images/generations")
request = api_server.ImageGenerationRequest(prompt="a cat", model="demo-model")
monkeypatch.setattr(api_server, "get_stage_type", lambda _stage_cfg: "diffusion")

with pytest.raises(HTTPException) as exc_info:
asyncio.run(api_server.generate_images(request, raw_request))
Expand Down Expand Up @@ -709,17 +711,15 @@ def test_engine_error_json_response_includes_request_and_stage_fields(monkeypatc
in P0.2 is not a false red.
"""
app = FastAPI()
# TODO(P0.2): retarget this call if ``_register_omni_exception_handlers``
# leaves api_server (red here is a rename, not a surface regression).
api_server._register_omni_exception_handlers(app)
serve_errors._register_omni_exception_handlers(app)
app.state.engine_client = SimpleNamespace(errored=True, engine=SimpleNamespace(is_alive=lambda: False))
app.state.server = object()
app.state.args = SimpleNamespace(log_error_stack=False)

req = _request_for(app, method="POST", path="/v1/chat/completions")
req.state.request_metadata = SimpleNamespace(request_id="req-123")

monkeypatch.setattr(api_server, "terminate_if_errored", lambda **_kwargs: None)
monkeypatch.setattr(serve_errors, "terminate_if_errored", lambda **_kwargs: None)

exc = EngineGenerateError("boom")
exc.error_stage_id = 2 # type: ignore[attr-defined]
Expand All @@ -740,14 +740,12 @@ def test_engine_dead_error_handler_registered_returns_json(monkeypatch) -> None:
JSON with ``request_id``.
"""
app = FastAPI()
# TODO(P0.2): retarget this call if ``_register_omni_exception_handlers``
# leaves api_server (red here is a rename, not a surface regression).
api_server._register_omni_exception_handlers(app)
serve_errors._register_omni_exception_handlers(app)
app.state.engine_client = SimpleNamespace(errored=True, engine=SimpleNamespace(is_alive=lambda: False))
app.state.server = object()
app.state.args = SimpleNamespace(log_error_stack=False)

monkeypatch.setattr(api_server, "terminate_if_errored", lambda **_kwargs: None)
monkeypatch.setattr(serve_errors, "terminate_if_errored", lambda **_kwargs: None)

handler = app.exception_handlers[EngineDeadError]
req = _request_for(app)
Expand All @@ -766,7 +764,7 @@ async def test_pure_diffusion_app_state_key_snapshot(monkeypatch) -> None:
models), wires a handler that must stay None, or leaves a required
handler as None.
"""
stage = SimpleNamespace(engine_args={})
stage = SimpleNamespace(stage_type="diffusion", engine_args={})
engine = _FakeEngineClient(stage_configs=[stage])

def _for_diffusion_factory(label: str):
Expand All @@ -776,7 +774,6 @@ def _factory(cls, *args, **kwargs):

return _factory

monkeypatch.setattr(api_server, "get_stage_type", lambda _cfg: "diffusion")
monkeypatch.setattr(api_server.OmniOpenAIServingChat, "for_diffusion", _for_diffusion_factory("chat"))
monkeypatch.setattr(api_server.OmniOpenAIServingChatBatch, "for_diffusion", _for_diffusion_factory("chat_batch"))
monkeypatch.setattr(
Expand Down Expand Up @@ -811,7 +808,7 @@ def _factory(cls, *args, **kwargs):

@pytest.mark.asyncio
async def test_pure_diffusion_speech_forwards_media_access_args(monkeypatch) -> None:
stage = SimpleNamespace(engine_args={})
stage = SimpleNamespace(stage_type="diffusion", engine_args={})
engine = _FakeEngineClient(stage_configs=[stage])
speech_kwargs = {}

Expand All @@ -827,7 +824,6 @@ def _speech_factory(cls, *args, **kwargs):
speech_kwargs.update(kwargs)
return _marker("speech")

monkeypatch.setattr(api_server, "get_stage_type", lambda _cfg: "diffusion")
monkeypatch.setattr(api_server.OmniOpenAIServingChat, "for_diffusion", _for_diffusion_factory("chat"))
monkeypatch.setattr(api_server.OmniOpenAIServingChatBatch, "for_diffusion", _for_diffusion_factory("chat_batch"))
monkeypatch.setattr(
Expand Down
12 changes: 7 additions & 5 deletions tests/entrypoints/openai_api/test_image_server.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project
"""
Tests for async image generation API endpoints.

Expand All @@ -24,11 +24,13 @@
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.api_server import router
from vllm_omni.entrypoints.openai.image_api_utils import (
encode_image_base64,
parse_size,
)
from vllm_omni.entrypoints.openai.images.helpers import _check_max_generated_image_size
from vllm_omni.entrypoints.openai.models.serving import _DiffusionServingModels
from vllm_omni.entrypoints.openai.serving_chat import OmniOpenAIServingChat
from vllm_omni.errors import GuardrailViolationError
from vllm_omni.inputs.data import OmniDiffusionSamplingParams
Expand Down Expand Up @@ -219,7 +221,7 @@ def test_client(mock_async_diffusion):
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
from vllm_omni.entrypoints.openai.models.serving import _DiffusionServingModels

app.state.openai_serving_models = _DiffusionServingModels(
[BaseModelPath(name="Qwen/Qwen-Image", model_path="Qwen/Qwen-Image")]
Expand Down Expand Up @@ -2126,7 +2128,7 @@ 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
from vllm_omni.entrypoints.openai.images.helpers import _normalize_image

# Test PIL Image input
img = Image.new("RGB", (64, 64), color="red")
Expand Down Expand Up @@ -2163,7 +2165,7 @@ 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
from vllm_omni.entrypoints.openai.images.helpers import _extract_images_from_result

# Test empty result
class EmptyResult:
Expand Down
4 changes: 2 additions & 2 deletions tests/entrypoints/openai_api/test_omni_sleep_wakeup.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project

import dataclasses

Expand Down Expand Up @@ -264,7 +264,7 @@ def pure_diffusion_app(pure_diffusion_engine, mocker):
mocker.patch(
"vllm_omni.entrypoints.openai.api_server.OmniOpenAIServingSpeech"
).for_diffusion.return_value = mocker.MagicMock()
mocker.patch("vllm_omni.entrypoints.openai.api_server._DiffusionServingModels")
mocker.patch("vllm_omni.entrypoints.openai.models.serving._DiffusionServingModels")

loop = asyncio.new_event_loop()
try:
Expand Down
Original file line number Diff line number Diff line change
@@ -1,3 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project

import logging
from inspect import Signature, signature
from types import SimpleNamespace
Expand All @@ -18,6 +21,7 @@
from vllm_omni.entrypoints.openai.serving_audio_generate import (
OmniOpenAIServingAudioGenerate,
)
from vllm_omni.entrypoints.serve.utils import errors as serve_errors
from vllm_omni.inputs.data import OmniDiffusionSamplingParams
from vllm_omni.outputs import OmniRequestOutput

Expand Down Expand Up @@ -515,7 +519,7 @@ def test_api_server_engine_error_response_includes_request_and_stage_id(
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 patch.object(serve_errors, "terminate_if_errored") as terminate_mock:
with TestClient(app) as client:
response = client.post("/v1/audio/generate", json={"input": "Hello"})

Expand Down
10 changes: 6 additions & 4 deletions tests/entrypoints/openai_api/test_serving_speech.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,7 @@
from vllm_omni.diffusion.sched.step_scheduler import StepScheduler
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 import errors as openai_errors
from vllm_omni.entrypoints.openai import serving_speech as serving_speech_module
from vllm_omni.entrypoints.openai.audio_utils_mixin import AudioMixin
from vllm_omni.entrypoints.openai.protocol.audio import (
Expand All @@ -58,6 +59,7 @@
from vllm_omni.entrypoints.openai.tts_adapters.ming_tts import MingTTSAdapter
from vllm_omni.entrypoints.openai.tts_adapters.qwen3_tts import Qwen3TTSAdapter, Qwen3TTSCodecLimitError
from vllm_omni.entrypoints.openai.tts_adapters.voxtral import VoxtralTTSAdapter
from vllm_omni.entrypoints.serve.utils import errors as serve_errors
from vllm_omni.inputs.data import OmniDiffusionSamplingParams
from vllm_omni.model_executor.models.fish_speech.prompt_utils import (
FISH_TEXT_ONLY_SYSTEM_PROMPT,
Expand Down Expand Up @@ -3783,7 +3785,7 @@ def _fake_create_error_response(message, err_type="BadRequestError", status_code

fake_base = mocker.MagicMock()
fake_base.create_error_response.side_effect = _fake_create_error_response
mocker.patch.object(api_server_module, "base", return_value=fake_base)
mocker.patch.object(openai_errors, "base", return_value=fake_base)
return fake_base


Expand Down Expand Up @@ -4098,7 +4100,7 @@ def test_api_server_create_speech_engine_error_response_includes_request_and_sta
)
)

terminate_mock = mocker.patch.object(api_server_module, "terminate_if_errored")
terminate_mock = mocker.patch.object(serve_errors, "terminate_if_errored")

raw_request = _make_api_server_request(handler, path="/v1/audio/speech")
raw_request.app.state.args = SimpleNamespace(log_error_stack=False)
Expand Down Expand Up @@ -4130,8 +4132,8 @@ def test_omni_engine_error_handler_includes_request_and_stage_id(mocker: MockerF
)
app.state.server = SimpleNamespace()

terminate_mock = mocker.patch.object(api_server_module, "terminate_if_errored")
api_server_module._register_omni_exception_handlers(app)
terminate_mock = mocker.patch.object(serve_errors, "terminate_if_errored")
serve_errors._register_omni_exception_handlers(app)

@app.get("/boom")
async def boom(request: Request):
Expand Down
24 changes: 16 additions & 8 deletions tests/entrypoints/openai_api/test_video_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,6 +37,12 @@
from vllm_omni.entrypoints.openai.serving_video import OmniOpenAIServingVideo, ReferenceImage
from vllm_omni.entrypoints.openai.storage import LocalStorageManager
from vllm_omni.entrypoints.openai.stores import AsyncDictStore, TaskRegistry
from vllm_omni.entrypoints.openai.video.generation import helpers as video_generation_helpers
from vllm_omni.entrypoints.openai.video.generation.helpers import (
MINIMAX_H3_MAX_REFERENCE_IMAGE_BYTES,
_read_upload_limited,
_reference_video_decode_spec,
)
from vllm_omni.errors import GuardrailViolationError
from vllm_omni.inputs.data import OmniDiffusionSamplingParams

Expand Down Expand Up @@ -192,6 +198,8 @@ def isolated_video_backends(tmp_path, monkeypatch):
monkeypatch.setattr(api_server, "VIDEO_STORE", store)
monkeypatch.setattr(api_server, "VIDEO_TASKS", tasks)
monkeypatch.setattr(api_server, "STORAGE_MANAGER", storage)
monkeypatch.setattr(video_generation_helpers, "VIDEO_STORE", store)
monkeypatch.setattr(video_generation_helpers, "STORAGE_MANAGER", storage)
return store, tasks, storage


Expand Down Expand Up @@ -245,7 +253,7 @@ async def fake_init_app_state(engine_client, state, args):
monkeypatch.setattr(api_server, "build_openai_app", lambda args, supported_tasks: FastAPI())
monkeypatch.setattr(api_server, "serve_http", fake_serve_http)
monkeypatch.setattr(api_server.STORAGE_MANAGER, "start", fake_storage_start)
monkeypatch.setattr(api_server, "_get_vllm_config", fake_get_vllm_config)
monkeypatch.setattr(api_server.openai_app_state, "_get_vllm_config", fake_get_vllm_config)
monkeypatch.setattr(api_server, "omni_init_app_state", fake_init_app_state)
monkeypatch.setattr(api_server, "get_uvicorn_log_config", lambda args: None)

Expand Down Expand Up @@ -978,7 +986,7 @@ def test_cosmos3_reference_video_limit_uses_v2v_condition_frames():
extra_params={"condition_frame_indexes_vision": [0, 2]},
)

spec = api_server._reference_video_decode_spec(request, _cosmos3_stage_configs())
spec = _reference_video_decode_spec(request, _cosmos3_stage_configs())
assert spec.max_frames == 9
assert spec.keep == "first"

Expand All @@ -990,7 +998,7 @@ def test_cosmos3_reference_video_limit_preserves_action_frames():
extra_params={"action_mode": "inverse_dynamics", "action_chunk_size": 16},
)

assert api_server._reference_video_decode_spec(request, _cosmos3_stage_configs()).max_frames == 17
assert _reference_video_decode_spec(request, _cosmos3_stage_configs()).max_frames == 17


def test_cosmos3_reference_video_limit_caps_condition_frames_to_output_frames():
Expand All @@ -1000,7 +1008,7 @@ def test_cosmos3_reference_video_limit_caps_condition_frames_to_output_frames():
extra_params={"condition_frame_indexes_vision": [0, 20]},
)

assert api_server._reference_video_decode_spec(request, _cosmos3_stage_configs()).max_frames == 5
assert _reference_video_decode_spec(request, _cosmos3_stage_configs()).max_frames == 5


def test_s2v_video_generation_with_audio_reference_form(test_client, mocker: MockerFixture):
Expand Down Expand Up @@ -1681,15 +1689,15 @@ def test_h3_multipart_maps_pillow_pixel_limit_error(field, test_client, monkeypa
@pytest.mark.asyncio
async def test_h3_upload_limit_checks_declared_size_before_read():
class OversizedUpload:
size = api_server.MINIMAX_H3_MAX_REFERENCE_IMAGE_BYTES + 1
size = MINIMAX_H3_MAX_REFERENCE_IMAGE_BYTES + 1

async def read(self, _size):
raise AssertionError("the oversized upload must be rejected before reading")

with pytest.raises(HTTPException, match="size limit"):
await api_server._read_upload_limited(
await _read_upload_limited(
OversizedUpload(),
max_bytes=api_server.MINIMAX_H3_MAX_REFERENCE_IMAGE_BYTES,
max_bytes=MINIMAX_H3_MAX_REFERENCE_IMAGE_BYTES,
)


Expand Down Expand Up @@ -2438,7 +2446,7 @@ def test_cosmos3_control_upload_rejects_existing_control_source(test_client):
)
def test_cosmos3_control_upload_rejects_invalid_size(control_bytes, message, test_client, monkeypatch):
test_client.app.state.openai_serving_video._engine_client.model_class_name = "Cosmos3OmniDiffusersPipeline"
monkeypatch.setattr(api_server, "CONTROL_REFERENCE_MAX_BYTES", 3)
monkeypatch.setattr(video_generation_helpers, "CONTROL_REFERENCE_MAX_BYTES", 3)

response = test_client.post(
"/v1/videos/sync",
Expand Down
19 changes: 2 additions & 17 deletions vllm_omni/config/endpoint_policy.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,16 +2,16 @@
# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project
"""Endpoint restriction policy for omni pipelines."""

from collections.abc import Set
from dataclasses import dataclass
from enum import Enum
from typing import NamedTuple

from fastapi import FastAPI, Request
from fastapi.responses import JSONResponse
from starlette.routing import Route
from vllm.entrypoints.serve.exception_handling.error_response import create_error_response

from vllm_omni.entrypoints.serve.utils.routes import remove_route_from_app


class RouteTarget(NamedTuple):
"""A server path & supported methods."""
Expand Down Expand Up @@ -53,21 +53,6 @@ async def rejection_handler(raw_request: Request):
return rejection_handler


def remove_route_from_app(
app: FastAPI,
path: str,
methods: Set[str],
) -> None:
"""Remove routes matching a path and one of the given HTTP methods."""
routes_to_remove = [
route
for route in app.routes
if isinstance(route, Route) and route.path == path and route.methods is not None and route.methods & methods
]
for route in routes_to_remove:
app.routes.remove(route)


def shutdown_unsupported_routes(
app: FastAPI,
endpoint_restrictions: tuple[EndpointRestriction, ...],
Expand Down
Loading
Loading