From b1cd51ac8a504aabdf35b9062f2ba8e168907139 Mon Sep 17 00:00:00 2001 From: Aaron Boxer Date: Fri, 17 Apr 2026 09:05:34 -0400 Subject: [PATCH] CI: expand test suite coverage these new test cover issues such as thread safety, VAD thresholding, message routing, error handling etc. that weren't covered by existing tests. Mocking is used to avoid dependencies on GPU, ONNX etc. --- pytest.ini | 5 + tests/test_base_backend.py | 387 ++++++++++++++++++++++++++++++++++ tests/test_client_extended.py | 257 ++++++++++++++++++++++ tests/test_server_extended.py | 217 +++++++++++++++++++ tests/test_utils.py | 134 ++++++++++++ tests/test_vad_extended.py | 131 ++++++++++++ 6 files changed, 1131 insertions(+) create mode 100644 pytest.ini create mode 100644 tests/test_base_backend.py create mode 100644 tests/test_client_extended.py create mode 100644 tests/test_server_extended.py create mode 100644 tests/test_utils.py create mode 100644 tests/test_vad_extended.py diff --git a/pytest.ini b/pytest.ini new file mode 100644 index 0000000..4ecb1ad --- /dev/null +++ b/pytest.ini @@ -0,0 +1,5 @@ +[pytest] +testpaths = tests +python_files = test_*.py +python_classes = Test* +python_functions = test_* diff --git a/tests/test_base_backend.py b/tests/test_base_backend.py new file mode 100644 index 0000000..8f90b77 --- /dev/null +++ b/tests/test_base_backend.py @@ -0,0 +1,387 @@ +import json +import queue +import threading +import time +import unittest +from unittest.mock import MagicMock, patch + +import numpy as np + +from whisper_live.backend.base import ServeClientBase + + +class ConcreteServeClient(ServeClientBase): + """Concrete subclass for testing the abstract base class.""" + + def __init__(self, **kwargs): + super().__init__(**kwargs) + self.language = "en" + + def transcribe_audio(self, input_sample): + return None + + def handle_transcription_output(self, result, duration): + pass + + +class TestServeClientBaseInit(unittest.TestCase): + def test_default_values(self): + ws = MagicMock() + client = ConcreteServeClient(client_uid="test-uid", websocket=ws) + self.assertEqual(client.client_uid, "test-uid") + self.assertEqual(client.send_last_n_segments, 10) + self.assertAlmostEqual(client.no_speech_thresh, 0.45) + self.assertFalse(client.clip_audio) + self.assertEqual(client.same_output_threshold, 10) + self.assertIsNone(client.frames_np) + self.assertAlmostEqual(client.timestamp_offset, 0.0) + self.assertFalse(client.exit) + self.assertEqual(client.transcript, []) + + def test_custom_values(self): + ws = MagicMock() + q = queue.Queue() + client = ConcreteServeClient( + client_uid="uid2", + websocket=ws, + send_last_n_segments=5, + no_speech_thresh=0.6, + clip_audio=True, + same_output_threshold=20, + translation_queue=q, + ) + self.assertEqual(client.send_last_n_segments, 5) + self.assertAlmostEqual(client.no_speech_thresh, 0.6) + self.assertTrue(client.clip_audio) + self.assertEqual(client.same_output_threshold, 20) + self.assertIs(client.translation_queue, q) + + +class TestAddFrames(unittest.TestCase): + def setUp(self): + self.ws = MagicMock() + self.client = ConcreteServeClient(client_uid="test", websocket=self.ws) + + def test_first_frame_initializes_buffer(self): + frame = np.array([0.1, 0.2, 0.3], dtype=np.float32) + self.client.add_frames(frame) + np.testing.assert_array_equal(self.client.frames_np, frame) + + def test_subsequent_frames_concatenated(self): + frame1 = np.array([0.1, 0.2], dtype=np.float32) + frame2 = np.array([0.3, 0.4], dtype=np.float32) + self.client.add_frames(frame1) + self.client.add_frames(frame2) + expected = np.array([0.1, 0.2, 0.3, 0.4], dtype=np.float32) + np.testing.assert_array_equal(self.client.frames_np, expected) + + def test_buffer_trimmed_at_45_seconds(self): + # 45 seconds + 1 sample at 16kHz = 720001 samples + self.client.frames_np = np.zeros(45 * 16000 + 1, dtype=np.float32) + self.client.add_frames(np.array([1.0], dtype=np.float32)) + # after trimming 30s, buffer should be ~15s + 1 original + 1 new + expected_len = (45 * 16000 + 1) - (30 * 16000) + 1 + self.assertEqual(self.client.frames_np.shape[0], expected_len) + self.assertAlmostEqual(self.client.frames_offset, 30.0) + + def test_timestamp_offset_updated_on_trim(self): + self.client.frames_np = np.zeros(45 * 16000 + 1, dtype=np.float32) + self.client.timestamp_offset = 5.0 # behind frames_offset after trim + self.client.add_frames(np.array([1.0], dtype=np.float32)) + # timestamp_offset should be bumped to at least frames_offset + self.assertGreaterEqual(self.client.timestamp_offset, self.client.frames_offset) + + +class TestAddFramesThreadSafety(unittest.TestCase): + def test_concurrent_add_frames(self): + ws = MagicMock() + client = ConcreteServeClient(client_uid="test", websocket=ws) + errors = [] + + def add_many(): + try: + for _ in range(100): + client.add_frames(np.random.randn(160).astype(np.float32)) + except Exception as e: + errors.append(e) + + threads = [threading.Thread(target=add_many) for _ in range(4)] + for t in threads: + t.start() + for t in threads: + t.join() + + self.assertEqual(errors, []) + self.assertIsNotNone(client.frames_np) + + +class TestGetAudioChunkForProcessing(unittest.TestCase): + def setUp(self): + self.ws = MagicMock() + self.client = ConcreteServeClient(client_uid="test", websocket=self.ws) + + def test_empty_buffer_returns_empty(self): + self.client.frames_np = np.array([], dtype=np.float32) + chunk, duration = self.client.get_audio_chunk_for_processing() + self.assertEqual(duration, 0.0) + self.assertEqual(chunk.shape[0], 0) + + def test_full_buffer_no_offset(self): + audio = np.random.randn(16000).astype(np.float32) # 1 second + self.client.frames_np = audio + chunk, duration = self.client.get_audio_chunk_for_processing() + self.assertAlmostEqual(duration, 1.0) + np.testing.assert_array_equal(chunk, audio) + + def test_with_offset(self): + audio = np.random.randn(32000).astype(np.float32) # 2 seconds + self.client.frames_np = audio + self.client.timestamp_offset = 1.0 # skip first second + chunk, duration = self.client.get_audio_chunk_for_processing() + self.assertAlmostEqual(duration, 1.0) + self.assertEqual(chunk.shape[0], 16000) + + +class TestClipAudioIfNoValidSegment(unittest.TestCase): + def setUp(self): + self.ws = MagicMock() + self.client = ConcreteServeClient( + client_uid="test", websocket=self.ws, clip_audio=True + ) + + def test_clips_when_chunk_exceeds_25s(self): + # 30 seconds of audio with no valid segments + self.client.frames_np = np.zeros(30 * 16000, dtype=np.float32) + self.client.timestamp_offset = 0.0 + self.client.frames_offset = 0.0 + self.client.clip_audio_if_no_valid_segment() + # offset should have advanced to leave ~5s of remaining audio + expected_offset = (30 * 16000 / 16000) - 5 + self.assertAlmostEqual(self.client.timestamp_offset, expected_offset, places=1) + + def test_no_clip_when_short(self): + self.client.frames_np = np.zeros(10 * 16000, dtype=np.float32) + self.client.timestamp_offset = 0.0 + self.client.frames_offset = 0.0 + self.client.clip_audio_if_no_valid_segment() + self.assertAlmostEqual(self.client.timestamp_offset, 0.0) + + +class TestPrepareSegments(unittest.TestCase): + def setUp(self): + self.ws = MagicMock() + self.client = ConcreteServeClient( + client_uid="test", websocket=self.ws, send_last_n_segments=3 + ) + + def test_empty_transcript_no_last(self): + segments = self.client.prepare_segments() + self.assertEqual(segments, []) + + def test_empty_transcript_with_last(self): + last = {"start": "0.000", "end": "1.000", "text": "hello", "completed": False} + segments = self.client.prepare_segments(last_segment=last) + self.assertEqual(len(segments), 1) + self.assertEqual(segments[0]["text"], "hello") + + def test_fewer_than_n_segments(self): + self.client.transcript = [ + {"start": "0.000", "end": "1.000", "text": "a", "completed": True}, + {"start": "1.000", "end": "2.000", "text": "b", "completed": True}, + ] + segments = self.client.prepare_segments() + self.assertEqual(len(segments), 2) + + def test_more_than_n_segments_truncated(self): + self.client.transcript = [ + {"start": f"{i}.000", "end": f"{i+1}.000", "text": f"seg{i}", "completed": True} + for i in range(10) + ] + segments = self.client.prepare_segments() + self.assertEqual(len(segments), 3) + self.assertEqual(segments[0]["text"], "seg7") + + def test_last_segment_appended(self): + self.client.transcript = [ + {"start": "0.000", "end": "1.000", "text": "a", "completed": True}, + ] + last = {"start": "1.000", "end": "2.000", "text": "in progress", "completed": False} + segments = self.client.prepare_segments(last_segment=last) + self.assertEqual(len(segments), 2) + self.assertEqual(segments[-1]["text"], "in progress") + + +class TestFormatSegment(unittest.TestCase): + def setUp(self): + self.ws = MagicMock() + self.client = ConcreteServeClient(client_uid="test", websocket=self.ws) + + def test_format(self): + seg = self.client.format_segment(1.234, 5.678, "hello world", completed=True) + self.assertEqual(seg["start"], "1.234") + self.assertEqual(seg["end"], "5.678") + self.assertEqual(seg["text"], "hello world") + self.assertTrue(seg["completed"]) + + def test_format_not_completed(self): + seg = self.client.format_segment(0.0, 1.0, "text") + self.assertFalse(seg["completed"]) + + +class TestSendTranscriptionToClient(unittest.TestCase): + def setUp(self): + self.ws = MagicMock() + self.client = ConcreteServeClient(client_uid="test-uid", websocket=self.ws) + + def test_sends_json(self): + segments = [{"start": "0.000", "end": "1.000", "text": "hi", "completed": True}] + self.client.send_transcription_to_client(segments) + self.ws.send.assert_called_once() + sent = json.loads(self.ws.send.call_args[0][0]) + self.assertEqual(sent["uid"], "test-uid") + self.assertEqual(len(sent["segments"]), 1) + + def test_send_failure_logged_not_raised(self): + self.ws.send.side_effect = ConnectionError("broken pipe") + # should not raise + self.client.send_transcription_to_client([]) + + +class TestDisconnect(unittest.TestCase): + def test_sends_disconnect_message(self): + ws = MagicMock() + client = ConcreteServeClient(client_uid="uid1", websocket=ws) + client.disconnect() + sent = json.loads(ws.send.call_args[0][0]) + self.assertEqual(sent["uid"], "uid1") + self.assertEqual(sent["message"], "DISCONNECT") + + +class TestCleanup(unittest.TestCase): + def test_sets_exit_flag(self): + ws = MagicMock() + client = ConcreteServeClient(client_uid="uid1", websocket=ws) + self.assertFalse(client.exit) + client.cleanup() + self.assertTrue(client.exit) + + +class TestUpdateSegments(unittest.TestCase): + """Tests for the core update_segments() logic.""" + + def setUp(self): + self.ws = MagicMock() + self.client = ConcreteServeClient( + client_uid="test", + websocket=self.ws, + no_speech_thresh=0.45, + same_output_threshold=3, + ) + self.client.frames_np = np.zeros(16000 * 5, dtype=np.float32) + + def _make_segment(self, start, end, text, no_speech_prob=0.0): + seg = MagicMock() + seg.start = start + seg.end = end + seg.text = text + seg.no_speech_prob = no_speech_prob + return seg + + def test_single_segment_becomes_last(self): + segs = [self._make_segment(0.0, 1.0, " hello")] + last = self.client.update_segments(segs, duration=2.0) + self.assertIsNotNone(last) + self.assertIn("hello", last["text"]) + self.assertFalse(last["completed"]) + self.assertEqual(len(self.client.transcript), 0) + + def test_multiple_segments_completes_all_but_last(self): + segs = [ + self._make_segment(0.0, 1.0, " first"), + self._make_segment(1.0, 2.0, " second"), + ] + last = self.client.update_segments(segs, duration=3.0) + self.assertEqual(len(self.client.transcript), 1) + self.assertTrue(self.client.transcript[0]["completed"]) + self.assertIn("first", self.client.transcript[0]["text"]) + self.assertIsNotNone(last) + self.assertIn("second", last["text"]) + + def test_high_no_speech_prob_skipped(self): + segs = [ + self._make_segment(0.0, 1.0, " noise", no_speech_prob=0.9), + self._make_segment(1.0, 2.0, " also noise", no_speech_prob=0.9), + ] + last = self.client.update_segments(segs, duration=3.0) + self.assertEqual(len(self.client.transcript), 0) + self.assertIsNone(last) + + def test_segment_with_start_gte_end_skipped(self): + segs = [ + self._make_segment(1.0, 0.5, " backwards"), + self._make_segment(1.5, 2.0, " normal"), + ] + last = self.client.update_segments(segs, duration=3.0) + self.assertEqual(len(self.client.transcript), 0) + self.assertIsNotNone(last) + + def test_repeated_output_triggers_completion(self): + seg = self._make_segment(0.0, 1.0, " repeated") + for _ in range(self.client.same_output_threshold + 2): + last = self.client.update_segments([seg], duration=2.0) + # after enough repeats, should be added to transcript + self.assertTrue(len(self.client.transcript) >= 1) + + def test_translation_queue_receives_completed(self): + q = queue.Queue() + self.client.translation_queue = q + segs = [ + self._make_segment(0.0, 1.0, " first"), + self._make_segment(1.0, 2.0, " second"), + ] + self.client.update_segments(segs, duration=3.0) + self.assertFalse(q.empty()) + item = q.get_nowait() + self.assertIn("first", item["text"]) + + def test_timestamp_offset_advances(self): + segs = [ + self._make_segment(0.0, 1.0, " first"), + self._make_segment(1.0, 2.0, " second"), + ] + self.client.update_segments(segs, duration=3.0) + self.assertGreater(self.client.timestamp_offset, 0.0) + + +class TestGetSegmentHelpers(unittest.TestCase): + def setUp(self): + self.ws = MagicMock() + self.client = ConcreteServeClient(client_uid="test", websocket=self.ws) + + def test_get_segment_no_speech_prob_attr(self): + seg = MagicMock() + seg.no_speech_prob = 0.3 + self.assertAlmostEqual(self.client.get_segment_no_speech_prob(seg), 0.3) + + def test_get_segment_no_speech_prob_fallback(self): + seg = MagicMock(spec=[]) # no attributes + self.assertEqual(self.client.get_segment_no_speech_prob(seg), 0) + + def test_get_segment_start_uses_start(self): + seg = MagicMock() + seg.start = 1.5 + self.assertAlmostEqual(self.client.get_segment_start(seg), 1.5) + + def test_get_segment_end_uses_end(self): + seg = MagicMock() + seg.end = 3.0 + self.assertAlmostEqual(self.client.get_segment_end(seg), 3.0) + + def test_get_segment_start_fallback_to_start_ts(self): + seg = MagicMock(spec=["start_ts"]) + seg.start_ts = 2.0 + self.assertAlmostEqual(self.client.get_segment_start(seg), 2.0) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_client_extended.py b/tests/test_client_extended.py new file mode 100644 index 0000000..087bae7 --- /dev/null +++ b/tests/test_client_extended.py @@ -0,0 +1,257 @@ +import json +import time +import unittest +from unittest.mock import patch, MagicMock, PropertyMock + +from whisper_live.client import Client, TranscriptionTeeClient + + +class TestClientStatusMessages(unittest.TestCase): + """Tests for Client.handle_status_messages() and on_message() branches.""" + + @patch("whisper_live.client.websocket.WebSocketApp") + @patch("whisper_live.client.pyaudio.PyAudio") + def setUp(self, mock_pyaudio, mock_websocket): + mock_pyaudio.return_value.open.return_value = MagicMock() + self.client = Client(host="localhost", port=9090, lang="en") + + def tearDown(self): + self.client.close_websocket() + + def test_wait_status(self): + msg = {"uid": self.client.uid, "status": "WAIT", "message": 5.0} + self.client.handle_status_messages(msg) + self.assertTrue(self.client.waiting) + + def test_error_status(self): + msg = {"uid": self.client.uid, "status": "ERROR", "message": "model not found"} + self.client.handle_status_messages(msg) + self.assertTrue(self.client.server_error) + + def test_warning_status_no_side_effects(self): + msg = {"uid": self.client.uid, "status": "WARNING", "message": "fallback backend"} + self.client.handle_status_messages(msg) + self.assertFalse(self.client.server_error) + self.assertFalse(self.client.waiting) + + def test_on_message_wrong_uid_ignored(self): + msg = json.dumps({"uid": "wrong-uid", "segments": [{"start": 0, "end": 1, "text": "hi", "completed": True}]}) + self.client.on_message(MagicMock(), msg) + self.assertEqual(len(self.client.transcript), 0) + + def test_on_message_disconnect(self): + self.client.recording = True + msg = json.dumps({"uid": self.client.uid, "message": "DISCONNECT"}) + self.client.on_message(MagicMock(), msg) + self.assertFalse(self.client.recording) + + def test_on_message_server_ready(self): + msg = json.dumps({ + "uid": self.client.uid, + "message": "SERVER_READY", + "backend": "faster_whisper", + }) + self.client.on_message(MagicMock(), msg) + self.assertTrue(self.client.recording) + self.assertEqual(self.client.server_backend, "faster_whisper") + + def test_on_message_language_detection(self): + msg = json.dumps({ + "uid": self.client.uid, + "language": "fr", + "language_prob": 0.95, + }) + self.client.on_message(MagicMock(), msg) + self.assertEqual(self.client.language, "fr") + + +class TestClientTranslationFlow(unittest.TestCase): + """Tests for the translation-related client functionality.""" + + @patch("whisper_live.client.websocket.WebSocketApp") + @patch("whisper_live.client.pyaudio.PyAudio") + def setUp(self, mock_pyaudio, mock_websocket): + mock_pyaudio.return_value.open.return_value = MagicMock() + self.client = Client( + host="localhost", + port=9090, + lang="en", + enable_translation=True, + target_language="es", + ) + # simulate SERVER_READY so server_backend is set + ready_msg = json.dumps({ + "uid": self.client.uid, + "message": "SERVER_READY", + "backend": "faster_whisper", + }) + self.client.on_message(MagicMock(), ready_msg) + + def tearDown(self): + self.client.close_websocket() + + def test_on_open_includes_translation_fields(self): + mock_ws = MagicMock() + self.client.on_open(mock_ws) + sent = json.loads(mock_ws.send.call_args[0][0]) + self.assertTrue(sent["enable_translation"]) + self.assertEqual(sent["target_language"], "es") + + def test_translated_segments_processed(self): + msg = json.dumps({ + "uid": self.client.uid, + "translated_segments": [ + {"start": "0.000", "end": "1.000", "text": "Hola mundo", "completed": True}, + ], + }) + self.client.on_message(MagicMock(), msg) + self.assertEqual(len(self.client.translated_transcript), 1) + self.assertEqual(self.client.translated_transcript[0]["text"], "Hola mundo") + + def test_translation_callback_invoked(self): + callback = MagicMock() + self.client.translation_callback = callback + msg = json.dumps({ + "uid": self.client.uid, + "translated_segments": [ + {"start": "0.000", "end": "1.000", "text": "Hola", "completed": True}, + ], + }) + self.client.on_message(MagicMock(), msg) + callback.assert_called_once() + + def test_translation_callback_exception_handled(self): + callback = MagicMock(side_effect=RuntimeError("callback broke")) + self.client.translation_callback = callback + msg = json.dumps({ + "uid": self.client.uid, + "translated_segments": [ + {"start": "0.000", "end": "1.000", "text": "Hola", "completed": True}, + ], + }) + # should not raise + self.client.on_message(MagicMock(), msg) + + +class TestClientTranscriptionCallback(unittest.TestCase): + """Tests for the transcription callback feature.""" + + @patch("whisper_live.client.websocket.WebSocketApp") + @patch("whisper_live.client.pyaudio.PyAudio") + def setUp(self, mock_pyaudio, mock_websocket): + mock_pyaudio.return_value.open.return_value = MagicMock() + self.callback = MagicMock() + self.client = Client( + host="localhost", + port=9090, + lang="en", + transcription_callback=self.callback, + ) + ready_msg = json.dumps({ + "uid": self.client.uid, + "message": "SERVER_READY", + "backend": "faster_whisper", + }) + self.client.on_message(MagicMock(), ready_msg) + + def tearDown(self): + self.client.close_websocket() + + def test_callback_receives_text_and_segments(self): + msg = json.dumps({ + "uid": self.client.uid, + "segments": [ + {"start": "0.000", "end": "1.000", "text": "Hello", "completed": True}, + ], + }) + self.client.on_message(MagicMock(), msg) + self.callback.assert_called_once() + text_arg, segments_arg = self.callback.call_args[0] + self.assertIn("Hello", text_arg) + self.assertIsInstance(segments_arg, list) + + def test_callback_exception_does_not_crash(self): + self.callback.side_effect = ValueError("boom") + msg = json.dumps({ + "uid": self.client.uid, + "segments": [ + {"start": "0.000", "end": "1.000", "text": "Test", "completed": True}, + ], + }) + # should not raise + self.client.on_message(MagicMock(), msg) + + +class TestClientSrtWriting(unittest.TestCase): + """Tests for Client.write_srt_file() edge cases.""" + + @patch("whisper_live.client.websocket.WebSocketApp") + @patch("whisper_live.client.pyaudio.PyAudio") + def setUp(self, mock_pyaudio, mock_websocket): + mock_pyaudio.return_value.open.return_value = MagicMock() + self.client = Client(host="localhost", port=9090, lang="en") + self.client.server_backend = "faster_whisper" + + def tearDown(self): + self.client.close_websocket() + import os + for f in ["test_out.srt"]: + if os.path.exists(f): + os.remove(f) + + def test_write_srt_empty_transcript_with_last_segment(self): + self.client.transcript = [] + self.client.last_segment = {"start": "0.000", "end": "1.000", "text": "final"} + self.client.write_srt_file("test_out.srt") + self.assertEqual(len(self.client.transcript), 1) + self.assertEqual(self.client.transcript[0]["text"], "final") + + def test_write_srt_appends_last_segment_if_different(self): + self.client.transcript = [{"start": "0.000", "end": "1.000", "text": "first"}] + self.client.last_segment = {"start": "1.000", "end": "2.000", "text": "second"} + self.client.write_srt_file("test_out.srt") + self.assertEqual(len(self.client.transcript), 2) + + def test_write_srt_no_duplicate_last_segment(self): + self.client.transcript = [{"start": "0.000", "end": "1.000", "text": "same"}] + self.client.last_segment = {"start": "0.000", "end": "1.000", "text": "same"} + self.client.write_srt_file("test_out.srt") + self.assertEqual(len(self.client.transcript), 1) + + +class TestWaitBeforeDisconnect(unittest.TestCase): + """Tests for Client.wait_before_disconnect().""" + + @patch("whisper_live.client.websocket.WebSocketApp") + @patch("whisper_live.client.pyaudio.PyAudio") + def setUp(self, mock_pyaudio, mock_websocket): + mock_pyaudio.return_value.open.return_value = MagicMock() + self.client = Client(host="localhost", port=9090, lang="en") + + def tearDown(self): + self.client.close_websocket() + + def test_raises_if_no_response(self): + self.client.last_response_received = None + with self.assertRaises(AssertionError): + self.client.wait_before_disconnect() + + def test_returns_immediately_if_timeout_elapsed(self): + self.client.last_response_received = time.time() - 100 + self.client.disconnect_if_no_response_for = 15 + start = time.time() + self.client.wait_before_disconnect() + elapsed = time.time() - start + self.assertLess(elapsed, 1.0) + + +class TestTeeClientEdgeCases(unittest.TestCase): + """Edge cases for TranscriptionTeeClient.""" + + def test_empty_clients_raises(self): + with self.assertRaises(Exception): + TranscriptionTeeClient([]) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_server_extended.py b/tests/test_server_extended.py new file mode 100644 index 0000000..93daf5f --- /dev/null +++ b/tests/test_server_extended.py @@ -0,0 +1,217 @@ +import json +import time +import unittest +from unittest import mock +from unittest.mock import MagicMock, patch + +from whisper_live.server import TranscriptionServer, BackendType, ClientManager + + +class TestClientManagerAddRemove(unittest.TestCase): + def setUp(self): + self.cm = ClientManager(max_clients=2, max_connection_time=60) + + def test_add_and_get_client(self): + ws = MagicMock() + client = MagicMock() + self.cm.add_client(ws, client) + self.assertIs(self.cm.get_client(ws), client) + + def test_get_nonexistent_client(self): + ws = MagicMock() + self.assertFalse(self.cm.get_client(ws)) + + def test_remove_client_calls_cleanup(self): + ws = MagicMock() + client = MagicMock() + self.cm.add_client(ws, client) + self.cm.remove_client(ws) + client.cleanup.assert_called_once() + self.assertNotIn(ws, self.cm.clients) + self.assertNotIn(ws, self.cm.start_times) + + def test_remove_nonexistent_client_no_error(self): + ws = MagicMock() + self.cm.remove_client(ws) # should not raise + + +class TestClientManagerServerFull(unittest.TestCase): + def setUp(self): + self.cm = ClientManager(max_clients=1, max_connection_time=60) + + def test_not_full_returns_false(self): + ws = MagicMock() + options = {"uid": "test"} + self.assertFalse(self.cm.is_server_full(ws, options)) + + def test_full_sends_wait_and_returns_true(self): + ws1 = MagicMock() + self.cm.add_client(ws1, MagicMock()) + + ws2 = MagicMock() + options = {"uid": "new-client"} + self.assertTrue(self.cm.is_server_full(ws2, options)) + ws2.send.assert_called_once() + sent = json.loads(ws2.send.call_args[0][0]) + self.assertEqual(sent["status"], "WAIT") + self.assertEqual(sent["uid"], "new-client") + + +class TestClientManagerTimeout(unittest.TestCase): + def setUp(self): + self.cm = ClientManager(max_clients=4, max_connection_time=10) + + def test_not_timed_out(self): + ws = MagicMock() + client = MagicMock() + self.cm.add_client(ws, client) + self.assertFalse(self.cm.is_client_timeout(ws)) + + def test_timed_out(self): + ws = MagicMock() + client = MagicMock() + self.cm.add_client(ws, client) + self.cm.start_times[ws] = time.time() - 20 + self.assertTrue(self.cm.is_client_timeout(ws)) + client.disconnect.assert_called_once() + + +class TestClientManagerGetWaitTime(unittest.TestCase): + def test_no_clients_returns_zero(self): + cm = ClientManager(max_clients=4, max_connection_time=600) + self.assertEqual(cm.get_wait_time(), 0) + + def test_single_client_wait_time(self): + cm = ClientManager(max_clients=4, max_connection_time=600) + ws = MagicMock() + cm.add_client(ws, MagicMock()) + cm.start_times[ws] = time.time() - 300 + wait = cm.get_wait_time() + self.assertAlmostEqual(wait, 5.0, places=0) + + def test_multiple_clients_returns_minimum(self): + cm = ClientManager(max_clients=4, max_connection_time=600) + ws1, ws2 = MagicMock(), MagicMock() + cm.add_client(ws1, MagicMock()) + cm.add_client(ws2, MagicMock()) + cm.start_times[ws1] = time.time() - 100 + cm.start_times[ws2] = time.time() - 500 + wait = cm.get_wait_time() + # ws2 has 100s remaining = ~1.67 minutes + self.assertAlmostEqual(wait, 100 / 60, places=0) + + +class TestBackendType(unittest.TestCase): + def test_valid_types(self): + valid = BackendType.valid_types() + self.assertIn("faster_whisper", valid) + self.assertIn("tensorrt", valid) + self.assertIn("openvino", valid) + + def test_is_valid(self): + self.assertTrue(BackendType.is_valid("faster_whisper")) + self.assertFalse(BackendType.is_valid("nonexistent")) + + def test_type_checks(self): + self.assertTrue(BackendType.FASTER_WHISPER.is_faster_whisper()) + self.assertFalse(BackendType.FASTER_WHISPER.is_tensorrt()) + self.assertTrue(BackendType.TENSORRT.is_tensorrt()) + self.assertTrue(BackendType.OPENVINO.is_openvino()) + + def test_enum_from_string(self): + bt = BackendType("faster_whisper") + self.assertEqual(bt, BackendType.FASTER_WHISPER) + + def test_invalid_enum_raises(self): + with self.assertRaises(ValueError): + BackendType("invalid_backend") + + +class TestTranscriptionServerInit(unittest.TestCase): + def test_defaults(self): + server = TranscriptionServer() + self.assertIsNone(server.client_manager) + self.assertTrue(server.use_vad) + self.assertFalse(server.single_model) + self.assertIsNone(server.batch_config) + + def test_run_invalid_backend_raises(self): + server = TranscriptionServer() + with self.assertRaises(ValueError): + server.run(host="localhost", port=9090, backend="nonexistent") + + def test_run_invalid_trt_path_raises(self): + server = TranscriptionServer() + with self.assertRaises(ValueError): + server.run( + host="localhost", + port=9090, + backend="tensorrt", + whisper_tensorrt_path="/nonexistent/path", + ) + + +class TestTranscriptionServerGetAudio(unittest.TestCase): + def setUp(self): + self.server = TranscriptionServer() + + def test_end_of_audio_returns_false(self): + ws = MagicMock() + ws.recv.return_value = b"END_OF_AUDIO" + result = self.server.get_audio_from_websocket(ws) + self.assertFalse(result) + + def test_valid_audio_returns_numpy(self): + import numpy as np + ws = MagicMock() + audio = np.array([0.1, 0.2, 0.3], dtype=np.float32) + ws.recv.return_value = audio.tobytes() + result = self.server.get_audio_from_websocket(ws) + np.testing.assert_array_almost_equal(result, audio) + + +class TestTranscriptionServerHandleNewConnection(unittest.TestCase): + def setUp(self): + self.server = TranscriptionServer() + self.server.client_manager = ClientManager(max_clients=4, max_connection_time=600) + self.server.cache_path = "~/.cache/whisper-live/" + self.server.backend = BackendType.FASTER_WHISPER + + @mock.patch("websockets.WebSocketCommonProtocol") + def test_invalid_json_returns_false(self, mock_ws): + mock_ws.recv.return_value = "not valid json {{" + result = self.server.handle_new_connection(mock_ws, None, None, False) + self.assertFalse(result) + + @mock.patch("websockets.WebSocketCommonProtocol") + def test_server_full_returns_false(self, mock_ws): + # Fill server + for i in range(4): + self.server.client_manager.add_client(MagicMock(), MagicMock()) + + mock_ws.recv.return_value = json.dumps({ + "uid": "test", + "language": "en", + "task": "transcribe", + "model": "tiny.en", + }) + result = self.server.handle_new_connection(mock_ws, None, None, False) + self.assertFalse(result) + + +class TestTranscriptionServerCleanup(unittest.TestCase): + def setUp(self): + self.server = TranscriptionServer() + self.server.client_manager = ClientManager(max_clients=4, max_connection_time=600) + + def test_cleanup_removes_client(self): + ws = MagicMock() + client = MagicMock() + self.server.client_manager.add_client(ws, client) + self.server.cleanup(ws) + self.assertNotIn(ws, self.server.client_manager.clients) + client.cleanup.assert_called_once() + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_utils.py b/tests/test_utils.py new file mode 100644 index 0000000..2dc42e8 --- /dev/null +++ b/tests/test_utils.py @@ -0,0 +1,134 @@ +import os +import tempfile +import unittest +from io import StringIO +from unittest.mock import patch + +from whisper_live.utils import format_time, create_srt_file, print_transcript + + +class TestFormatTime(unittest.TestCase): + def test_zero(self): + self.assertEqual(format_time(0), "00:00:00,000") + + def test_seconds_only(self): + self.assertEqual(format_time(5.0), "00:00:05,000") + + def test_fractional_seconds(self): + self.assertEqual(format_time(1.5), "00:00:01,500") + + def test_minutes(self): + self.assertEqual(format_time(65.0), "00:01:05,000") + + def test_hours(self): + self.assertEqual(format_time(3661.123), "01:01:01,123") + + def test_millisecond_precision(self): + self.assertEqual(format_time(0.001), "00:00:00,001") + + def test_large_value(self): + # float precision: int((86399.999 - 86399) * 1000) may be 998 or 999 + result = format_time(86399.999) + self.assertIn(result, ("23:59:59,998", "23:59:59,999")) + + def test_rounding_edge(self): + result = format_time(0.9999) + # 0.9999 -> int(s%60)=0, milliseconds=int(0.9999*1000)=999 + self.assertEqual(result, "00:00:00,999") + + +class TestCreateSrtFile(unittest.TestCase): + def test_single_segment(self): + segments = [{"start": "0.000", "end": "1.500", "text": "Hello world"}] + with tempfile.NamedTemporaryFile(mode="w", suffix=".srt", delete=False) as f: + path = f.name + try: + create_srt_file(segments, path) + with open(path, "r", encoding="utf-8") as f: + content = f.read() + self.assertIn("1\n", content) + self.assertIn("00:00:00,000 --> 00:00:01,500", content) + self.assertIn("Hello world", content) + finally: + os.remove(path) + + def test_multiple_segments(self): + segments = [ + {"start": "0.000", "end": "1.000", "text": "First"}, + {"start": "1.000", "end": "2.500", "text": "Second"}, + {"start": "2.500", "end": "4.000", "text": "Third"}, + ] + with tempfile.NamedTemporaryFile(mode="w", suffix=".srt", delete=False) as f: + path = f.name + try: + create_srt_file(segments, path) + with open(path, "r", encoding="utf-8") as f: + content = f.read() + self.assertIn("1\n", content) + self.assertIn("2\n", content) + self.assertIn("3\n", content) + self.assertIn("First", content) + self.assertIn("Third", content) + finally: + os.remove(path) + + def test_empty_segments(self): + with tempfile.NamedTemporaryFile(mode="w", suffix=".srt", delete=False) as f: + path = f.name + try: + create_srt_file([], path) + with open(path, "r", encoding="utf-8") as f: + content = f.read() + self.assertEqual(content, "") + finally: + os.remove(path) + + def test_unicode_text(self): + segments = [{"start": "0.000", "end": "1.000", "text": "日本語テスト"}] + with tempfile.NamedTemporaryFile(mode="w", suffix=".srt", delete=False) as f: + path = f.name + try: + create_srt_file(segments, path) + with open(path, "r", encoding="utf-8") as f: + content = f.read() + self.assertIn("日本語テスト", content) + finally: + os.remove(path) + + +class TestPrintTranscript(unittest.TestCase): + @patch("sys.stdout", new_callable=StringIO) + def test_print_plain_text(self, mock_stdout): + text = ["Hello", " world"] + print_transcript(text) + output = mock_stdout.getvalue() + self.assertIn("Hello world", output) + + @patch("sys.stdout", new_callable=StringIO) + def test_print_with_timestamps(self, mock_stdout): + text = [ + {"start": 0.0, "end": 1.0, "text": "Hello"}, + {"start": 1.0, "end": 2.0, "text": "world"}, + ] + print_transcript(text, timestamps=True) + output = mock_stdout.getvalue() + self.assertIn("[0.0 -> 1.0]", output) + self.assertIn("Hello", output) + + @patch("sys.stdout", new_callable=StringIO) + def test_print_translated(self, mock_stdout): + text = ["Bonjour", "le monde"] + print_transcript(text, translated=True) + output = mock_stdout.getvalue() + self.assertIn("Bonjour le monde", output) + + @patch("sys.stdout", new_callable=StringIO) + def test_print_empty(self, mock_stdout): + print_transcript([]) + output = mock_stdout.getvalue() + # empty text joined is empty string, should not crash + self.assertEqual(output.strip(), "") + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_vad_extended.py b/tests/test_vad_extended.py new file mode 100644 index 0000000..333e410 --- /dev/null +++ b/tests/test_vad_extended.py @@ -0,0 +1,131 @@ +import unittest +from unittest.mock import patch, MagicMock + +import numpy as np +import torch + +from whisper_live.vad import VoiceActivityDetection, VoiceActivityDetector + + +class TestVoiceActivityDetectionValidation(unittest.TestCase): + """Tests for VoiceActivityDetection input validation without requiring the ONNX model.""" + + @patch.object(VoiceActivityDetection, "__init__", lambda self, **kw: None) + def setUp(self): + self.vad = VoiceActivityDetection() + self.vad.sample_rates = [8000, 16000] + + def test_1d_input_unsqueezed(self): + x = torch.randn(512) + x_out, sr_out = self.vad._validate_input(x, 16000) + self.assertEqual(x_out.dim(), 2) + self.assertEqual(sr_out, 16000) + + def test_3d_input_raises(self): + x = torch.randn(1, 1, 512) + with self.assertRaises(ValueError): + self.vad._validate_input(x, 16000) + + def test_unsupported_sample_rate_raises(self): + x = torch.randn(1, 512) + with self.assertRaises(ValueError): + self.vad._validate_input(x, 44100) + + def test_too_short_audio_raises(self): + x = torch.randn(1, 1) + with self.assertRaises(ValueError): + self.vad._validate_input(x, 16000) + + def test_downsample_multiple_of_16k(self): + x = torch.randn(1, 512 * 3) + x_out, sr_out = self.vad._validate_input(x, 48000) + self.assertEqual(sr_out, 16000) + self.assertEqual(x_out.shape[1], 512) + + +class TestVoiceActivityDetectionStateReset(unittest.TestCase): + """Tests for VoiceActivityDetection.reset_states().""" + + @patch.object(VoiceActivityDetection, "__init__", lambda self, **kw: None) + def setUp(self): + self.vad = VoiceActivityDetection() + + def test_reset_creates_correct_shapes(self): + self.vad.reset_states(batch_size=4) + self.assertEqual(self.vad._state.shape, (2, 4, 128)) + self.assertEqual(self.vad._context.shape[0], 0) + self.assertEqual(self.vad._last_sr, 0) + self.assertEqual(self.vad._last_batch_size, 0) + + def test_reset_default_batch_size(self): + self.vad.reset_states() + self.assertEqual(self.vad._state.shape, (2, 1, 128)) + + +class TestVoiceActivityDetectionDownload(unittest.TestCase): + """Tests for the model download function.""" + + @patch("os.path.exists", return_value=True) + def test_skips_download_if_exists(self, mock_exists): + path = VoiceActivityDetection.download() + self.assertTrue(path.endswith("silero_vad.onnx")) + + @patch("os.path.exists", return_value=False) + @patch("subprocess.run") + @patch("os.makedirs") + def test_downloads_if_missing(self, mock_makedirs, mock_run, mock_exists): + path = VoiceActivityDetection.download() + mock_run.assert_called_once() + self.assertIn("silero_vad.onnx", path) + + @patch("os.path.exists", return_value=False) + @patch("subprocess.run", side_effect=Exception("wget not found")) + @patch("os.makedirs") + def test_handles_download_failure(self, mock_makedirs, mock_run, mock_exists): + # should not raise, just prints an error + with self.assertRaises(Exception): + VoiceActivityDetection.download() + + +class TestVoiceActivityDetectorThreshold(unittest.TestCase): + """Tests for VoiceActivityDetector threshold behavior.""" + + @patch.object(VoiceActivityDetection, "__init__", lambda self, **kw: None) + def test_above_threshold_returns_true(self): + detector = VoiceActivityDetector.__new__(VoiceActivityDetector) + detector.model = VoiceActivityDetection() + detector.threshold = 0.5 + detector.frame_rate = 16000 + + mock_probs = torch.tensor([[0.9, 0.8, 0.7]]) + with patch.object(detector.model, "audio_forward", return_value=mock_probs): + result = detector(np.random.randn(16000).astype(np.float32)) + self.assertTrue(result) + + @patch.object(VoiceActivityDetection, "__init__", lambda self, **kw: None) + def test_below_threshold_returns_false(self): + detector = VoiceActivityDetector.__new__(VoiceActivityDetector) + detector.model = VoiceActivityDetection() + detector.threshold = 0.5 + detector.frame_rate = 16000 + + mock_probs = torch.tensor([[0.1, 0.2, 0.3]]) + with patch.object(detector.model, "audio_forward", return_value=mock_probs): + result = detector(np.random.randn(16000).astype(np.float32)) + self.assertFalse(result) + + @patch.object(VoiceActivityDetection, "__init__", lambda self, **kw: None) + def test_custom_threshold(self): + detector = VoiceActivityDetector.__new__(VoiceActivityDetector) + detector.model = VoiceActivityDetection() + detector.threshold = 0.95 + detector.frame_rate = 16000 + + mock_probs = torch.tensor([[0.9]]) + with patch.object(detector.model, "audio_forward", return_value=mock_probs): + result = detector(np.random.randn(16000).astype(np.float32)) + self.assertFalse(result) + + +if __name__ == "__main__": + unittest.main()