From 9ee1f9c1edf9db7ae6b9b3d895d7d7e6a0f273fc Mon Sep 17 00:00:00 2001 From: Palash Debnath <4178343+debpalash@users.noreply.github.com> Date: Thu, 17 Sep 2026 15:03:03 +0530 Subject: [PATCH] fix(dubbing): make disconnect cleanup single-shot --- backend/api/routers/dub_core.py | 5 +++++ tests/test_dub_transcribe.py | 32 ++++++++++++++++++++++++++++++-- 2 files changed, 35 insertions(+), 2 deletions(-) diff --git a/backend/api/routers/dub_core.py b/backend/api/routers/dub_core.py index c7baab3a..11049b83 100644 --- a/backend/api/routers/dub_core.py +++ b/backend/api/routers/dub_core.py @@ -888,6 +888,7 @@ class _ASRWorkLifetime: import threading self._lock = threading.Lock() self._closed = threading.Event() + self._cleaned = False def run(self, fn): with self._lock: @@ -902,6 +903,10 @@ class _ASRWorkLifetime: def cleanup(self, fn): self.stop() with self._lock: + if self._cleaned: + return + # Claim cleanup before invoking even a non-idempotent unload. + self._cleaned = True return fn() diff --git a/tests/test_dub_transcribe.py b/tests/test_dub_transcribe.py index e6a1175a..02092ad9 100644 --- a/tests/test_dub_transcribe.py +++ b/tests/test_dub_transcribe.py @@ -974,6 +974,7 @@ def test_stream_cleanup_waits_for_native_work_and_rejects_late_work(): from api.routers.dub_core import _ASRWorkLifetime lifetime = _ASRWorkLifetime() started, release, cleaned = threading.Event(), threading.Event(), threading.Event() + cleanup_started = threading.Event() events = [] def native(): started.set() @@ -986,9 +987,13 @@ def test_stream_cleanup_waits_for_native_work_and_rejects_late_work(): work = pool.submit(lifetime.run, native) assert started.wait(5) lifetime.stop() - removal = pool.submit(lifetime.cleanup, cleanup) + def remove(): + cleanup_started.set() + lifetime.cleanup(cleanup) + removal = pool.submit(remove) try: - assert not cleaned.wait(0.05) + assert cleanup_started.wait(5) + assert not cleaned.is_set() finally: release.set() work.result(timeout=5) @@ -996,3 +1001,26 @@ def test_stream_cleanup_waits_for_native_work_and_rejects_late_work(): assert events == ["native finished", "unloaded"] with pytest.raises(RuntimeError, match="stream has ended"): lifetime.run(lambda: pytest.fail("late work accessed unloaded model")) + + +def test_stream_unload_is_single_shot_during_disconnect(): + import threading + from concurrent.futures import ThreadPoolExecutor + from api.routers.dub_core import _ASRWorkLifetime + lifetime = _ASRWorkLifetime() + started, release = threading.Event(), threading.Event() + calls = [] + def unload(): + calls.append("unload") + assert len(calls) == 1 + started.set() + assert release.wait(5) + with ThreadPoolExecutor(max_workers=2) as pool: + normal = pool.submit(lifetime.cleanup, unload) + assert started.wait(5) + lifetime.stop() + disconnected = pool.submit(lifetime.cleanup, unload) + release.set() + normal.result(timeout=5) + disconnected.result(timeout=5) + assert calls == ["unload"]