aboutsummaryrefslogtreecommitdiff
path: root/tests
diff options
context:
space:
mode:
authorhistoria <historiavg@proton.me>2026-08-18 17:58:04 -0400
committerhistoria <historiavg@proton.me>2026-08-18 17:58:04 -0400
commite56754498f1c6b2a9dabb62529f783e73fae8e6b (patch)
tree20ea86a16d6c51c4bfd793cea56fd76ae9985671 /tests
parent50f1825f05972e3685c55beb10c288899959b2e5 (diff)
downloadtts-audiobook-generator-e56754498f1c6b2a9dabb62529f783e73fae8e6b.tar.gz
feat: language parameter for potential accent tuning
Diffstat (limited to 'tests')
-rw-r--r--tests/cover_test.pngbin6285 -> 6801 bytes
-rw-r--r--tests/gen_test_cover.py2
-rw-r--r--tests/test_converter.py22
-rw-r--r--tests/test_tts.py155
4 files changed, 178 insertions, 1 deletions
diff --git a/tests/cover_test.png b/tests/cover_test.png
index c6c4bc6..0c252db 100644
--- a/tests/cover_test.png
+++ b/tests/cover_test.png
Binary files differ
diff --git a/tests/gen_test_cover.py b/tests/gen_test_cover.py
index 7c9347e..292469a 100644
--- a/tests/gen_test_cover.py
+++ b/tests/gen_test_cover.py
@@ -3,6 +3,6 @@ from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from converter.cover import generate_cover
-p = generate_cover('Your Book Title Here',
+p = generate_cover('The Count of Monte Cristo',
Path(__file__).resolve().parent / 'cover_test.png')
print('written:', p)
diff --git a/tests/test_converter.py b/tests/test_converter.py
index 91ce699..351849a 100644
--- a/tests/test_converter.py
+++ b/tests/test_converter.py
@@ -34,6 +34,27 @@ class ConfigurationValidationTests(unittest.TestCase):
with self.assertRaises(ValueError):
AudiobookConverter(output_format="wma")
+ def test_unknown_language_rejected(self):
+ with self.assertRaises(ValueError):
+ AudiobookConverter(language="klingon")
+
+ def test_language_defaults_to_config(self):
+ with patch("converter.converter.QwenTTSClient") as mock_tts:
+ AudiobookConverter()
+ self.assertEqual(mock_tts.call_args.kwargs["language"], "English")
+
+ def test_output_format_defaults_to_config(self):
+ with patch("converter.converter.QwenTTSClient"):
+ converter = AudiobookConverter()
+ self.assertEqual(converter.output_format, config.AUDIO_FORMAT)
+ self.assertEqual(config.AUDIO_FORMAT, "m4b")
+
+ def test_language_normalized_before_tts_client(self):
+ with patch("converter.converter.QwenTTSClient") as mock_tts:
+ converter = AudiobookConverter(language="ja")
+ self.assertEqual(converter.language, "Japanese")
+ self.assertEqual(mock_tts.call_args.kwargs["language"], "Japanese")
+
class FindExistingOutputsTests(unittest.TestCase):
def setUp(self):
@@ -127,6 +148,7 @@ class RunOverwritePromptTests(unittest.TestCase):
self.converter.speed = 1.0
self.converter.single_file = False
self.converter.output_format = "mp3"
+ self.converter.language = "English"
self.converted = []
self.converter.convert_book = (
lambda file_path, output_name=None:
diff --git a/tests/test_tts.py b/tests/test_tts.py
new file mode 100644
index 0000000..b605daa
--- /dev/null
+++ b/tests/test_tts.py
@@ -0,0 +1,155 @@
+"""Tests for the Qwen TTS client wrapper (language handling and payloads)."""
+
+import tempfile
+import unittest
+from pathlib import Path
+from unittest.mock import MagicMock, patch
+
+from converter import config
+from converter.tts import QwenTTSClient, normalize_language
+
+
+class NormalizeLanguageTests(unittest.TestCase):
+ def test_display_names_case_insensitive(self):
+ self.assertEqual(normalize_language("english"), "English")
+ self.assertEqual(normalize_language("ENGLISH"), "English")
+ self.assertEqual(normalize_language(" Japanese "), "Japanese")
+
+ def test_auto_accepted(self):
+ self.assertEqual(normalize_language("auto"), "Auto")
+ self.assertEqual(normalize_language("Auto"), "Auto")
+
+ def test_iso_aliases(self):
+ self.assertEqual(normalize_language("en"), "English")
+ self.assertEqual(normalize_language("ja"), "Japanese")
+ self.assertEqual(normalize_language("zh"), "Chinese")
+ self.assertEqual(normalize_language("ko"), "Korean")
+ self.assertEqual(normalize_language("de"), "German")
+ self.assertEqual(normalize_language("fr"), "French")
+ self.assertEqual(normalize_language("ru"), "Russian")
+ self.assertEqual(normalize_language("pt"), "Portuguese")
+ self.assertEqual(normalize_language("es"), "Spanish")
+ self.assertEqual(normalize_language("it"), "Italian")
+
+ def test_all_supported_languages_round_trip(self):
+ for name in config.TTS_LANGUAGES:
+ self.assertEqual(normalize_language(name.lower()), name)
+
+ def test_unknown_language_rejected_with_guidance(self):
+ with self.assertRaises(ValueError) as ctx:
+ normalize_language("klingon")
+ message = str(ctx.exception)
+ self.assertIn("klingon", message)
+ self.assertIn("English", message)
+
+ def test_none_and_empty_rejected(self):
+ with self.assertRaises(ValueError):
+ normalize_language(None)
+ with self.assertRaises(ValueError):
+ normalize_language(" ")
+
+
+class QwenTTSClientLanguageTests(unittest.TestCase):
+ """Language validation and defaults, without touching the network."""
+
+ def _make_client(self, **kwargs):
+ with patch.object(QwenTTSClient, "_connect"):
+ return QwenTTSClient(**kwargs)
+
+ def test_default_follows_config_for_each_mode(self):
+ custom = self._make_client(voice_mode=config.VOICE_MODE_CUSTOM)
+ self.assertEqual(custom.language, config.CUSTOM_VOICE_LANGUAGE)
+ clone = self._make_client(voice_mode=config.VOICE_MODE_CLONE,
+ voice_clone_ref_audio="ref.wav")
+ self.assertEqual(clone.language, config.VOICE_CLONE_LANGUAGE)
+
+ def test_explicit_language_normalized(self):
+ client = self._make_client(voice_mode=config.VOICE_MODE_CUSTOM, language="ja")
+ self.assertEqual(client.language, "Japanese")
+
+ def test_invalid_language_fails_before_connect(self):
+ with patch.object(QwenTTSClient, "_connect") as mock_connect:
+ with self.assertRaises(ValueError):
+ QwenTTSClient(language="klingon")
+ mock_connect.assert_not_called()
+
+
+class PayloadLanguageTests(unittest.TestCase):
+ """The language must reach the API payload in every endpoint variant."""
+
+ def setUp(self):
+ self._tmp = tempfile.TemporaryDirectory()
+ self.ref_audio = Path(self._tmp.name) / "reference.wav"
+ self.ref_audio.write_bytes(b"x")
+
+ def tearDown(self):
+ self._tmp.cleanup()
+
+ def _custom_client(self, language, endpoint, api_info=None):
+ client = QwenTTSClient.__new__(QwenTTSClient)
+ client.voice_mode = config.VOICE_MODE_CUSTOM
+ client.language = language
+ client.api_info = api_info if api_info is not None else {
+ "named_endpoints": {endpoint: {}}
+ }
+ client.client = MagicMock()
+ return client
+
+ def _clone_client(self, language, endpoint, api_info=None, ref_text="hello"):
+ client = QwenTTSClient.__new__(QwenTTSClient)
+ client.voice_mode = config.VOICE_MODE_CLONE
+ client.language = language
+ client.voice_clone_ref_audio = str(self.ref_audio)
+ client.voice_clone_ref_text = ref_text
+ client.clone_api_info = api_info if api_info is not None else {
+ "named_endpoints": {endpoint: {}}
+ }
+ client.clone_client = MagicMock()
+ client._ref_audio_filedata = {"dummy": "payload"}
+ return client
+
+ def test_custom_voice_run_instruct_uses_language(self):
+ client = self._custom_client("Japanese", "/run_instruct")
+ client._generate_custom_voice("text")
+ kwargs = client.client.predict.call_args.kwargs
+ self.assertEqual(kwargs["lang_disp"], "Japanese")
+
+ def test_custom_voice_alt_endpoint_uses_language(self):
+ client = self._custom_client("French", "/run_custom_voice")
+ client._generate_custom_voice("text")
+ kwargs = client.client.predict.call_args.kwargs
+ self.assertEqual(kwargs["language"], "French")
+
+ def test_voice_clone_run_voice_clone_uses_language(self):
+ client = self._clone_client("Japanese", "/run_voice_clone")
+ client._generate_voice_clone("text")
+ kwargs = client.clone_client.predict.call_args.kwargs
+ self.assertEqual(kwargs["lang_disp"], "Japanese")
+
+ def test_voice_clone_alt_endpoint_uses_language(self):
+ client = self._clone_client("Korean", "/generate_voice_clone")
+ client._generate_voice_clone("text")
+ kwargs = client.clone_client.predict.call_args.kwargs
+ self.assertEqual(kwargs["language"], "Korean")
+
+ def test_voice_clone_alt_endpoint_includes_optional_params(self):
+ api_info = {
+ "named_endpoints": {
+ "/generate_voice_clone": {
+ "parameters": [
+ {"parameter_name": "model_size"},
+ {"parameter_name": "seed"},
+ ]
+ }
+ }
+ }
+ client = self._clone_client("English", "/generate_voice_clone", api_info=api_info)
+ client._generate_voice_clone("text")
+ kwargs = client.clone_client.predict.call_args.kwargs
+ self.assertEqual(kwargs["model_size"], config.VOICE_CLONE_MODEL_SIZE)
+ self.assertEqual(kwargs["seed"], config.VOICE_CLONE_SEED)
+ self.assertNotIn("max_chunk_chars", kwargs)
+
+
+if __name__ == "__main__":
+ unittest.main()