aboutsummaryrefslogtreecommitdiff
path: root/app/tests/test_converter.py
diff options
context:
space:
mode:
Diffstat (limited to 'app/tests/test_converter.py')
-rw-r--r--app/tests/test_converter.py69
1 files changed, 69 insertions, 0 deletions
diff --git a/app/tests/test_converter.py b/app/tests/test_converter.py
index 0b0eb0d..85ae6ca 100644
--- a/app/tests/test_converter.py
+++ b/app/tests/test_converter.py
@@ -21,7 +21,11 @@ from converter.clients import (
from converter import converter as converter_mod
from converter.converter import (
AudiobookConverter,
+ ChunkClampCancelled,
+ chunk_clamp_message,
+ chunk_clamp_needed,
find_existing_outputs,
+ prompt_chunk_clamp,
prompt_overwrite,
setup_logging,
)
@@ -554,6 +558,71 @@ class PromptOverwriteTests(unittest.TestCase):
self.assertIn("overwrite them", prompt_text)
+class ChunkClampPromptTests(unittest.TestCase):
+ """The chunk-cap popup for models that cannot narrate a full
+ CHUNK_SIZE sub-request (Higgs)."""
+
+ def _entry(self):
+ from backends.sglomni.catalog import entry_by_key
+ return entry_by_key("higgs_audio_v3_tts")
+
+ def test_uncapped_models_need_no_clamp(self):
+ from backends.sglomni.catalog import entry_by_key
+ self.assertFalse(chunk_clamp_needed(None))
+ self.assertFalse(chunk_clamp_needed(entry_by_key("zonos2")))
+
+ def test_higgs_needs_a_clamp_at_the_default_chunk_size(self):
+ self.assertTrue(chunk_clamp_needed(self._entry()))
+
+ def test_message_names_the_model_and_the_cap(self):
+ lines = chunk_clamp_message(self._entry())
+ text = " ".join(lines)
+ self.assertIn("Higgs Audio v3 TTS", text)
+ self.assertIn("80 words", text)
+ self.assertIn("cut off mid-sentence", text)
+
+ def test_clamped_models_at_or_below_the_cap_need_no_clamp(self):
+ with patch.object(config, "CHUNK_SIZE", 80):
+ self.assertFalse(chunk_clamp_needed(self._entry()))
+
+ def test_prompt_answers(self):
+ entry = self._entry()
+ with patch("builtins.input", return_value=""):
+ self.assertEqual(prompt_chunk_clamp(entry), 80)
+ with patch("builtins.input", return_value="s"):
+ self.assertEqual(prompt_chunk_clamp(entry), 80)
+ with patch("builtins.input", return_value="t"):
+ self.assertIsNone(prompt_chunk_clamp(entry))
+ with patch("builtins.input", return_value="cancel"):
+ with self.assertRaises(ChunkClampCancelled):
+ prompt_chunk_clamp(entry)
+
+ def test_prompt_invalid_answer_reasked(self):
+ with patch("builtins.input", side_effect=["maybe", "t"]) as mock_input:
+ self.assertIsNone(prompt_chunk_clamp(self._entry()))
+ self.assertEqual(mock_input.call_count, 2)
+
+ def test_prompt_eof_clamps_for_unattended_runs(self):
+ with patch("builtins.input", side_effect=EOFError):
+ self.assertEqual(prompt_chunk_clamp(self._entry()), 80)
+
+ def test_ask_callback_replaces_the_console(self):
+ calls = []
+
+ def ask(lines, words):
+ calls.append((lines, words))
+ return "clamp"
+
+ self.assertEqual(prompt_chunk_clamp(self._entry(), ask=ask), 80)
+ self.assertEqual(len(calls), 1)
+ self.assertEqual(calls[0][1], 80)
+
+ def test_ask_anyway_returns_no_clamp(self):
+ self.assertIsNone(prompt_chunk_clamp(
+ self._entry(), ask=lambda lines, words: "anyway"))
+
+
+
class PreflightOverwritesTests(unittest.TestCase):
"""The pre-flight overwrite check runs without a TTS server connection."""