vignette/scripts/test_melotts_server.py
2026-08-09 18:22:03 +09:00

141 lines
5.6 KiB
Python

from __future__ import annotations
import importlib.util
import io
import math
import sys
import unittest
import wave
from pathlib import Path
SCRIPT_PATH = Path(__file__).with_name("melotts-server.py")
SPEC = importlib.util.spec_from_file_location("melotts_server", SCRIPT_PATH)
assert SPEC is not None and SPEC.loader is not None
MODULE = importlib.util.module_from_spec(SPEC)
sys.modules[SPEC.name] = MODULE
SPEC.loader.exec_module(MODULE)
class FakeSynthesizer:
sample_rate = 44_100
def __init__(self) -> None:
self.calls: list[tuple[str, float]] = []
def synthesize(self, text: str, *, speed: float):
self.calls.append((text, speed))
return [0.0, 0.5, -0.5, 1.0, -1.0]
class ClampSpeedTest(unittest.TestCase):
def test_in_range_speed_is_kept(self) -> None:
self.assertEqual(MODULE.clamp_speed(1.25), 1.25)
def test_out_of_range_speed_is_folded(self) -> None:
self.assertEqual(MODULE.clamp_speed(99), MODULE.MAX_SPEED)
self.assertEqual(MODULE.clamp_speed(-99), MODULE.MIN_SPEED)
def test_non_numeric_and_nan_fall_back_to_default(self) -> None:
for value in ("fast", None, [], float("nan"), float("inf")):
with self.subTest(value=value):
self.assertEqual(MODULE.clamp_speed(value), MODULE.DEFAULT_SPEED)
class ParseRequestTest(unittest.TestCase):
def test_valid_request_is_normalized(self) -> None:
request = MODULE.parse_tts_request({"text": " 안녕하세요 ", "speed": 1.1})
self.assertEqual(request.text, "안녕하세요")
self.assertAlmostEqual(request.speed, 1.1)
def test_speed_defaults_when_absent(self) -> None:
request = MODULE.parse_tts_request({"text": "안녕"})
self.assertEqual(request.speed, MODULE.DEFAULT_SPEED)
def test_empty_or_missing_text_fails_closed(self) -> None:
for payload in ({}, {"text": ""}, {"text": " "}, {"text": 5}, []):
with self.subTest(payload=payload):
with self.assertRaises(MODULE.TtsError):
MODULE.parse_tts_request(payload)
def test_oversized_text_is_rejected_with_413(self) -> None:
with self.assertRaises(MODULE.TtsError) as ctx:
MODULE.parse_tts_request({"text": "" * (MODULE.MAX_TEXT_CHARS + 1)})
self.assertEqual(ctx.exception.http_status, 413)
self.assertEqual(ctx.exception.code, "text_too_long")
class EncodeWavTest(unittest.TestCase):
def _read(self, payload: bytes):
with wave.open(io.BytesIO(payload), "rb") as handle:
return handle.getnchannels(), handle.getsampwidth(), handle.getframerate(), handle.readframes(handle.getnframes())
def test_wav_header_is_mono_16bit(self) -> None:
payload = MODULE.encode_wav([0.0, 0.25], 22_050)
channels, width, rate, frames = self._read(payload)
self.assertEqual((channels, width, rate), (1, 2, 22_050))
self.assertEqual(len(frames), 4)
def test_samples_are_clipped_into_int16_range(self) -> None:
payload = MODULE.encode_wav([2.0, -2.0], 16_000)
_, _, _, frames = self._read(payload)
self.assertEqual(
int.from_bytes(frames[0:2], "little", signed=True), 32_767
)
self.assertEqual(
int.from_bytes(frames[2:4], "little", signed=True), -32_767
)
def test_nan_samples_become_silence(self) -> None:
payload = MODULE.encode_wav([math.nan], 16_000)
_, _, _, frames = self._read(payload)
self.assertEqual(int.from_bytes(frames[0:2], "little", signed=True), 0)
def test_empty_audio_still_produces_a_valid_wav(self) -> None:
payload = MODULE.encode_wav([], 16_000)
channels, width, rate, frames = self._read(payload)
self.assertEqual((channels, width, rate), (1, 2, 16_000))
self.assertEqual(frames, b"")
def test_invalid_sample_rate_fails_closed(self) -> None:
with self.assertRaises(MODULE.TtsError):
MODULE.encode_wav([0.0], 0)
class HealthPayloadTest(unittest.TestCase):
def test_health_declares_the_permissive_license_and_no_reference(self) -> None:
payload = MODULE.health_payload(FakeSynthesizer(), speakers=["KR"])
self.assertEqual(payload["status"], "ok")
self.assertEqual(payload["provider"], "melotts")
self.assertEqual(payload["model"], "melotts-korean")
self.assertEqual(payload["license"], "MIT")
self.assertEqual(payload["language"], "KR")
self.assertEqual(payload["sample_rate"], 44_100)
self.assertEqual(payload["speakers"], ["KR"])
self.assertIn("no-external-reference", payload["reference_policy"])
class CliTest(unittest.TestCase):
def test_host_is_loopback_only(self) -> None:
parser = MODULE.build_parser()
with self.assertRaises(SystemExit):
parser.parse_args(["--host", "0.0.0.0"])
def test_defaults_target_the_reserved_sidecar_port(self) -> None:
args = MODULE.build_parser().parse_args([])
self.assertEqual(args.port, MODULE.DEFAULT_PORT)
self.assertEqual(args.language, "KR")
class SynthesisPipelineTest(unittest.TestCase):
def test_request_to_wav_round_trip(self) -> None:
synthesizer = FakeSynthesizer()
request = MODULE.parse_tts_request({"text": "안녕하세요", "speed": 5})
audio = synthesizer.synthesize(request.text, speed=request.speed)
payload = MODULE.encode_wav(audio, synthesizer.sample_rate)
self.assertEqual(synthesizer.calls, [("안녕하세요", MODULE.MAX_SPEED)])
self.assertTrue(payload.startswith(b"RIFF"))
if __name__ == "__main__":
unittest.main()