115 lines
4.5 KiB
Python
115 lines
4.5 KiB
Python
# SPDX-License-Identifier: Apache-2.0
|
|
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
|
"""Tests for chat endpoint speaker validation."""
|
|
|
|
import asyncio
|
|
from types import SimpleNamespace
|
|
|
|
import pytest
|
|
from pytest_mock import MockerFixture
|
|
|
|
from vllm_omni.entrypoints.openai.utils import (
|
|
get_supported_speakers_from_hf_config,
|
|
validate_requested_speaker,
|
|
)
|
|
|
|
pytestmark = [pytest.mark.core_model, pytest.mark.cpu]
|
|
|
|
|
|
@pytest.fixture
|
|
def serving_chat():
|
|
from vllm_omni.entrypoints.openai.serving_chat import OmniOpenAIServingChat
|
|
|
|
instance = object.__new__(OmniOpenAIServingChat)
|
|
instance._supported_speakers = None
|
|
instance.has_kv_connector = False
|
|
return instance
|
|
|
|
|
|
def _make_hf_config(mocker: MockerFixture, *, speaker_id: dict | None = None, spk_id: dict | None = None):
|
|
hf_config = mocker.MagicMock()
|
|
talker_config = mocker.MagicMock()
|
|
talker_config.speaker_id = speaker_id
|
|
talker_config.spk_id = spk_id
|
|
hf_config.talker_config = talker_config
|
|
return hf_config
|
|
|
|
|
|
def test_validate_requested_speaker_accepts_case_insensitive_value():
|
|
supported = {"vivian", "ethan"}
|
|
assert validate_requested_speaker("Vivian", supported) == "vivian"
|
|
assert validate_requested_speaker(" vivian ", supported) == "vivian"
|
|
|
|
|
|
def test_validate_requested_speaker_rejects_invalid_value_with_supported_list():
|
|
supported = {"vivian", "ethan"}
|
|
with pytest.raises(ValueError, match="Invalid speaker 'uncle_fu'. Supported: ethan, vivian"):
|
|
validate_requested_speaker("uncle_fu", supported)
|
|
|
|
|
|
def test_validate_requested_speaker_skips_validation_when_supported_empty():
|
|
assert validate_requested_speaker("anything", set()) == "anything"
|
|
assert validate_requested_speaker(" ", {"vivian"}) is None
|
|
|
|
|
|
def test_get_supported_speakers_from_hf_config_uses_spk_id_fallback(mocker: MockerFixture):
|
|
hf_config = _make_hf_config(mocker, speaker_id=None, spk_id={"Serena": 0})
|
|
assert get_supported_speakers_from_hf_config(hf_config) == {"serena"}
|
|
|
|
|
|
def test_get_supported_speakers_caches_normalized_keys(mocker: MockerFixture, serving_chat):
|
|
serving_chat.model_config = mocker.MagicMock()
|
|
serving_chat.model_config.hf_config = _make_hf_config(mocker, speaker_id={"Vivian": 0, "Ethan": 1})
|
|
|
|
assert serving_chat._get_supported_speakers() == {"vivian", "ethan"}
|
|
|
|
# Cached value should be reused even if the config changes afterwards.
|
|
serving_chat.model_config.hf_config.talker_config.speaker_id = {"Serena": 2}
|
|
assert serving_chat._get_supported_speakers() == {"vivian", "ethan"}
|
|
|
|
|
|
def test_create_chat_completion_converts_value_error_to_error_response(mocker: MockerFixture, serving_chat):
|
|
serving_chat._diffusion_mode = False
|
|
serving_chat._check_model = mocker.AsyncMock(return_value=None)
|
|
serving_chat.engine_client = mocker.MagicMock(errored=False)
|
|
serving_chat._maybe_get_adapters = mocker.MagicMock(return_value=None)
|
|
serving_chat.models = mocker.MagicMock()
|
|
serving_chat.models.model_name.return_value = "test-model"
|
|
serving_chat.renderer = mocker.MagicMock()
|
|
serving_chat.renderer.get_tokenizer.return_value = mocker.MagicMock()
|
|
serving_chat.reasoning_parser_cls = None
|
|
serving_chat.tool_parser = None
|
|
serving_chat.parser_cls = None
|
|
serving_chat.use_harmony = False
|
|
serving_chat.enable_auto_tools = False
|
|
serving_chat.exclude_tools_when_tool_choice_none = False
|
|
serving_chat.trust_request_chat_template = False
|
|
serving_chat.chat_template = None
|
|
serving_chat.chat_template_content_format = "string"
|
|
serving_chat.default_chat_template_kwargs = {}
|
|
serving_chat.online_renderer = mocker.MagicMock()
|
|
serving_chat.online_renderer.validate_chat_template.return_value = None
|
|
serving_chat._effective_chat_template_kwargs = mocker.MagicMock(return_value={})
|
|
serving_chat._preprocess_chat = mocker.AsyncMock(
|
|
side_effect=ValueError("Invalid speaker 'uncle_fu'. Supported: ethan, vivian")
|
|
)
|
|
serving_chat.create_error_response = mocker.MagicMock(return_value="error-response")
|
|
|
|
request = SimpleNamespace(
|
|
tool_choice=None,
|
|
tools=None,
|
|
chat_template=None,
|
|
chat_template_kwargs=None,
|
|
reasoning_effort=None,
|
|
messages=[],
|
|
add_generation_prompt=False,
|
|
continue_final_message=False,
|
|
add_special_tokens=False,
|
|
request_id="speaker-test",
|
|
)
|
|
|
|
result = asyncio.run(serving_chat.create_chat_completion(request))
|
|
|
|
assert result == "error-response"
|
|
serving_chat.create_error_response.assert_called_once_with("Invalid speaker 'uncle_fu'. Supported: ethan, vivian")
|