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([])
|
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__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
+30
-10
@@ -47,6 +47,8 @@ class Client:
|
|||||||
enable_diarization=False,
|
enable_diarization=False,
|
||||||
max_speakers=10,
|
max_speakers=10,
|
||||||
word_timestamps=False,
|
word_timestamps=False,
|
||||||
|
max_retries=0,
|
||||||
|
retry_delay=5,
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
Initializes a Client instance for audio recording and streaming to a server.
|
Initializes a Client instance for audio recording and streaming to a server.
|
||||||
@@ -109,20 +111,17 @@ class Client:
|
|||||||
self.enable_diarization = enable_diarization
|
self.enable_diarization = enable_diarization
|
||||||
self.max_speakers = max_speakers
|
self.max_speakers = max_speakers
|
||||||
self.word_timestamps = word_timestamps
|
self.word_timestamps = word_timestamps
|
||||||
|
self.max_retries = max_retries
|
||||||
|
self.retry_delay = retry_delay
|
||||||
|
self._retry_count = 0
|
||||||
self.audio_bytes = None
|
self.audio_bytes = None
|
||||||
|
|
||||||
if host is not None and port is not 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_protocol = 'wss' if self.use_wss else "ws"
|
||||||
socket_url = f"{socket_protocol}://{host}:{port}"
|
self.socket_url = f"{socket_protocol}://{host}:{port}"
|
||||||
self.client_socket = websocket.WebSocketApp(
|
self._create_websocket()
|
||||||
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
|
|
||||||
),
|
|
||||||
)
|
|
||||||
else:
|
else:
|
||||||
print("[ERROR]: No host or port specified.")
|
print("[ERROR]: No host or port specified.")
|
||||||
return
|
return
|
||||||
@@ -138,6 +137,18 @@ class Client:
|
|||||||
self.translated_transcript = []
|
self.translated_transcript = []
|
||||||
print("[INFO]: * recording")
|
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):
|
def handle_status_messages(self, message_data):
|
||||||
"""Handles server status messages."""
|
"""Handles server status messages."""
|
||||||
status = message_data["status"]
|
status = message_data["status"]
|
||||||
@@ -280,6 +291,15 @@ class Client:
|
|||||||
self.recording = False
|
self.recording = False
|
||||||
self.waiting = 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):
|
def on_open(self, ws):
|
||||||
"""
|
"""
|
||||||
Callback function called when the WebSocket connection is successfully opened.
|
Callback function called when the WebSocket connection is successfully opened.
|
||||||
|
|||||||
Reference in New Issue
Block a user