From c1420cba0d3b5e8a82206693153582fd3894f937 Mon Sep 17 00:00:00 2001 From: makaveli10 Date: Thu, 15 Feb 2024 12:16:58 +0530 Subject: [PATCH] add tests for server exception handling --- tests/test_server.py | 34 ++++++++++++++++++++++++++++++++++ 1 file changed, 34 insertions(+) diff --git a/tests/test_server.py b/tests/test_server.py index f563f9a..cd14bb3 100644 --- a/tests/test_server.py +++ b/tests/test_server.py @@ -6,6 +6,8 @@ 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 TranscriptionClient from whisper.normalizers import EnglishTextNormalizer @@ -101,3 +103,35 @@ class TestServerInferenceAccuracy(unittest.TestCase): references=[gt_normalized] ) self.assertLess(wer, 0.05) + + +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: Unexpected error" in message for message in log.output))