diff options
Diffstat (limited to 'tests/test_tts.py')
| -rw-r--r-- | tests/test_tts.py | 431 |
1 files changed, 417 insertions, 14 deletions
diff --git a/tests/test_tts.py b/tests/test_tts.py index 64b4846..813332d 100644 --- a/tests/test_tts.py +++ b/tests/test_tts.py @@ -1,5 +1,6 @@ """Tests for the TTS client wrappers (language handling and payloads).""" +import io import json import tempfile import unittest @@ -9,7 +10,12 @@ from unittest.mock import MagicMock, patch from converter import config, tts from converter.converter import AudiobookConverter -from converter.tts import FasterTTSClient, QwenTTSClient, normalize_language +from converter.tts import ( + AudioCppTTSClient, + FasterTTSClient, + QwenTTSClient, + normalize_language, +) class NormalizeLanguageTests(unittest.TestCase): @@ -518,47 +524,444 @@ class QwenTTSClientGenerateTests(unittest.TestCase): mock_generate.assert_not_called() -class FasterModeWiringTests(unittest.TestCase): - """AudiobookConverter wiring for the --faster backend.""" +class AudioCppTTSClientHealthTests(unittest.TestCase): + """Connection behavior of the audio.cpp client.""" - def test_faster_mode_uses_faster_client_without_reference(self): + @staticmethod + def _json_response(payload): + response = MagicMock() + response.__enter__.return_value = response + response.read.return_value = json.dumps(payload).encode("utf-8") + return response + + def _get_responses(self, health=None, models=None, voices=None): + """Side effect dispatching GET responses by URL.""" + def _dispatch(request, **_kwargs): + url = request if isinstance(request, str) else request.full_url + if url.endswith("/health"): + return self._json_response(health if health is not None + else {"status": "ok"}) + if url.endswith("/v1/models"): + return self._json_response(models if models is not None else + {"data": [{"id": config.AUDIOCPP_MODEL_ID}]}) + if "/v1/audio/voices" in url: + if voices is Exception: + raise Exception("voices endpoint down") + return self._json_response(voices if voices is not None + else {"voices": ["narrator"]}) + raise AssertionError(f"unexpected URL: {url}") + return _dispatch + + def _client(self, voice=None, language=None, **kwargs): + with patch("converter.tts.urllib.request.urlopen", + side_effect=self._get_responses(**kwargs)): + return AudioCppTTSClient(voice=voice, language=language) + + def test_unreachable_server_raises_with_readme_pointer(self): + import urllib.error + with patch("converter.tts.urllib.request.urlopen", + side_effect=urllib.error.URLError("Connection refused")): + with self.assertRaises(RuntimeError) as ctx: + AudioCppTTSClient() + message = str(ctx.exception) + self.assertIn("not reachable", message) + self.assertIn("README", message) + + def test_unhealthy_status_raises(self): + with self.assertRaises(RuntimeError) as ctx: + self._client(health={"status": "starting"}) + self.assertIn("starting", str(ctx.exception)) + + def test_unknown_model_id_raises_with_configured_ids(self): + with self.assertRaises(RuntimeError) as ctx: + self._client(models={"data": [{"id": "pocket-tts"}, {"id": "other"}]}) + message = str(ctx.exception) + self.assertIn(config.AUDIOCPP_MODEL_ID, message) + self.assertIn("pocket-tts", message) + self.assertIn("other", message) + + def test_healthy_server_speaker_mode_defaults(self): + client = self._client() + self.assertEqual(client.api_url, config.AUDIOCPP_API_URL.rstrip("/")) + self.assertEqual(client.model_id, config.AUDIOCPP_MODEL_ID) + self.assertEqual(client.language, config.LANGUAGE) + self.assertEqual(client.voice, "Vivian") + self.assertFalse(client.preset_mode) + + def test_speaker_mode_uses_configured_speaker(self): + with patch.object(config, "SPEAKER", "uncle_fu"): + client = self._client() + self.assertEqual(client.voice, "Uncle Fu") + + def test_preset_mode_uses_requested_voice(self): + client = self._client(voice="narrator") + self.assertEqual(client.voice, "narrator") + self.assertTrue(client.preset_mode) + + def test_preset_mode_validates_voice_against_server_list(self): + with self.assertRaises(RuntimeError) as ctx: + self._client(voice="ghost", voices={"voices": ["narrator", "obama"]}) + message = str(ctx.exception) + self.assertIn("ghost", message) + self.assertIn("narrator", message) + self.assertIn("obama", message) + + def test_preset_mode_skips_validation_when_voices_endpoint_fails(self): + client = self._client(voice="narrator", voices=Exception) + self.assertEqual(client.voice, "narrator") + + def test_invalid_language_fails_before_connect(self): + with patch("converter.tts.urllib.request.urlopen") as mock_urlopen: + with self.assertRaises(ValueError): + AudioCppTTSClient(language="klingon") + mock_urlopen.assert_not_called() + + def test_explicit_language_normalized(self): + client = self._client(language="ja") + self.assertEqual(client.language, "Japanese") + + def test_seed_resolved_once_per_run(self): + with patch.object(config, "CONSTANT_SEED", True), \ + patch.object(config, "SEED", -1): + client = self._client() + self.assertGreaterEqual(client._seed, 0) + + def test_preset_mode_routes_to_clone_model_when_configured(self): + with patch.object(config, "AUDIOCPP_CLONE_MODEL_ID", "qwen3-tts-clone"): + client = self._client( + voice="narrator", + models={"data": [{"id": "qwen3-tts"}, {"id": "qwen3-tts-clone"}]}) + self.assertEqual(client.model_id, "qwen3-tts-clone") + + def test_preset_mode_falls_back_when_clone_model_not_on_server(self): + with patch.object(config, "AUDIOCPP_CLONE_MODEL_ID", "qwen3-tts-clone"), \ + self.assertLogs("converter.tts", level="WARNING") as logs: + client = self._client( + voice="narrator", + models={"data": [{"id": "qwen3-tts"}, {"id": "pocket-tts"}]}) + self.assertEqual(client.model_id, config.AUDIOCPP_MODEL_ID) + self.assertTrue(any("qwen3-tts-clone" in line for line in logs.output)) + + def test_clone_model_id_ignored_for_speaker_mode(self): + with patch.object(config, "AUDIOCPP_CLONE_MODEL_ID", "qwen3-tts-clone"): + client = self._client( + models={"data": [{"id": "qwen3-tts"}, {"id": "qwen3-tts-clone"}]}) + self.assertEqual(client.model_id, config.AUDIOCPP_MODEL_ID) + + def test_clone_model_id_equal_to_primary_is_noop(self): + with patch.object(config, "AUDIOCPP_CLONE_MODEL_ID", + config.AUDIOCPP_MODEL_ID): + client = self._client(voice="narrator") + self.assertEqual(client.model_id, config.AUDIOCPP_MODEL_ID) + + +class AudioCppTTSClientRequestTests(unittest.TestCase): + """The /v1/audio/speech payload and response validation.""" + + def setUp(self): + self._tmp = tempfile.TemporaryDirectory() + self._chunks = patch.object(tts, "CHUNKS_FOLDER", Path(self._tmp.name)) + self._chunks.start() + self._sleep = patch("converter.tts.time.sleep") + self._sleep.start() + + def tearDown(self): + self._sleep.stop() + self._chunks.stop() + self._tmp.cleanup() + + @staticmethod + def _make_client(preset_mode=False, voice="Vivian", language="English", seed=-1): + client = AudioCppTTSClient.__new__(AudioCppTTSClient) + client.api_url = "http://127.0.0.1:8080" + client.model_id = config.AUDIOCPP_MODEL_ID + client.preset_mode = preset_mode + client.voice = voice + client.language = language + client._seed = seed + return client + + @staticmethod + def _wav_bytes(frames=b"\x01\x00" * 10, rate=tts.SAMPLE_RATE): + buffer = io.BytesIO() + with wave.open(buffer, "wb") as wav_file: + wav_file.setnchannels(1) + wav_file.setsampwidth(2) + wav_file.setframerate(rate) + wav_file.writeframes(frames) + return buffer.getvalue() + + def _post_response(self, body): + response = MagicMock() + response.__enter__.return_value = response + response.read.return_value = body + return response + + def test_payload_includes_model_input_voice_language_and_seed(self): + client = self._make_client(preset_mode=True, voice="narrator", + language="Japanese", seed=1234) + with patch("converter.tts.urllib.request.urlopen", + return_value=self._post_response(self._wav_bytes())) as mock_urlopen: + client._request_wav("Hello world.") + request = mock_urlopen.call_args[0][0] + self.assertEqual(request.full_url, + "http://127.0.0.1:8080/v1/audio/speech") + payload = json.loads(request.data.decode("utf-8")) + self.assertEqual(payload["model"], config.AUDIOCPP_MODEL_ID) + self.assertEqual(payload["input"], "Hello world.") + self.assertEqual(payload["voice"], "narrator") + self.assertEqual(payload["language"], "Japanese") + self.assertEqual(payload["seed"], 1234) + self.assertNotIn("instructions", payload) + + def test_speaker_mode_sends_instruct(self): + client = self._make_client(preset_mode=False) + with patch("converter.tts.urllib.request.urlopen", + return_value=self._post_response(self._wav_bytes())) as mock_urlopen: + client._request_wav("Hello.") + payload = json.loads(mock_urlopen.call_args[0][0].data.decode("utf-8")) + self.assertEqual(payload["instructions"], config.INSTRUCT) + + def test_non_wav_response_rejected(self): + client = self._make_client() + for body in (b"", b"RIFFxxxx", b"MP3DATA-MP3DATA", b"RIFF\x00\x00\x00\x00mpeg"): + with patch("converter.tts.urllib.request.urlopen", + return_value=self._post_response(body)): + with self.assertRaises(RuntimeError): + client._request_wav("Hello.") + + def test_http_error_body_surfaced(self): + import urllib.error + client = self._make_client() + error = urllib.error.HTTPError( + "http://127.0.0.1:8080/v1/audio/speech", 500, + "Server Error", {}, io.BytesIO(b'{"error":"bad voice"}')) + with patch("converter.tts.urllib.request.urlopen", side_effect=error): + with self.assertRaises(RuntimeError) as ctx: + client._request_wav("Hello.") + self.assertIn("500", str(ctx.exception)) + self.assertIn("bad voice", str(ctx.exception)) + + def test_transient_failure_is_retried(self): + client = self._make_client() + wav = self._wav_bytes() + with patch.object(client, "_request_wav", + side_effect=[RuntimeError("boom"), wav]) as mock_request: + result = client.generate_chunk("Hello.", 1) + self.assertIsNotNone(result) + self.assertEqual(mock_request.call_count, 2) + + def test_exhausted_retries_fail_the_chunk(self): + client = self._make_client() + with patch.object(client, "_request_wav", + side_effect=RuntimeError("down")) as mock_request: + result = client.generate_chunk("Hello.", 1) + self.assertIsNone(result) + self.assertEqual(mock_request.call_count, config.MAX_RETRIES) + + def test_empty_text_fails_the_chunk(self): + client = self._make_client() + with patch.object(client, "_request_wav") as mock_request: + result = client.generate_chunk(" ", 1) + self.assertIsNone(result) + mock_request.assert_not_called() + + def test_generate_chunk_writes_valid_wav(self): + client = self._make_client() + frames = b"\x01\x00" * 100 + with patch.object(client, "_request_wav", return_value=self._wav_bytes(frames)): + result = client.generate_chunk("Hello world.", 1) + self.assertIsNotNone(result) + path = Path(result) + self.assertEqual(path.name, "chunk_0001.wav") + with wave.open(str(path), "rb") as wav_file: + self.assertEqual(wav_file.getnchannels(), 1) + self.assertEqual(wav_file.getsampwidth(), 2) + self.assertEqual(wav_file.getframerate(), tts.SAMPLE_RATE) + self.assertEqual(wav_file.readframes(wav_file.getnframes()), frames) + + def test_long_text_is_subchunked_and_concatenated_in_order(self): + client = self._make_client() + sentences = [" ".join(f"word{i}" for i in range(6)) + "." for _ in range(3)] + text = " ".join(sentences) + parts = [self._wav_bytes(b"\x01\x00" * 10), + self._wav_bytes(b"\x02\x00" * 20), + self._wav_bytes(b"\x03\x00" * 30)] + with patch.object(tts, "MAX_REQUEST_WORDS", 10), \ + patch.object(client, "_request_wav", side_effect=parts) as mock_request: + result = client.generate_chunk(text, 1) + self.assertEqual(mock_request.call_count, 3) + with wave.open(str(Path(result)), "rb") as wav_file: + self.assertEqual(wav_file.readframes(wav_file.getnframes()), + b"\x01\x00" * 10 + b"\x02\x00" * 20 + b"\x03\x00" * 30) + + def test_stale_chunk_files_are_removed(self): + stale = Path(self._tmp.name) / "chunk_0001.mp3" + stale.write_bytes(b"old") + client = self._make_client() + with patch.object(client, "_request_wav", return_value=self._wav_bytes()): + client.generate_chunk("Hello.", 1) + remaining = sorted(path.name for path in Path(self._tmp.name).glob("chunk_0001.*")) + self.assertEqual(remaining, ["chunk_0001.wav"]) + + +class AudioCppTTSClientTruncationTests(unittest.TestCase): + """Audio far shorter than its text implies fails the request.""" + + def setUp(self): + self._tmp = tempfile.TemporaryDirectory() + self._chunks = patch.object(tts, "CHUNKS_FOLDER", Path(self._tmp.name)) + self._chunks.start() + + def tearDown(self): + self._chunks.stop() + self._tmp.cleanup() + + def _make_client(self): + client = AudioCppTTSClient.__new__(AudioCppTTSClient) + client.api_url = "http://127.0.0.1:8080" + client.model_id = config.AUDIOCPP_MODEL_ID + client.preset_mode = True + client.voice = "narrator" + client.language = "English" + client._seed = -1 + return client + + @staticmethod + def _wav_bytes(frames): + buffer = io.BytesIO() + with wave.open(buffer, "wb") as wav_file: + wav_file.setnchannels(1) + wav_file.setsampwidth(2) + wav_file.setframerate(tts.SAMPLE_RATE) + wav_file.writeframes(frames) + return buffer.getvalue() + + def test_truncated_wav_fails_the_chunk(self): + client = self._make_client() + text = " ".join(f"word{i}" for i in range(12)) + wav = self._wav_bytes(b"\x01\x00" * 24) # 0.001s for ~4.8s of speech + with patch.object(client, "_request_wav", return_value=wav), \ + self.assertLogs("converter.tts", level="ERROR") as logs: + result = client.generate_chunk(text, 1) + self.assertIsNone(result) + self.assertTrue(any("truncated" in line for line in logs.output)) + + def test_full_length_wav_passes(self): + client = self._make_client() + text = " ".join(f"word{i}" for i in range(12)) + # 12 words -> expected 4.8s, half is 2.4s -> 2.5s of audio passes. + wav = self._wav_bytes(b"\x01\x00" * int(2.5 * tts.SAMPLE_RATE)) + with patch.object(client, "_request_wav", return_value=wav): + result = client.generate_chunk(text, 1) + self.assertIsNotNone(result) + + +class BackendWiringTests(unittest.TestCase): + """AudiobookConverter wiring for the --backend selector.""" + + def test_faster_backend_uses_faster_client_without_reference(self): with patch("converter.converter.FasterTTSClient") as mock_faster, \ - patch("converter.converter.QwenTTSClient") as mock_qwen: + patch("converter.converter.QwenTTSClient") as mock_qwen, \ + patch("converter.converter.AudioCppTTSClient") as mock_audiocpp: AudiobookConverter(voice_mode=tts.VOICE_MODE_CLONE, - faster=True, faster_voice="narrator") + backend=tts.BACKEND_FASTER, voice="narrator") mock_faster.assert_called_once_with(voice="narrator") mock_qwen.assert_not_called() + mock_audiocpp.assert_not_called() + + def test_audiocpp_backend_with_voice_uses_audiocpp_client(self): + with patch("converter.converter.FasterTTSClient") as mock_faster, \ + patch("converter.converter.QwenTTSClient") as mock_qwen, \ + patch("converter.converter.AudioCppTTSClient") as mock_audiocpp: + AudiobookConverter(voice_mode=tts.VOICE_MODE_CLONE, + backend=tts.BACKEND_AUDIOCPP, voice="narrator", + language="ja") + mock_audiocpp.assert_called_once_with(voice="narrator", language="Japanese") + mock_faster.assert_not_called() + mock_qwen.assert_not_called() + + def test_audiocpp_backend_without_voice_uses_audiocpp_client(self): + with patch("converter.converter.AudioCppTTSClient") as mock_audiocpp: + AudiobookConverter(voice_mode=tts.VOICE_MODE_CUSTOM, + backend=tts.BACKEND_AUDIOCPP) + mock_audiocpp.assert_called_once_with(voice=None, language=config.LANGUAGE) - def test_non_faster_clone_mode_still_requires_reference(self): + def test_gradio_backend_uses_qwen_client(self): + with patch("converter.converter.FasterTTSClient") as mock_faster, \ + patch("converter.converter.QwenTTSClient") as mock_qwen, \ + patch("converter.converter.AudioCppTTSClient") as mock_audiocpp: + AudiobookConverter(voice_mode=tts.VOICE_MODE_CUSTOM) + mock_qwen.assert_called_once() + mock_faster.assert_not_called() + mock_audiocpp.assert_not_called() + + def test_gradio_clone_mode_still_requires_reference(self): with patch("converter.converter.QwenTTSClient"): with self.assertRaises(ValueError): AudiobookConverter(voice_mode=tts.VOICE_MODE_CLONE) - def test_faster_mode_still_validates_other_settings(self): + def test_audiocpp_clone_mode_does_not_require_reference(self): + # Cloning is server-side for the audiocpp backend, so the + # clone-mode voice can be selected without local reference audio. + with patch("converter.converter.AudioCppTTSClient"): + converter = AudiobookConverter(voice_mode=tts.VOICE_MODE_CLONE, + backend=tts.BACKEND_AUDIOCPP, + voice="narrator") + self.assertIsNone(converter.voice_clone_ref_audio) + + def test_faster_backend_still_validates_other_settings(self): with patch("converter.converter.FasterTTSClient"): with self.assertRaises(ValueError): - AudiobookConverter(faster=True, speed=0) + AudiobookConverter(backend=tts.BACKEND_FASTER, speed=0) + with self.assertRaises(ValueError): + AudiobookConverter(backend=tts.BACKEND_FASTER, language="klingon") + + def test_audiocpp_backend_still_validates_other_settings(self): + with patch("converter.converter.AudioCppTTSClient"): + with self.assertRaises(ValueError): + AudiobookConverter(backend=tts.BACKEND_AUDIOCPP, speed=0) with self.assertRaises(ValueError): - AudiobookConverter(faster=True, language="klingon") + AudiobookConverter(backend=tts.BACKEND_AUDIOCPP, language="klingon") - def _faster_converter(self, faster_voice=None): + def _faster_converter(self, voice=None): with patch("converter.converter.FasterTTSClient"): return AudiobookConverter(voice_mode=tts.VOICE_MODE_CLONE, - faster=True, faster_voice=faster_voice) + backend=tts.BACKEND_FASTER, voice=voice) + + def _audiocpp_converter(self, voice=None): + with patch("converter.converter.AudioCppTTSClient"): + return AudiobookConverter( + voice_mode=tts.VOICE_MODE_CLONE if voice else tts.VOICE_MODE_CUSTOM, + backend=tts.BACKEND_AUDIOCPP, voice=voice) def test_narrator_tag_uses_faster_voice_name(self): - converter = self._faster_converter(faster_voice="male_richard_poe") + converter = self._faster_converter(voice="male_richard_poe") self.assertEqual(converter._narrator_tag(), "male_richard_poe") def test_narrator_tag_falls_back_to_config_voice(self): converter = self._faster_converter() self.assertEqual(converter._narrator_tag(), config.FASTER_VOICE) + def test_narrator_tag_audiocpp_uses_voice_name(self): + converter = self._audiocpp_converter(voice="female_narrator") + self.assertEqual(converter._narrator_tag(), "female_narrator") + + def test_narrator_tag_audiocpp_falls_back_to_speaker(self): + converter = self._audiocpp_converter() + self.assertEqual(converter._narrator_tag(), "Vivian") + def test_banner_and_narrator_work_without_reference_audio(self): - converter = self._faster_converter(faster_voice="male_richard_poe") + converter = self._faster_converter(voice="male_richard_poe") converter._print_banner() # must not raise (regression: Path(None)) self.assertIsNone(converter.voice_clone_ref_audio) + def test_audiocpp_banner_prints_without_reference_audio(self): + converter = self._audiocpp_converter(voice="narrator") + converter._print_banner() # must not raise + converter = self._audiocpp_converter() + converter._print_banner() + def test_non_faster_narrator_tag_unchanged(self): with tempfile.TemporaryDirectory() as tmp: ref = Path(tmp) / "ref.wav" |
