8396763444
MagicMock auto-attributes are not JSON serializable. Set language, language_probability, and duration explicitly. Also exclude metadata events from segment count assertion.
490 lines
17 KiB
Python
490 lines
17 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()
|
|
|
|
|
|
class TestStreamTranscription(unittest.TestCase):
|
|
"""Tests for the SSE streaming endpoint (stream=true)."""
|
|
|
|
def _make_app(self):
|
|
"""Create a FastAPI app with the transcribe endpoint that has streaming support."""
|
|
from fastapi import FastAPI, UploadFile, Form
|
|
from fastapi.testclient import TestClient
|
|
from starlette.responses import StreamingResponse
|
|
import os
|
|
import tempfile
|
|
import shutil
|
|
|
|
app = FastAPI()
|
|
server = TranscriptionServer()
|
|
|
|
@app.post("/v1/audio/transcriptions")
|
|
async def transcribe(
|
|
file: UploadFile,
|
|
stream: bool = Form(default=False),
|
|
language: str = Form(default=None),
|
|
response_format: str = Form(default="json"),
|
|
):
|
|
if stream:
|
|
return server._stream_transcription(
|
|
file, language, None, 0.0, None, None
|
|
)
|
|
return {"text": "non-streamed"}
|
|
|
|
return app
|
|
|
|
@patch("whisper_live.server.WhisperModel")
|
|
def test_stream_returns_sse_content_type(self, mock_model_cls):
|
|
mock_seg = MagicMock()
|
|
mock_seg.id = 0
|
|
mock_seg.start = 0.0
|
|
mock_seg.end = 1.0
|
|
mock_seg.text = " hello "
|
|
mock_seg.words = []
|
|
|
|
mock_info = MagicMock()
|
|
mock_info.language = "en"
|
|
mock_info.language_probability = 0.98
|
|
mock_info.duration = 1.0
|
|
|
|
mock_model = MagicMock()
|
|
mock_model.transcribe.return_value = (iter([mock_seg]), mock_info)
|
|
mock_model_cls.return_value = mock_model
|
|
|
|
import io
|
|
from fastapi.testclient import TestClient
|
|
|
|
app = self._make_app()
|
|
client = TestClient(app)
|
|
resp = client.post(
|
|
"/v1/audio/transcriptions",
|
|
files={"file": ("test.wav", io.BytesIO(b"\x00" * 100), "audio/wav")},
|
|
data={"stream": "true"},
|
|
)
|
|
self.assertEqual(resp.status_code, 200)
|
|
self.assertIn("text/event-stream", resp.headers.get("content-type", ""))
|
|
|
|
@patch("whisper_live.server.WhisperModel")
|
|
def test_stream_yields_segment_and_done(self, mock_model_cls):
|
|
mock_seg = MagicMock()
|
|
mock_seg.id = 0
|
|
mock_seg.start = 0.0
|
|
mock_seg.end = 1.5
|
|
mock_seg.text = " hello world "
|
|
mock_seg.words = []
|
|
|
|
mock_info = MagicMock()
|
|
mock_info.language = "en"
|
|
mock_info.language_probability = 0.95
|
|
mock_info.duration = 1.5
|
|
mock_model = MagicMock()
|
|
mock_model.transcribe.return_value = (iter([mock_seg]), mock_info)
|
|
mock_model_cls.return_value = mock_model
|
|
|
|
import io
|
|
from fastapi.testclient import TestClient
|
|
|
|
app = self._make_app()
|
|
client = TestClient(app)
|
|
resp = client.post(
|
|
"/v1/audio/transcriptions",
|
|
files={"file": ("test.wav", io.BytesIO(b"\x00" * 100), "audio/wav")},
|
|
data={"stream": "true"},
|
|
)
|
|
body = resp.text
|
|
self.assertIn('"text": "hello world"', body)
|
|
self.assertIn("[DONE]", body)
|
|
|
|
@patch("whisper_live.server.WhisperModel")
|
|
def test_stream_multiple_segments(self, mock_model_cls):
|
|
segs = []
|
|
for i in range(3):
|
|
s = MagicMock()
|
|
s.id = i
|
|
s.start = float(i)
|
|
s.end = float(i + 1)
|
|
s.text = f" segment {i} "
|
|
s.words = []
|
|
segs.append(s)
|
|
|
|
mock_info = MagicMock()
|
|
mock_info.language = "en"
|
|
mock_info.language_probability = 0.99
|
|
mock_info.duration = 3.0
|
|
mock_model = MagicMock()
|
|
mock_model.transcribe.return_value = (iter(segs), mock_info)
|
|
mock_model_cls.return_value = mock_model
|
|
|
|
import io
|
|
from fastapi.testclient import TestClient
|
|
|
|
app = self._make_app()
|
|
client = TestClient(app)
|
|
resp = client.post(
|
|
"/v1/audio/transcriptions",
|
|
files={"file": ("test.wav", io.BytesIO(b"\x00" * 100), "audio/wav")},
|
|
data={"stream": "true"},
|
|
)
|
|
body = resp.text
|
|
events = [line for line in body.split("\n") if line.startswith("data: ") and "[DONE]" not in line and '"type": "metadata"' not in line]
|
|
self.assertEqual(len(events), 3)
|
|
for i, event in enumerate(events):
|
|
data = json.loads(event.removeprefix("data: "))
|
|
self.assertEqual(data["text"], f"segment {i}")
|
|
|
|
@patch("whisper_live.server.WhisperModel", side_effect=RuntimeError("model error"))
|
|
def test_stream_error_yields_error_event(self, mock_model_cls):
|
|
import io
|
|
from fastapi.testclient import TestClient
|
|
|
|
app = self._make_app()
|
|
client = TestClient(app)
|
|
resp = client.post(
|
|
"/v1/audio/transcriptions",
|
|
files={"file": ("test.wav", io.BytesIO(b"\x00" * 100), "audio/wav")},
|
|
data={"stream": "true"},
|
|
)
|
|
body = resp.text
|
|
self.assertIn('"error"', body)
|
|
self.assertIn("model error", body)
|
|
|
|
def test_non_stream_still_works(self):
|
|
import io
|
|
from fastapi.testclient import TestClient
|
|
|
|
app = self._make_app()
|
|
client = TestClient(app)
|
|
resp = client.post(
|
|
"/v1/audio/transcriptions",
|
|
files={"file": ("test.wav", io.BytesIO(b"\x00" * 100), "audio/wav")},
|
|
data={"stream": "false"},
|
|
)
|
|
self.assertEqual(resp.status_code, 200)
|
|
self.assertEqual(resp.json()["text"], "non-streamed")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|