aboutsummaryrefslogtreecommitdiff
path: root/tests/test_tts.py
diff options
context:
space:
mode:
Diffstat (limited to 'tests/test_tts.py')
-rw-r--r--tests/test_tts.py431
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"