Fix idle-client busy-wait before first audio frame

This commit is contained in:
David Maier
2026-06-26 14:57:03 +02:00
parent 471c3fd6b4
commit 056774ea50
2 changed files with 99 additions and 3 deletions
+91 -1
View File
@@ -3,7 +3,7 @@ import queue
import threading
import time
import unittest
from unittest.mock import MagicMock, patch
from unittest.mock import MagicMock
import numpy as np
@@ -24,6 +24,21 @@ class ConcreteServeClient(ServeClientBase):
pass
class WaitTrackingEvent:
"""Threading event that records when wait() is entered."""
def __init__(self):
self._event = threading.Event()
self.wait_started = threading.Event()
def wait(self, timeout=None):
self.wait_started.set()
return self._event.wait(timeout)
def set(self):
self._event.set()
class TestServeClientBaseInit(unittest.TestCase):
def test_default_values(self):
ws = MagicMock()
@@ -266,6 +281,81 @@ class TestCleanup(unittest.TestCase):
self.assertTrue(client.exit)
class TestSpeechToTextWaitingBehavior(unittest.TestCase):
"""Tests the first-frame wait behavior in speech_to_text()."""
def setUp(self):
self.ws = MagicMock()
self.client = ConcreteServeClient(client_uid="test", websocket=self.ws)
self.client.frames_ready = WaitTrackingEvent()
self.transcribe_called = threading.Event()
self.thread_started = threading.Event()
self.cpu_used = None
self.thread = None
def tearDown(self):
if self.thread is not None and self.thread.is_alive():
self.client.exit = True
# release wait() directly so a broken cleanup() cannot hang the test process
self.client.frames_ready.set()
self.thread.join(timeout=1.0)
def _start_speech_thread(self, target=None):
self.thread = threading.Thread(target=target or self.client.speech_to_text)
self.thread.start()
return self.thread
def _join_speech_thread(self):
self.thread.join(timeout=1.0)
return not self.thread.is_alive()
def _transcribe_once(self, input_sample):
# mark the first processing step after wait and stop the loop
self.transcribe_called.set()
self.client.exit = True
return []
def _measure_waiting_cpu(self):
# measure CPU consumed by speech_to_text loop while it waits for the first frame
self.thread_started.set()
start_cpu = time.thread_time()
self.client.speech_to_text()
self.cpu_used = time.thread_time() - start_cpu
def test_waits_for_first_frame_before_transcribing(self):
self.client.transcribe_audio = MagicMock(side_effect=self._transcribe_once)
self._start_speech_thread()
self.assertTrue(self.client.frames_ready.wait_started.wait(timeout=1.0))
self.assertFalse(self.transcribe_called.is_set())
self.client.add_frames(np.zeros(self.client.RATE, dtype=np.float32))
self.assertTrue(self.transcribe_called.wait(timeout=1.0))
self.assertTrue(self._join_speech_thread())
def test_cleanup_unblocks_waiting_thread_without_audio(self):
self.client.transcribe_audio = MagicMock()
self._start_speech_thread()
self.assertTrue(self.client.frames_ready.wait_started.wait(timeout=1.0))
self.client.cleanup()
self.assertTrue(self._join_speech_thread())
self.assertTrue(self.client.exit)
self.client.transcribe_audio.assert_not_called()
def test_waiting_for_first_frame_uses_negligible_thread_cpu(self):
self._start_speech_thread(target=self._measure_waiting_cpu)
self.assertTrue(self.thread_started.wait(timeout=1.0))
time.sleep(0.25)
self.client.cleanup()
self.assertTrue(self._join_speech_thread())
self.assertIsNotNone(self.cpu_used)
self.assertLess(self.cpu_used, 0.05)
class TestTrimTranscript(unittest.TestCase):
def setUp(self):
self.ws = MagicMock()
+8 -2
View File
@@ -81,13 +81,15 @@ class ServeClientBase(object):
# threading
self.lock = threading.Lock()
self.frames_ready = threading.Event()
def speech_to_text(self):
"""
Process an audio stream in an infinite loop, continuously transcribing the speech.
This method continuously receives audio frames, performs real-time transcription, and sends
transcribed segments to the client via a WebSocket connection.
transcribed segments to the client via a WebSocket connection. The loop blocks until the first
audio frame arrives when a client is connected but still idle.
If the client's language is not detected, it waits for 30 seconds of audio input to make a language prediction.
It utilizes the Whisper ASR model to transcribe the audio, continuously processing and streaming results. Segments
@@ -103,6 +105,7 @@ class ServeClientBase(object):
break
if self.frames_np is None:
self.frames_ready.wait()
continue
if self.clip_audio:
@@ -170,7 +173,8 @@ class ServeClientBase(object):
This method is responsible for maintaining the audio stream buffer, allowing the continuous addition
of audio frames as they are received. It also ensures that the buffer does not exceed a specified size
to prevent excessive memory usage.
to prevent excessive memory usage. When the first frame arrives, it also wakes the transcription
thread so processing can begin.
If the buffer size exceeds a threshold (45 seconds of audio data), it discards the oldest 30 seconds
of audio data to maintain a reasonable buffer size. If the buffer is empty, it initializes it with the provided
@@ -194,6 +198,7 @@ class ServeClientBase(object):
else:
self.frames_np = np.concatenate((self.frames_np, frame_np), axis=0)
self.lock.release()
self.frames_ready.set()
def clip_audio_if_no_valid_segment(self):
"""
@@ -323,6 +328,7 @@ class ServeClientBase(object):
"""
logging.info("Cleaning up.")
self.exit = True
self.frames_ready.set()
def get_segment_no_speech_prob(self, segment):
return getattr(segment, "no_speech_prob", 0)