fix(mlx): drain the pool in a finally so a failing in-flight op releases it too; fix the mid-generation test's call count

Review follow-ups: the post-op drain now runs in a finally block for the
TTS, STT and LLM folded callables, so an exception escaping the generation
(e.g. the voice-clone fallback failing) still releases the buffers the
worker held. The mid-generation test asserted one drain call but the
path legitimately produces two (unload_model's own and the post-op one);
a second test covers the error path.
This commit is contained in:
jamiepine
2026-10-04 00:25:53 +00:00
committed by capy-ai-staging[bot]
parent cffe24ffd2
commit fb67ad436c
3 changed files with 51 additions and 22 deletions
+27 -1
View File
@@ -107,4 +107,30 @@ async def test_mlx_tts_generate_drains_cache_when_unloaded_mid_generation(monkey
assert len(audio) == 2
assert backend.model is None
mock_clear.assert_called_once()
# Once from unload_model() itself (reference drop) and once more after the
# in-flight generation released the model it had bound locally.
assert mock_clear.call_count == 2
@pytest.mark.asyncio
async def test_mlx_tts_generate_drains_cache_when_unloaded_and_generation_fails(monkeypatch):
"""The post-generation drain also runs when the in-flight op raises."""
mock_clear = MagicMock()
monkeypatch.setattr(mlx_backend, "empty_mlx_cache", mock_clear)
backend = mlx_backend.MLXTTSBackend(model_size="0.6B")
class _FakeModel:
def generate(self, text, **kwargs):
backend.unload_model()
raise RuntimeError("boom")
yield # pragma: no cover - makes this a generator
backend.model = _FakeModel()
backend._current_model_size = "0.6B"
with pytest.raises(RuntimeError, match="boom"):
await backend.generate("hello", {}, "en")
assert backend.model is None
assert mock_clear.call_count == 2