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:
Aaron Boxer
2026-04-17 09:33:52 -04:00
committed by Aaron Boxer
parent 52005b94cb
commit c5ec7f4a99
2 changed files with 78 additions and 10 deletions
+48
View File
@@ -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
View File
@@ -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.