diff --git a/tests/test_client_extended.py b/tests/test_client_extended.py index 087bae7..1fedabf 100644 --- a/tests/test_client_extended.py +++ b/tests/test_client_extended.py @@ -253,5 +253,53 @@ class TestTeeClientEdgeCases(unittest.TestCase): TranscriptionTeeClient([]) +class TestClientReconnect(unittest.TestCase): + """Tests for reconnection logic.""" + + @patch("whisper_live.client.websocket.WebSocketApp") + @patch("whisper_live.client.pyaudio.PyAudio") + def test_reconnect_on_close(self, mock_pyaudio, mock_websocket): + mock_pyaudio.return_value.open.return_value = MagicMock() + client = Client(host="localhost", port=9090, lang="en", max_retries=2, retry_delay=0) + initial_socket = client.client_socket + client.on_close(MagicMock(), 1006, "abnormal closure") + self.assertEqual(client._retry_count, 1) + # A new websocket should have been created + self.assertIsNotNone(client.client_socket) + client.close_websocket() + + @patch("whisper_live.client.websocket.WebSocketApp") + @patch("whisper_live.client.pyaudio.PyAudio") + def test_no_reconnect_on_server_error(self, mock_pyaudio, mock_websocket): + mock_pyaudio.return_value.open.return_value = MagicMock() + client = Client(host="localhost", port=9090, lang="en", max_retries=2, retry_delay=0) + client.server_error = True + client.on_close(MagicMock(), 1000, "normal") + self.assertEqual(client._retry_count, 0) + client.close_websocket() + + @patch("whisper_live.client.websocket.WebSocketApp") + @patch("whisper_live.client.pyaudio.PyAudio") + def test_no_reconnect_when_max_retries_zero(self, mock_pyaudio, mock_websocket): + mock_pyaudio.return_value.open.return_value = MagicMock() + client = Client(host="localhost", port=9090, lang="en", max_retries=0, retry_delay=0) + client.on_close(MagicMock(), 1006, "abnormal closure") + self.assertEqual(client._retry_count, 0) + client.close_websocket() + + @patch("whisper_live.client.websocket.WebSocketApp") + @patch("whisper_live.client.pyaudio.PyAudio") + def test_stops_after_max_retries(self, mock_pyaudio, mock_websocket): + mock_pyaudio.return_value.open.return_value = MagicMock() + client = Client(host="localhost", port=9090, lang="en", max_retries=2, retry_delay=0) + client.on_close(MagicMock(), 1006, "closed") + client.on_close(MagicMock(), 1006, "closed") + self.assertEqual(client._retry_count, 2) + # third close should NOT retry + client.on_close(MagicMock(), 1006, "closed") + self.assertEqual(client._retry_count, 2) + client.close_websocket() + + if __name__ == "__main__": unittest.main() diff --git a/whisper_live/client.py b/whisper_live/client.py index 9a271fb..b30826f 100644 --- a/whisper_live/client.py +++ b/whisper_live/client.py @@ -47,6 +47,8 @@ class Client: enable_diarization=False, max_speakers=10, word_timestamps=False, + max_retries=0, + retry_delay=5, ): """ Initializes a Client instance for audio recording and streaming to a server. @@ -109,20 +111,17 @@ class Client: self.enable_diarization = enable_diarization self.max_speakers = max_speakers self.word_timestamps = word_timestamps + self.max_retries = max_retries + self.retry_delay = retry_delay + self._retry_count = 0 self.audio_bytes = None if host is not None and port is not None: + self.host = host + self.port = port socket_protocol = 'wss' if self.use_wss else "ws" - socket_url = f"{socket_protocol}://{host}:{port}" - self.client_socket = websocket.WebSocketApp( - socket_url, - on_open=lambda ws: self.on_open(ws), - on_message=lambda ws, message: self.on_message(ws, message), - on_error=lambda ws, error: self.on_error(ws, error), - on_close=lambda ws, close_status_code, close_msg: self.on_close( - ws, close_status_code, close_msg - ), - ) + self.socket_url = f"{socket_protocol}://{host}:{port}" + self._create_websocket() else: print("[ERROR]: No host or port specified.") return @@ -138,6 +137,18 @@ class Client: self.translated_transcript = [] print("[INFO]: * recording") + def _create_websocket(self): + """Creates a new WebSocketApp instance.""" + self.client_socket = websocket.WebSocketApp( + self.socket_url, + on_open=lambda ws: self.on_open(ws), + on_message=lambda ws, message: self.on_message(ws, message), + on_error=lambda ws, error: self.on_error(ws, error), + on_close=lambda ws, close_status_code, close_msg: self.on_close( + ws, close_status_code, close_msg + ), + ) + def handle_status_messages(self, message_data): """Handles server status messages.""" status = message_data["status"] @@ -280,6 +291,15 @@ class Client: self.recording = False self.waiting = False + if self.max_retries > 0 and self._retry_count < self.max_retries and not self.server_error: + self._retry_count += 1 + print(f"[INFO]: Reconnecting ({self._retry_count}/{self.max_retries}) in {self.retry_delay}s...") + time.sleep(self.retry_delay) + self._create_websocket() + self.ws_thread = threading.Thread(target=self.client_socket.run_forever) + self.ws_thread.daemon = True + self.ws_thread.start() + def on_open(self, ws): """ Callback function called when the WebSocket connection is successfully opened.