aboutsummaryrefslogtreecommitdiff
path: root/app/tests
diff options
context:
space:
mode:
Diffstat (limited to 'app/tests')
-rw-r--r--app/tests/test_tts.py98
1 files changed, 98 insertions, 0 deletions
diff --git a/app/tests/test_tts.py b/app/tests/test_tts.py
index 87f98df..f99e257 100644
--- a/app/tests/test_tts.py
+++ b/app/tests/test_tts.py
@@ -1284,6 +1284,104 @@ class AudioCppTTSClientTruncationTests(unittest.TestCase):
self.assertIsNotNone(result)
+class AudioCppUnloadModelsTests(unittest.TestCase):
+ """Before generating, the client asks the server to drop loaded models."""
+
+ @staticmethod
+ def _client():
+ client = AudioCppTTSClient.__new__(AudioCppTTSClient)
+ client.api_url = "http://127.0.0.1:8080"
+ return client
+
+ @staticmethod
+ def _response(body):
+ response = MagicMock()
+ response.__enter__.return_value = response
+ response.read.return_value = body
+ return response
+
+ def test_posts_to_unload_all_models(self):
+ client = self._client()
+ with patch("converter.tts.urllib.request.urlopen",
+ return_value=self._response(b'{"unloaded": ["qwen"]}')) as mock_urlopen:
+ client._unload_server_models()
+ request = mock_urlopen.call_args[0][0]
+ self.assertEqual(request.full_url,
+ "http://127.0.0.1:8080/v1/tasks/unload_all_models")
+ self.assertEqual(request.method, "POST")
+ self.assertEqual(request.data, b"")
+
+ def test_reports_unloaded_ids(self):
+ client = self._client()
+ buf = io.StringIO()
+ with patch("converter.tts.urllib.request.urlopen",
+ return_value=self._response(b'{"unloaded": ["a", "b"]}')), \
+ redirect_stdout(buf):
+ client._unload_server_models()
+ self.assertIn("Unloaded 2 model(s)", buf.getvalue())
+ self.assertIn("a, b", buf.getvalue())
+
+ def test_no_loaded_models_is_silent(self):
+ client = self._client()
+ buf = io.StringIO()
+ with patch("converter.tts.urllib.request.urlopen",
+ return_value=self._response(b'{"unloaded": []}')), \
+ redirect_stdout(buf):
+ client._unload_server_models()
+ self.assertEqual(buf.getvalue(), "")
+
+ def test_http_error_warns_and_continues(self):
+ client = self._client()
+ buf = io.StringIO()
+ with patch("converter.tts.urllib.request.urlopen",
+ side_effect=tts.urllib.error.HTTPError(
+ "http://127.0.0.1:8080/v1/tasks/unload_all_models",
+ 404, "Not Found", None, io.BytesIO())), \
+ redirect_stdout(buf):
+ client._unload_server_models()
+ out = buf.getvalue()
+ self.assertIn("[WARNING]", out)
+ self.assertIn("404", out)
+
+ def test_connection_error_warns_and_continues(self):
+ client = self._client()
+ buf = io.StringIO()
+ with patch("converter.tts.urllib.request.urlopen",
+ side_effect=tts.urllib.error.URLError("refused")), \
+ redirect_stdout(buf):
+ client._unload_server_models()
+ self.assertIn("[WARNING]", buf.getvalue())
+
+ def test_connect_unloads_before_returning(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
+ client.family = "qwen3_tts"
+ client.task = tts.AUDIOCPP_TASK_TTS
+ client.profile = tts.AUDIOCPP_FAMILY_PROFILES["qwen3_tts"]
+ client.design_mode = False
+ client.instruction_voice = False
+ client.instructions = ""
+ with patch.object(client, "_check_health"), \
+ patch.object(client, "_list_models",
+ return_value=[{"id": client.model_id,
+ "family": "qwen3_tts",
+ "task": "tts"}]), \
+ patch.object(client, "_auto_pick_model_id"), \
+ patch.object(client, "_select_model"), \
+ patch.object(client, "_require_model_id"), \
+ patch.object(client, "_resolve_family"), \
+ patch.object(client, "_resolve_task"), \
+ patch.object(client, "_check_voice"), \
+ patch.object(client, "_unload_server_models") as mock_unload:
+ client._connect()
+ mock_unload.assert_called_once()
+
+
class BackendWiringTests(unittest.TestCase):
"""AudiobookConverter wiring for the --backend selector."""