Add reconnect logic to WebSocket client
- New params: max_retries (default 0), retry_delay (default 5s) - On unexpected close, retries up to max_retries times - Does not retry on server_error (server rejected connection) - Extracted _create_websocket() helper for reuse - Added 4 unit tests for reconnect behavior
This commit is contained in:
@@ -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()
|
||||
|
||||
+30
-10
@@ -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.
|
||||
|
||||
Reference in New Issue
Block a user