diff options
Diffstat (limited to 'app/tests')
| -rw-r--r-- | app/tests/test_tts.py | 98 |
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.""" |
