Files
WhisperLive/tests/test_server_extended.py
T
Aaron Boxer 18bce1864a Validate server parameters on startup
- max_clients must be >= 1
- max_connection_time must be > 0
- batch_max_size must be >= 1 (when batch enabled)
- batch_window_ms must be >= 0 (when batch enabled)
- Added 5 new tests for parameter validation
2026-04-17 09:30:19 -04:00

330 lines
11 KiB
Python

import json
import time
import threading
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 TestClientManagerThreadSafety(unittest.TestCase):
def test_concurrent_add_remove(self):
cm = ClientManager(max_clients=100, max_connection_time=600)
errors = []
def add_clients(start_idx):
try:
for i in range(50):
ws = MagicMock(name=f"ws-{start_idx}-{i}")
client = MagicMock(name=f"client-{start_idx}-{i}")
cm.add_client(ws, client)
except Exception as e:
errors.append(e)
def remove_clients():
try:
for _ in range(25):
with cm.lock:
if cm.clients:
ws = next(iter(cm.clients))
else:
continue
cm.remove_client(ws)
except Exception as e:
errors.append(e)
threads = [
threading.Thread(target=add_clients, args=(0,)),
threading.Thread(target=add_clients, args=(1,)),
threading.Thread(target=remove_clients),
threading.Thread(target=remove_clients),
]
for t in threads:
t.start()
for t in threads:
t.join()
self.assertEqual(errors, [])
def test_concurrent_get_client(self):
cm = ClientManager(max_clients=100, max_connection_time=600)
ws = MagicMock()
client = MagicMock()
cm.add_client(ws, client)
errors = []
results = []
def get_many():
try:
for _ in range(100):
results.append(cm.get_client(ws))
except Exception as e:
errors.append(e)
threads = [threading.Thread(target=get_many) for _ in range(4)]
for t in threads:
t.start()
for t in threads:
t.join()
self.assertEqual(errors, [])
self.assertTrue(all(r is client for r in results))
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",
)
def test_run_max_clients_zero_raises(self):
server = TranscriptionServer()
with self.assertRaises(ValueError):
server.run(host="localhost", port=9090, max_clients=0)
def test_run_max_clients_negative_raises(self):
server = TranscriptionServer()
with self.assertRaises(ValueError):
server.run(host="localhost", port=9090, max_clients=-1)
def test_run_max_connection_time_zero_raises(self):
server = TranscriptionServer()
with self.assertRaises(ValueError):
server.run(host="localhost", port=9090, max_connection_time=0)
def test_run_batch_max_size_zero_raises(self):
server = TranscriptionServer()
with self.assertRaises(ValueError):
server.run(host="localhost", port=9090, batch_enabled=True, batch_max_size=0)
def test_run_batch_window_ms_negative_raises(self):
server = TranscriptionServer()
with self.assertRaises(ValueError):
server.run(host="localhost", port=9090, batch_enabled=True, batch_window_ms=-1)
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)
def test_raw_pcm_input_normalizes_int16(self):
import numpy as np
self.server.raw_pcm_input = True
ws = MagicMock()
pcm = np.array([0, 16384, -16384, 32767], dtype=np.int16)
ws.recv.return_value = pcm.tobytes()
result = self.server.get_audio_from_websocket(ws)
expected = pcm.astype(np.float32) / 32768.0
np.testing.assert_array_almost_equal(result, expected)
self.assertTrue(result.dtype == np.float32)
self.assertTrue(np.all(result >= -1.0))
self.assertTrue(np.all(result <= 1.0))
def test_raw_pcm_input_off_reads_float32(self):
import numpy as np
self.server.raw_pcm_input = False
ws = MagicMock()
audio = np.array([0.5, -0.5], 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()