diff options
| author | historia <historiavg@proton.me> | 2026-08-20 22:58:52 -0400 |
|---|---|---|
| committer | historia <historiavg@proton.me> | 2026-08-20 22:58:52 -0400 |
| commit | 5c3df0a434059bd0d541bda35a51e49e3c44dd55 (patch) | |
| tree | 1d18e5f41ed9fc1184275a2a2b1a6555dffa4dc4 /tests/test_make_audiocpp_server_json.py | |
| parent | 0c197324f5444b448c285d2a57bd0a5834c2fc84 (diff) | |
| download | tts-audiobook-generator-5c3df0a434059bd0d541bda35a51e49e3c44dd55.tar.gz | |
feat: experimental support for non-qwen models
Diffstat (limited to 'tests/test_make_audiocpp_server_json.py')
| -rw-r--r-- | tests/test_make_audiocpp_server_json.py | 222 |
1 files changed, 210 insertions, 12 deletions
diff --git a/tests/test_make_audiocpp_server_json.py b/tests/test_make_audiocpp_server_json.py index bdc3604..f9fa794 100644 --- a/tests/test_make_audiocpp_server_json.py +++ b/tests/test_make_audiocpp_server_json.py @@ -20,6 +20,13 @@ FAKE_CONFIG = ( "CHUNK_SIZE = 250\n" ) +FAKE_CONFIG_WITH_MODEL_IDS = ( + 'AUDIOCPP_API_URL = "http://127.0.0.1:9999" # audio.cpp audiocpp_server\n' + "\n" + 'AUDIOCPP_MODEL_ID = "qwen" # server entry for speaker mode\n' + 'AUDIOCPP_CLONE_MODEL_ID = "qwen-clone"\n' +) + class FindWavFilesTests(unittest.TestCase): def setUp(self): @@ -117,6 +124,92 @@ class UpdateConfigPortTests(unittest.TestCase): 8080, config_path=Path(self._tmp.name) / "nope.py")) +class UpdateConfigModelIdsTests(unittest.TestCase): + def setUp(self): + self._tmp = tempfile.TemporaryDirectory() + self.config_path = Path(self._tmp.name) / "config.py" + self.config_path.write_text(FAKE_CONFIG_WITH_MODEL_IDS, + encoding="utf-8") + + def tearDown(self): + self._tmp.cleanup() + + def test_rewrites_both_ids_preserving_lines(self): + changed = make_server.update_config_model_ids( + "higgs", "higgs", config_path=self.config_path) + self.assertTrue(changed) + text = self.config_path.read_text(encoding="utf-8") + self.assertIn('AUDIOCPP_MODEL_ID = "higgs" # server entry for speaker mode', + text) + self.assertIn('AUDIOCPP_CLONE_MODEL_ID = "higgs"', text) + self.assertIn('AUDIOCPP_API_URL = "http://127.0.0.1:9999"', text) + + def test_clone_id_optional(self): + changed = make_server.update_config_model_ids( + "voxcpm2", config_path=self.config_path) + self.assertTrue(changed) + text = self.config_path.read_text(encoding="utf-8") + self.assertIn('AUDIOCPP_MODEL_ID = "voxcpm2"', text) + self.assertIn('AUDIOCPP_CLONE_MODEL_ID = "qwen-clone"', text) + + def test_returns_false_when_ids_unchanged(self): + changed = make_server.update_config_model_ids( + "qwen", "qwen-clone", config_path=self.config_path) + self.assertFalse(changed) + self.assertEqual(self.config_path.read_text(encoding="utf-8"), + FAKE_CONFIG_WITH_MODEL_IDS) + + def test_returns_false_when_lines_missing(self): + path = Path(self._tmp.name) / "other.py" + path.write_text('CHUNK_SIZE = 250\n', encoding="utf-8") + self.assertFalse(make_server.update_config_model_ids( + "higgs", "higgs", config_path=path)) + + def test_returns_false_when_file_missing(self): + self.assertFalse(make_server.update_config_model_ids( + "higgs", "higgs", + config_path=Path(self._tmp.name) / "nope.py")) + + +class BuildSingleFamilyServerConfigTests(unittest.TestCase): + def test_single_entry_with_presets(self): + presets = {"narrator": {"voice_ref": "/x.wav", + "reference_text": "hi"}} + server_config = make_server.build_single_family_server_config( + host="127.0.0.1", port=8080, backend="cuda", lazy_load=False, + family="higgs_audio_tts", model_id="higgs", + model_path="models/Higgs-Audio-v3-TTS-4B-GGUF", + voice_presets=presets) + self.assertEqual(server_config["host"], "127.0.0.1") + self.assertEqual(server_config["port"], 8080) + self.assertEqual(server_config["backend"], "cuda") + self.assertFalse(server_config["lazy_load"]) + self.assertEqual(len(server_config["models"]), 1) + entry = server_config["models"][0] + self.assertEqual(entry["id"], "higgs") + self.assertEqual(entry["family"], "higgs_audio_tts") + self.assertEqual(entry["path"], "models/Higgs-Audio-v3-TTS-4B-GGUF") + self.assertEqual(entry["task"], "tts") + self.assertEqual(entry["mode"], "offline") + self.assertEqual(entry["voice_presets"], presets) + + def test_no_presets_omits_key(self): + server_config = make_server.build_single_family_server_config( + host="127.0.0.1", port=8080, backend="cpu", lazy_load=True, + family="index_tts2", model_id="indextts2", + model_path="models/IndexTTS2-GGUF", voice_presets={}) + self.assertNotIn("voice_presets", server_config["models"][0]) + + def test_family_entries_reference_real_families(self): + for entry in make_server.FAMILY_ENTRIES: + if entry["key"] == make_server.FAMILY_QWEN3_TTS: + continue + self.assertIn("install", entry) + self.assertIn("default_id", entry) + self.assertIn("default_path", entry) + self.assertIn("family", entry) + + class BuildVoicePresetsTests(unittest.TestCase): def setUp(self): self._tmp = tempfile.TemporaryDirectory() @@ -279,7 +372,8 @@ class MainTests(unittest.TestCase): def _defaults(self, models="", host="", port="", backend="", lazy="", custom_path="", clone_path="", wav_dir="", confirm="y", prefix=()): - return list(prefix) + [models, host, port, backend, lazy, + # First input selects the model family (default: Qwen3-TTS). + return list(prefix) + ["", models, host, port, backend, lazy, custom_path, clone_path, wav_dir, confirm] def test_default_run_hosts_both_models(self): @@ -313,7 +407,7 @@ class MainTests(unittest.TestCase): def test_clone_only_with_positional_wav_dir(self): (self.folder / "narrator.wav").write_bytes(b"x") (self.folder / "alpha.wav").write_bytes(b"x") - inputs = ["3", "", "", "", "", "", "y"] + inputs = ["", "3", "", "", "", "", "", "y"] exit_code = self._run( [str(self.folder), "--output", str(self.output)], inputs=inputs, @@ -331,7 +425,7 @@ class MainTests(unittest.TestCase): "reference_text": "transcript of narrator.wav"}) def test_custom_only_single_model(self): - inputs = ["", "", "", "", "", "y"] + inputs = ["", "", "", "", "", "", "y"] exit_code = self._run( ["--output", str(self.output), "--models", "custom"], inputs=inputs) @@ -343,7 +437,7 @@ class MainTests(unittest.TestCase): def test_duplicate_ids_prompt_for_distinct_clone_id(self): with patch.object(config, "AUDIOCPP_MODEL_ID", "qwen"), \ patch.object(config, "AUDIOCPP_CLONE_MODEL_ID", "qwen"): - inputs = ["1", "qwen-clone-2", "", "", "", "", "", "", "", "y"] + inputs = ["", "1", "qwen-clone-2", "", "", "", "", "", "", "", "y"] exit_code = self._run(["--output", str(self.output)], inputs=inputs) self.assertEqual(exit_code, 0) @@ -362,7 +456,7 @@ class MainTests(unittest.TestCase): def test_port_sync_accepted_updates_config(self): with patch.object(config, "AUDIOCPP_API_URL", "http://127.0.0.1:9999"): - inputs = ["", "", "y", "", "", "", "", "", "y"] + inputs = ["", "", "", "y", "", "", "", "", "", "y"] exit_code = self._run(["--output", str(self.output), "--port", "8080"], inputs=inputs) @@ -375,7 +469,7 @@ class MainTests(unittest.TestCase): def test_port_sync_declined_keeps_config(self): with patch.object(config, "AUDIOCPP_API_URL", "http://127.0.0.1:9999"): - inputs = ["", "", "n", "", "", "", "", "", "y"] + inputs = ["", "", "", "n", "", "", "", "", "", "y"] exit_code = self._run(["--output", str(self.output), "--port", "8080"], inputs=inputs) @@ -394,7 +488,8 @@ class MainTests(unittest.TestCase): FAKE_CONFIG) def test_invalid_menu_choice_reprompts(self): - inputs = ["9", "", "", "", "", "", "", "", "", "y"] + # Family menu default, then an invalid models-menu choice retried. + inputs = ["", "9", "", "", "", "", "", "", "", "", "y"] exit_code = self._run(["--output", str(self.output)], inputs=inputs) self.assertEqual(exit_code, 0) @@ -435,11 +530,14 @@ class MainTests(unittest.TestCase): self.assertEqual(len(data["models"]), 2) def test_flags_skip_prompts(self): + # Family still asked (no --family flag); port 9000 differs from the + # config port so its sync prompt fires; custom/clone paths and the + # wav dir use their defaults. exit_code = self._run( ["--output", str(self.output), "--models", "both", "--host", "0.0.0.0", "--port", "9000", "--backend", "cpu", "--lazy-load"], - inputs=["y", "", "", "", "y"]) + inputs=["", "y", "", "", "", "y"]) self.assertEqual(exit_code, 0) self.assertIn('"http://127.0.0.1:9000"', self.fake_config.read_text(encoding="utf-8")) @@ -457,6 +555,105 @@ class MainTests(unittest.TestCase): self.assertEqual(ctx.exception.code, 2) +class NonQwenFamilyMainTests(unittest.TestCase): + """The --family flow for clone-only model families.""" + + def setUp(self): + self._tmp = tempfile.TemporaryDirectory() + self.folder = Path(self._tmp.name) + self.output = self.folder / "server.json" + self.fake_config = self.folder / "config.py" + self.fake_config.write_text(FAKE_CONFIG_WITH_MODEL_IDS, + encoding="utf-8") + patcher = patch.object(make_server, "CONFIG_PATH", self.fake_config) + patcher.start() + self.addCleanup(patcher.stop) + + def tearDown(self): + self._tmp.cleanup() + + def _run(self, argv, inputs=None, transcribe=None, whisper="faster_whisper"): + argv = ["make_audiocpp_server_json.py"] + argv + input_effect = inputs if inputs is not None else EOFError + transcribe_effect = transcribe if transcribe is not None else MagicMock() + with patch.object(sys, "argv", argv), \ + patch("builtins.input", side_effect=input_effect), \ + patch.object(make_server, "transcribe_reference_audio", + side_effect=transcribe_effect), \ + patch.object(make_server, "whisper_backend_available", + return_value=whisper): + return make_server.main() + + def test_higgs_family_run(self): + (self.folder / "narrator.wav").write_bytes(b"x") + # Inputs: model-id sync accepted, host, port, backend, lazy, confirm. + inputs = ["y", "", "", "", "", "y"] + exit_code = self._run( + [str(self.folder), "--output", str(self.output), + "--family", "higgs_audio_tts", "--model-id", "higgs", + "--model-path", "models/Higgs-Audio-v3-TTS-4B-GGUF"], + inputs=inputs, + transcribe=lambda path, model_name="base": "a transcript") + self.assertEqual(exit_code, 0) + data = json.loads(self.output.read_text(encoding="utf-8")) + self.assertEqual(len(data["models"]), 1) + entry = data["models"][0] + self.assertEqual(entry["id"], "higgs") + self.assertEqual(entry["family"], "higgs_audio_tts") + self.assertEqual(entry["path"], "models/Higgs-Audio-v3-TTS-4B-GGUF") + self.assertEqual(entry["task"], "tts") + self.assertEqual(entry["mode"], "offline") + self.assertEqual(entry["voice_presets"]["narrator"], + {"voice_ref": str((self.folder / "narrator.wav").resolve()), + "reference_text": "a transcript"}) + # Both converter model ids point at the single server entry. + self.assertIn('AUDIOCPP_MODEL_ID = "higgs"', + self.fake_config.read_text(encoding="utf-8")) + self.assertIn('AUDIOCPP_CLONE_MODEL_ID = "higgs"', + self.fake_config.read_text(encoding="utf-8")) + + def test_model_id_sync_declined_keeps_config(self): + # sync declined, host, port, backend, lazy, wav dir skipped, confirm + inputs = ["n", "", "", "", "", "", "y"] + exit_code = self._run( + ["--output", str(self.output), "--family", "voxcpm2", + "--model-id", "voxcpm2", "--model-path", "models/VoxCPM2-GGUF"], + inputs=inputs) + self.assertEqual(exit_code, 0) + text = self.fake_config.read_text(encoding="utf-8") + self.assertIn('AUDIOCPP_MODEL_ID = "qwen"', text) + self.assertIn('AUDIOCPP_CLONE_MODEL_ID = "qwen-clone"', text) + data = json.loads(self.output.read_text(encoding="utf-8")) + self.assertEqual(data["models"][0]["family"], "voxcpm2") + + def test_no_voice_presets_warns(self): + buf = io.StringIO() + # sync accepted, host, port, backend, lazy, wav dir skipped, confirm + with patch.object(sys, "argv", + ["make_audiocpp_server_json.py", + "--output", str(self.output), + "--family", "index_tts2", "--model-id", "indextts2", + "--model-path", "models/IndexTTS2-GGUF"]), \ + patch("builtins.input", side_effect=["y", "", "", "", "", "", "y"]), \ + patch.object(make_server, "transcribe_reference_audio"), \ + patch.object(make_server, "whisper_backend_available", + return_value="faster_whisper"), \ + redirect_stdout(buf): + code = make_server.main() + self.assertEqual(code, 0) + out = buf.getvalue() + self.assertIn("No voice presets were configured", out) + self.assertIn("model_manager_v2.py install index_tts2_q8_0", out) + data = json.loads(self.output.read_text(encoding="utf-8")) + self.assertNotIn("voice_presets", data["models"][0]) + + def test_models_flag_rejected_for_non_qwen_family(self): + with self.assertRaises(SystemExit) as ctx: + self._run(["--output", str(self.output), + "--family", "higgs_audio_tts", "--models", "both"]) + self.assertEqual(ctx.exception.code, 2) + + class TranscriptWarningTests(unittest.TestCase): """Empty transcripts and a missing Whisper backend produce loud warnings.""" @@ -489,10 +686,11 @@ class TranscriptWarningTests(unittest.TestCase): def test_empty_transcript_prints_loud_end_warning(self): (self.folder / "narrator.wav").write_bytes(b"x") (self.folder / "alpha.wav").write_bytes(b"x") - # Clone-only run (menu choice 3); transcribe returns None (empty). + # Qwen family default, clone-only run (menu choice 3); transcribe + # returns None (empty). code, out = self._run_capturing( [str(self.folder), "--output", str(self.output)], - inputs=["3", "", "", "", "", "", "y"], + inputs=["", "3", "", "", "", "", "", "y"], transcribe=lambda path, model_name="base": None, whisper="faster_whisper") self.assertEqual(code, 0) @@ -505,7 +703,7 @@ class TranscriptWarningTests(unittest.TestCase): (self.folder / "narrator.wav").write_bytes(b"x") code, out = self._run_capturing( [str(self.folder), "--output", str(self.output)], - inputs=["3", "", "", "", "", "", "y"], + inputs=["", "3", "", "", "", "", "", "y"], transcribe=lambda path, model_name="base": "a transcript", whisper=None) self.assertEqual(code, 0) @@ -516,7 +714,7 @@ class TranscriptWarningTests(unittest.TestCase): (self.folder / "narrator.wav").write_bytes(b"x") code, out = self._run_capturing( [str(self.folder), "--output", str(self.output)], - inputs=["3", "", "", "", "", "", "y"], + inputs=["", "3", "", "", "", "", "", "y"], transcribe=lambda path, model_name="base": "a real transcript", whisper="faster_whisper") self.assertEqual(code, 0) |
