Add configurable timeout for first-frame wait and improve thread-safety

This commit is contained in:
David Maier
2026-07-03 14:45:09 +02:00
parent 056774ea50
commit 5b577b34e4
2 changed files with 55 additions and 18 deletions
+38 -3
View File
@@ -3,7 +3,7 @@ import queue
import threading import threading
import time import time
import unittest import unittest
from unittest.mock import MagicMock from unittest.mock import MagicMock, patch
import numpy as np import numpy as np
@@ -38,6 +38,9 @@ class WaitTrackingEvent:
def set(self): def set(self):
self._event.set() self._event.set()
def __getattr__(self, name):
return getattr(self._event, name)
class TestServeClientBaseInit(unittest.TestCase): class TestServeClientBaseInit(unittest.TestCase):
def test_default_values(self): def test_default_values(self):
@@ -106,7 +109,6 @@ class TestAddFrames(unittest.TestCase):
# timestamp_offset should be bumped to at least frames_offset # timestamp_offset should be bumped to at least frames_offset
self.assertGreaterEqual(self.client.timestamp_offset, self.client.frames_offset) self.assertGreaterEqual(self.client.timestamp_offset, self.client.frames_offset)
class TestAddFramesThreadSafety(unittest.TestCase): class TestAddFramesThreadSafety(unittest.TestCase):
def test_concurrent_add_frames(self): def test_concurrent_add_frames(self):
ws = MagicMock() ws = MagicMock()
@@ -129,6 +131,18 @@ class TestAddFramesThreadSafety(unittest.TestCase):
self.assertEqual(errors, []) self.assertEqual(errors, [])
self.assertIsNotNone(client.frames_np) self.assertIsNotNone(client.frames_np)
def test_exception_releases_lock_without_signaling_frames_ready(self):
ws = MagicMock()
client = ConcreteServeClient(client_uid="test", websocket=ws)
client.frames_np = np.array([0.1], dtype=np.float32)
with patch("whisper_live.backend.base.np.concatenate", side_effect=RuntimeError("boom")):
with self.assertRaisesRegex(RuntimeError, "boom"):
client.add_frames(np.array([0.2], dtype=np.float32))
self.assertFalse(client.lock.locked())
self.assertFalse(client.frames_ready.is_set())
class TestGetAudioChunkForProcessing(unittest.TestCase): class TestGetAudioChunkForProcessing(unittest.TestCase):
def setUp(self): def setUp(self):
@@ -281,6 +295,16 @@ class TestCleanup(unittest.TestCase):
self.assertTrue(client.exit) self.assertTrue(client.exit)
def _supports_thread_time():
thread_time = getattr(time, "thread_time", None)
if thread_time is None:
return False
try:
thread_time()
except NotImplementedError:
return False
return True
class TestSpeechToTextWaitingBehavior(unittest.TestCase): class TestSpeechToTextWaitingBehavior(unittest.TestCase):
"""Tests the first-frame wait behavior in speech_to_text().""" """Tests the first-frame wait behavior in speech_to_text()."""
@@ -344,6 +368,17 @@ class TestSpeechToTextWaitingBehavior(unittest.TestCase):
self.assertTrue(self.client.exit) self.assertTrue(self.client.exit)
self.client.transcribe_audio.assert_not_called() self.client.transcribe_audio.assert_not_called()
def test_exit_flag_unblocks_waiting_thread_without_signal(self):
self.client.transcribe_audio = MagicMock()
self._start_speech_thread()
self.assertTrue(self.client.frames_ready.wait_started.wait(timeout=1.0))
self.client.exit = True
self.assertTrue(self._join_speech_thread())
self.client.transcribe_audio.assert_not_called()
@unittest.skipUnless(_supports_thread_time(), "time.thread_time() not supported")
def test_waiting_for_first_frame_uses_negligible_thread_cpu(self): def test_waiting_for_first_frame_uses_negligible_thread_cpu(self):
self._start_speech_thread(target=self._measure_waiting_cpu) self._start_speech_thread(target=self._measure_waiting_cpu)
self.assertTrue(self.thread_started.wait(timeout=1.0)) self.assertTrue(self.thread_started.wait(timeout=1.0))
@@ -353,7 +388,7 @@ class TestSpeechToTextWaitingBehavior(unittest.TestCase):
self.assertTrue(self._join_speech_thread()) self.assertTrue(self._join_speech_thread())
self.assertIsNotNone(self.cpu_used) self.assertIsNotNone(self.cpu_used)
self.assertLess(self.cpu_used, 0.05) self.assertLess(self.cpu_used, 0.1)
class TestTrimTranscript(unittest.TestCase): class TestTrimTranscript(unittest.TestCase):
+17 -15
View File
@@ -21,6 +21,8 @@ class ServeClientBase(object):
"""Duration threshold in seconds for clipping audio with no valid segments.""" """Duration threshold in seconds for clipping audio with no valid segments."""
CLIP_TAIL_DURATION_S = 5 CLIP_TAIL_DURATION_S = 5
"""Duration in seconds of audio to keep after clipping.""" """Duration in seconds of audio to keep after clipping."""
FIRST_FRAME_WAIT_TIMEOUT_S = 0.1
"""Interval in seconds for re-checking exit while waiting for the first audio frame."""
client_uid: str client_uid: str
"""A unique identifier for the client.""" """A unique identifier for the client."""
@@ -105,7 +107,8 @@ class ServeClientBase(object):
break break
if self.frames_np is None: if self.frames_np is None:
self.frames_ready.wait() while self.frames_np is None and not self.exit:
self.frames_ready.wait(timeout=self.FIRST_FRAME_WAIT_TIMEOUT_S)
continue continue
if self.clip_audio: if self.clip_audio:
@@ -184,20 +187,19 @@ class ServeClientBase(object):
frame_np (numpy.ndarray): The audio frame data as a NumPy array. frame_np (numpy.ndarray): The audio frame data as a NumPy array.
""" """
self.lock.acquire() with self.lock:
if self.frames_np is not None and self.frames_np.shape[0] > self.MAX_BUFFER_DURATION_S*self.RATE: if self.frames_np is not None and self.frames_np.shape[0] > self.MAX_BUFFER_DURATION_S*self.RATE:
self.frames_offset += float(self.BUFFER_TRIM_DURATION_S) self.frames_offset += float(self.BUFFER_TRIM_DURATION_S)
self.frames_np = self.frames_np[int(self.BUFFER_TRIM_DURATION_S*self.RATE):] self.frames_np = self.frames_np[int(self.BUFFER_TRIM_DURATION_S*self.RATE):]
# check timestamp offset(should be >= self.frame_offset) # check timestamp offset(should be >= self.frame_offset)
# this basically means that there is no speech as timestamp offset hasnt updated # this basically means that there is no speech as timestamp offset hasnt updated
# and is less than frame_offset # and is less than frame_offset
if self.timestamp_offset < self.frames_offset: if self.timestamp_offset < self.frames_offset:
self.timestamp_offset = self.frames_offset self.timestamp_offset = self.frames_offset
if self.frames_np is None: if self.frames_np is None:
self.frames_np = frame_np.copy() self.frames_np = frame_np.copy()
else: else:
self.frames_np = np.concatenate((self.frames_np, frame_np), axis=0) self.frames_np = np.concatenate((self.frames_np, frame_np), axis=0)
self.lock.release()
self.frames_ready.set() self.frames_ready.set()
def clip_audio_if_no_valid_segment(self): def clip_audio_if_no_valid_segment(self):