141 lines
5.6 KiB
Python
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()
|