First commit of files
This commit is contained in:
@@ -0,0 +1,156 @@
|
||||
import json
|
||||
import os
|
||||
import scipy
|
||||
import websocket
|
||||
import copy
|
||||
import unittest
|
||||
from unittest.mock import patch, MagicMock
|
||||
from whisper_live.client import Client, TranscriptionClient, TranscriptionTeeClient
|
||||
from whisper_live.utils import resample
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
class BaseTestCase(unittest.TestCase):
|
||||
@patch('whisper_live.client.websocket.WebSocketApp')
|
||||
@patch('whisper_live.client.pyaudio.PyAudio')
|
||||
def setUp(self, mock_pyaudio, mock_websocket):
|
||||
self.mock_pyaudio_instance = MagicMock()
|
||||
mock_pyaudio.return_value = self.mock_pyaudio_instance
|
||||
self.mock_stream = MagicMock()
|
||||
self.mock_pyaudio_instance.open.return_value = self.mock_stream
|
||||
|
||||
self.mock_ws_app = mock_websocket.return_value
|
||||
self.mock_ws_app.send = MagicMock()
|
||||
|
||||
self.client = TranscriptionClient(host='localhost', port=9090, lang="en").client
|
||||
|
||||
self.mock_pyaudio = mock_pyaudio
|
||||
self.mock_websocket = mock_websocket
|
||||
self.mock_audio_packet = b'\x00\x01\x02\x03'
|
||||
|
||||
def tearDown(self):
|
||||
self.client.close_websocket()
|
||||
self.mock_pyaudio.stop()
|
||||
self.mock_websocket.stop()
|
||||
del self.client
|
||||
|
||||
class TestClientWebSocketCommunication(BaseTestCase):
|
||||
def test_websocket_communication(self):
|
||||
expected_url = 'ws://localhost:9090'
|
||||
self.mock_websocket.assert_called()
|
||||
self.assertEqual(self.mock_websocket.call_args[0][0], expected_url)
|
||||
|
||||
|
||||
class TestClientCallbacks(BaseTestCase):
|
||||
def test_on_open(self):
|
||||
expected_message = json.dumps({
|
||||
"uid": self.client.uid,
|
||||
"language": self.client.language,
|
||||
"task": self.client.task,
|
||||
"model": self.client.model,
|
||||
"use_vad": True
|
||||
})
|
||||
self.client.on_open(self.mock_ws_app)
|
||||
self.mock_ws_app.send.assert_called_with(expected_message)
|
||||
|
||||
def test_on_message(self):
|
||||
message = json.dumps(
|
||||
{
|
||||
"uid": self.client.uid,
|
||||
"message": "SERVER_READY",
|
||||
"backend": "faster_whisper"
|
||||
}
|
||||
)
|
||||
self.client.on_message(self.mock_ws_app, message)
|
||||
|
||||
message = json.dumps({
|
||||
"uid": self.client.uid,
|
||||
"segments": [
|
||||
{"start": 0, "end": 1, "text": "Test transcript"},
|
||||
{"start": 1, "end": 2, "text": "Test transcript 2"},
|
||||
{"start": 2, "end": 3, "text": "Test transcript 3"}
|
||||
]
|
||||
})
|
||||
self.client.on_message(self.mock_ws_app, message)
|
||||
|
||||
# Assert that the transcript was updated correctly
|
||||
self.assertEqual(len(self.client.transcript), 2)
|
||||
self.assertEqual(self.client.transcript[1]['text'], "Test transcript 2")
|
||||
|
||||
def test_on_close(self):
|
||||
close_status_code = 1000
|
||||
close_msg = "Normal closure"
|
||||
self.client.on_close(self.mock_ws_app, close_status_code, close_msg)
|
||||
|
||||
self.assertFalse(self.client.recording)
|
||||
self.assertFalse(self.client.server_error)
|
||||
self.assertFalse(self.client.waiting)
|
||||
|
||||
def test_on_error(self):
|
||||
error_message = "Test Error"
|
||||
self.client.on_error(self.mock_ws_app, error_message)
|
||||
|
||||
self.assertTrue(self.client.server_error)
|
||||
self.assertEqual(self.client.error_message, error_message)
|
||||
|
||||
|
||||
class TestAudioResampling(unittest.TestCase):
|
||||
def test_resample_audio(self):
|
||||
original_audio = "assets/jfk.flac"
|
||||
expected_sr = 16000
|
||||
resampled_audio = resample(original_audio, expected_sr)
|
||||
|
||||
sr, _ = scipy.io.wavfile.read(resampled_audio)
|
||||
self.assertEqual(sr, expected_sr)
|
||||
|
||||
os.remove(resampled_audio)
|
||||
|
||||
|
||||
class TestSendingAudioPacket(BaseTestCase):
|
||||
def test_send_packet(self):
|
||||
self.client.send_packet_to_server(self.mock_audio_packet)
|
||||
self.client.client_socket.send.assert_called_with(self.mock_audio_packet, websocket.ABNF.OPCODE_BINARY)
|
||||
|
||||
class TestTee(BaseTestCase):
|
||||
@patch('whisper_live.client.websocket.WebSocketApp')
|
||||
@patch('whisper_live.client.pyaudio.PyAudio')
|
||||
def setUp(self, mock_audio, mock_websocket):
|
||||
super().setUp()
|
||||
self.client2 = Client(host='localhost', port=9090, lang="es", translate=False, srt_file_path="transcript.srt")
|
||||
self.client3 = Client(host='localhost', port=9090, lang="es", translate=True, srt_file_path="translation.srt")
|
||||
# need a separate mock for each websocket
|
||||
self.client3.client_socket = copy.deepcopy(self.client3.client_socket)
|
||||
self.tee = TranscriptionTeeClient([self.client2, self.client3])
|
||||
|
||||
def tearDown(self):
|
||||
self.tee.close_all_clients()
|
||||
del self.tee
|
||||
super().tearDown()
|
||||
|
||||
def test_invalid_constructor(self):
|
||||
with self.assertRaises(Exception) as context:
|
||||
TranscriptionTeeClient([])
|
||||
|
||||
def test_multicast_unconditional(self):
|
||||
self.tee.multicast_packet(self.mock_audio_packet, True)
|
||||
for client in self.tee.clients:
|
||||
client.client_socket.send.assert_called_with(self.mock_audio_packet, websocket.ABNF.OPCODE_BINARY)
|
||||
|
||||
def test_multicast_conditional(self):
|
||||
self.client2.recording = False
|
||||
self.client3.recording = True
|
||||
self.tee.multicast_packet(self.mock_audio_packet, False)
|
||||
self.client2.client_socket.send.assert_not_called()
|
||||
self.client3.client_socket.send.assert_called_with(self.mock_audio_packet, websocket.ABNF.OPCODE_BINARY)
|
||||
|
||||
def test_close_all(self):
|
||||
self.tee.close_all_clients()
|
||||
for client in self.tee.clients:
|
||||
client.client_socket.close.assert_called()
|
||||
|
||||
def test_write_all_srt(self):
|
||||
for client in self.tee.clients:
|
||||
client.server_backend = "faster_whisper"
|
||||
self.tee.write_all_clients_srt()
|
||||
self.assertTrue(Path("transcript.srt").is_file())
|
||||
self.assertTrue(Path("translation.srt").is_file())
|
||||
@@ -0,0 +1,150 @@
|
||||
import subprocess
|
||||
import time
|
||||
import json
|
||||
import unittest
|
||||
from unittest import mock
|
||||
|
||||
import numpy as np
|
||||
import evaluate
|
||||
|
||||
from websockets.exceptions import ConnectionClosed
|
||||
from whisper_live.server import TranscriptionServer
|
||||
from whisper_live.client import Client, TranscriptionClient, TranscriptionTeeClient
|
||||
from whisper.normalizers import EnglishTextNormalizer
|
||||
|
||||
|
||||
class TestTranscriptionServerInitialization(unittest.TestCase):
|
||||
def test_initialization(self):
|
||||
server = TranscriptionServer()
|
||||
self.assertEqual(server.client_manager.max_clients, 4)
|
||||
self.assertEqual(server.client_manager.max_connection_time, 600)
|
||||
self.assertDictEqual(server.client_manager.clients, {})
|
||||
self.assertDictEqual(server.client_manager.start_times, {})
|
||||
|
||||
|
||||
class TestGetWaitTime(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.server = TranscriptionServer()
|
||||
self.server.client_manager.start_times = {
|
||||
'client1': time.time() - 120,
|
||||
'client2': time.time() - 300
|
||||
}
|
||||
self.server.client_manager.max_connection_time = 600
|
||||
|
||||
def test_get_wait_time(self):
|
||||
expected_wait_time = (600 - (time.time() - self.server.client_manager.start_times['client2'])) / 60
|
||||
print(self.server.client_manager.get_wait_time(), expected_wait_time)
|
||||
self.assertAlmostEqual(self.server.client_manager.get_wait_time(), expected_wait_time, places=2)
|
||||
|
||||
|
||||
class TestServerConnection(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.server = TranscriptionServer()
|
||||
|
||||
@mock.patch('websockets.WebSocketCommonProtocol')
|
||||
def test_connection(self, mock_websocket):
|
||||
mock_websocket.recv.return_value = json.dumps({
|
||||
'uid': 'test_client',
|
||||
'language': 'en',
|
||||
'task': 'transcribe',
|
||||
'model': 'tiny.en'
|
||||
})
|
||||
self.server.recv_audio(mock_websocket, "faster_whisper")
|
||||
|
||||
@mock.patch('websockets.WebSocketCommonProtocol')
|
||||
def test_recv_audio_exception_handling(self, mock_websocket):
|
||||
mock_websocket.recv.side_effect = [json.dumps({
|
||||
'uid': 'test_client',
|
||||
'language': 'en',
|
||||
'task': 'transcribe',
|
||||
'model': 'tiny.en'
|
||||
}), np.array([1, 2, 3]).tobytes()]
|
||||
|
||||
with self.assertLogs(level="ERROR"):
|
||||
self.server.recv_audio(mock_websocket, "faster_whisper")
|
||||
|
||||
self.assertNotIn(mock_websocket, self.server.client_manager.clients)
|
||||
|
||||
|
||||
class TestServerInferenceAccuracy(unittest.TestCase):
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
cls.mock_pyaudio_patch = mock.patch('pyaudio.PyAudio')
|
||||
cls.mock_pyaudio = cls.mock_pyaudio_patch.start()
|
||||
cls.mock_pyaudio.return_value.open.return_value = mock.MagicMock()
|
||||
|
||||
cls.server_process = subprocess.Popen(["python", "run_server.py"])
|
||||
time.sleep(2)
|
||||
|
||||
@classmethod
|
||||
def tearDownClass(cls):
|
||||
cls.server_process.terminate()
|
||||
cls.server_process.wait()
|
||||
|
||||
def setUp(self):
|
||||
self.metric = evaluate.load("wer")
|
||||
self.normalizer = EnglishTextNormalizer()
|
||||
|
||||
def check_prediction(self, srt_path):
|
||||
gt = "And so my fellow Americans, ask not, what your country can do for you. Ask what you can do for your country!"
|
||||
with open(srt_path, "r") as f:
|
||||
lines = f.readlines()
|
||||
prediction = " ".join([line.strip() for line in lines[2::4]])
|
||||
prediction_normalized = self.normalizer(prediction)
|
||||
gt_normalized = self.normalizer(gt)
|
||||
|
||||
# calculate WER
|
||||
wer = self.metric.compute(
|
||||
predictions=[prediction_normalized],
|
||||
references=[gt_normalized]
|
||||
)
|
||||
self.assertLess(wer, 0.05)
|
||||
|
||||
def test_inference(self):
|
||||
client = TranscriptionClient(
|
||||
"localhost", "9090", model="base.en", lang="en",
|
||||
)
|
||||
client("assets/jfk.flac")
|
||||
self.check_prediction("output.srt")
|
||||
|
||||
def test_simultaneous_inference(self):
|
||||
client1 = Client(
|
||||
"localhost", "9090", model="base.en", lang="en", srt_file_path="transcript1.srt")
|
||||
client2 = Client(
|
||||
"localhost", "9090", model="base.en", lang="en", srt_file_path="transcript2.srt")
|
||||
tee = TranscriptionTeeClient([client1, client2])
|
||||
tee("assets/jfk.flac")
|
||||
self.check_prediction("transcript1.srt")
|
||||
self.check_prediction("transcript2.srt")
|
||||
|
||||
|
||||
class TestExceptionHandling(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.server = TranscriptionServer()
|
||||
|
||||
@mock.patch('websockets.WebSocketCommonProtocol')
|
||||
def test_connection_closed_exception(self, mock_websocket):
|
||||
mock_websocket.recv.side_effect = ConnectionClosed(1001, "testing connection closed")
|
||||
|
||||
with self.assertLogs(level="INFO") as log:
|
||||
self.server.recv_audio(mock_websocket, "faster_whisper")
|
||||
self.assertTrue(any("Connection closed by client" in message for message in log.output))
|
||||
|
||||
@mock.patch('websockets.WebSocketCommonProtocol')
|
||||
def test_json_decode_exception(self, mock_websocket):
|
||||
mock_websocket.recv.return_value = "invalid json"
|
||||
|
||||
with self.assertLogs(level="ERROR") as log:
|
||||
self.server.recv_audio(mock_websocket, "faster_whisper")
|
||||
self.assertTrue(any("Failed to decode JSON from client" in message for message in log.output))
|
||||
|
||||
@mock.patch('websockets.WebSocketCommonProtocol')
|
||||
def test_unexpected_exception_handling(self, mock_websocket):
|
||||
mock_websocket.recv.side_effect = RuntimeError("Unexpected error")
|
||||
|
||||
with self.assertLogs(level="ERROR") as log:
|
||||
self.server.recv_audio(mock_websocket, "faster_whisper")
|
||||
for message in log.output:
|
||||
print(message)
|
||||
print()
|
||||
self.assertTrue(any("Unexpected error" in message for message in log.output))
|
||||
@@ -0,0 +1,26 @@
|
||||
import unittest
|
||||
import numpy as np
|
||||
from whisper_live.tensorrt_utils import load_audio
|
||||
from whisper_live.vad import VoiceActivityDetector
|
||||
|
||||
|
||||
class TestVoiceActivityDetection(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.vad = VoiceActivityDetector()
|
||||
self.sample_rate = 16000
|
||||
|
||||
def generate_silence(self, duration_seconds):
|
||||
return np.zeros(int(self.sample_rate * duration_seconds), dtype=np.float32)
|
||||
|
||||
def load_speech_segment(self, filepath):
|
||||
return load_audio(filepath)
|
||||
|
||||
def test_vad_silence_detection(self):
|
||||
silence = self.generate_silence(3)
|
||||
is_speech_present = self.vad(silence.copy())
|
||||
self.assertFalse(is_speech_present, "VAD incorrectly identified silence as speech.")
|
||||
|
||||
def test_vad_speech_detection(self):
|
||||
audio_tensor = load_audio("assets/jfk.flac")
|
||||
is_speech_present = self.vad(audio_tensor)
|
||||
self.assertTrue(is_speech_present, "VAD failed to identify speech segment.")
|
||||
Reference in New Issue
Block a user