Extracted the logic for retrying connections, and create a new send_with_retry method inside WebSocketService.
This commit is contained in:
@@ -36,6 +36,7 @@ class WebsocketService(ABC):
|
|||||||
"""
|
"""
|
||||||
self._websocket: Optional[websockets.WebSocketClientProtocol] = None
|
self._websocket: Optional[websockets.WebSocketClientProtocol] = None
|
||||||
self._reconnect_on_error = reconnect_on_error
|
self._reconnect_on_error = reconnect_on_error
|
||||||
|
self._reconnect_in_progress: bool = False # Add this flag
|
||||||
|
|
||||||
async def _verify_connection(self) -> bool:
|
async def _verify_connection(self) -> bool:
|
||||||
"""Verify the websocket connection is active and responsive.
|
"""Verify the websocket connection is active and responsive.
|
||||||
@@ -66,6 +67,59 @@ class WebsocketService(ABC):
|
|||||||
await self._connect_websocket()
|
await self._connect_websocket()
|
||||||
return await self._verify_connection()
|
return await self._verify_connection()
|
||||||
|
|
||||||
|
async def _try_reconnect(
|
||||||
|
self,
|
||||||
|
max_retries: int = 3,
|
||||||
|
report_error: Optional[Callable[[ErrorFrame], Awaitable[None]]] = None,
|
||||||
|
) -> bool:
|
||||||
|
# Prevent concurrent reconnection attempts
|
||||||
|
if self._reconnect_in_progress:
|
||||||
|
logger.warning(f"{self} reconnect attempt aborted: already in progress")
|
||||||
|
return False
|
||||||
|
|
||||||
|
self._reconnect_in_progress = True
|
||||||
|
last_exception: Optional[Exception] = None
|
||||||
|
try:
|
||||||
|
for attempt in range(1, max_retries + 1):
|
||||||
|
try:
|
||||||
|
logger.warning(f"{self} reconnecting, attempt {attempt}")
|
||||||
|
if await self._reconnect_websocket(attempt):
|
||||||
|
logger.info(f"{self} reconnected successfully on attempt {attempt}")
|
||||||
|
return True
|
||||||
|
except Exception as e:
|
||||||
|
last_exception = e
|
||||||
|
logger.error(f"{self} reconnection attempt {attempt} failed: {e}")
|
||||||
|
if report_error:
|
||||||
|
await report_error(
|
||||||
|
ErrorFrame(f"{self} reconnection attempt {attempt} failed: {e}")
|
||||||
|
)
|
||||||
|
wait_time = exponential_backoff_time(attempt)
|
||||||
|
await asyncio.sleep(wait_time)
|
||||||
|
fatal_msg = f"{self} failed to reconnect after {max_retries} attempts"
|
||||||
|
if last_exception:
|
||||||
|
fatal_msg += f": {last_exception}"
|
||||||
|
logger.error(fatal_msg)
|
||||||
|
if report_error:
|
||||||
|
await report_error(ErrorFrame(fatal_msg, fatal=True))
|
||||||
|
return False
|
||||||
|
finally:
|
||||||
|
self._reconnect_in_progress = False
|
||||||
|
|
||||||
|
async def send_with_retry(self, message, report_error: Callable[[ErrorFrame], Awaitable[None]]):
|
||||||
|
"""Attempt to send a message, retrying after reconnect if necessary."""
|
||||||
|
try:
|
||||||
|
await self._websocket.send(message)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"{self} send failed: {e}, will try to reconnect")
|
||||||
|
# Try to reconnect before retrying
|
||||||
|
success = await self._try_reconnect(report_error=report_error)
|
||||||
|
if success:
|
||||||
|
logger.info(f"{self} reconnected successfully, will retry send the message")
|
||||||
|
# trying to send the message one more time
|
||||||
|
await self._websocket.send(message)
|
||||||
|
else:
|
||||||
|
logger.error(f"{self} send failed; unable to reconnect")
|
||||||
|
|
||||||
async def _receive_task_handler(self, report_error: Callable[[ErrorFrame], Awaitable[None]]):
|
async def _receive_task_handler(self, report_error: Callable[[ErrorFrame], Awaitable[None]]):
|
||||||
"""Handle websocket message receiving with automatic retry logic.
|
"""Handle websocket message receiving with automatic retry logic.
|
||||||
|
|
||||||
@@ -76,13 +130,9 @@ class WebsocketService(ABC):
|
|||||||
Args:
|
Args:
|
||||||
report_error: Callback function to report connection errors.
|
report_error: Callback function to report connection errors.
|
||||||
"""
|
"""
|
||||||
retry_count = 0
|
|
||||||
MAX_RETRIES = 3
|
|
||||||
|
|
||||||
while True:
|
while True:
|
||||||
try:
|
try:
|
||||||
await self._receive_messages()
|
await self._receive_messages()
|
||||||
retry_count = 0 # Reset counter on successful message receive
|
|
||||||
except ConnectionClosedOK as e:
|
except ConnectionClosedOK as e:
|
||||||
# Normal closure, don't retry
|
# Normal closure, don't retry
|
||||||
logger.debug(f"{self} connection closed normally: {e}")
|
logger.debug(f"{self} connection closed normally: {e}")
|
||||||
@@ -92,21 +142,9 @@ class WebsocketService(ABC):
|
|||||||
logger.error(message)
|
logger.error(message)
|
||||||
|
|
||||||
if self._reconnect_on_error:
|
if self._reconnect_on_error:
|
||||||
retry_count += 1
|
success = await self._try_reconnect(report_error=report_error)
|
||||||
if retry_count >= MAX_RETRIES:
|
if not success:
|
||||||
await report_error(ErrorFrame(message))
|
|
||||||
break
|
break
|
||||||
|
|
||||||
logger.warning(f"{self} connection error, will retry: {e}")
|
|
||||||
await report_error(ErrorFrame(message))
|
|
||||||
|
|
||||||
try:
|
|
||||||
if await self._reconnect_websocket(retry_count):
|
|
||||||
retry_count = 0 # Reset counter on successful reconnection
|
|
||||||
wait_time = exponential_backoff_time(retry_count)
|
|
||||||
await asyncio.sleep(wait_time)
|
|
||||||
except Exception as reconnect_error:
|
|
||||||
logger.error(f"{self} reconnection failed: {reconnect_error}")
|
|
||||||
else:
|
else:
|
||||||
await report_error(ErrorFrame(message))
|
await report_error(ErrorFrame(message))
|
||||||
break
|
break
|
||||||
|
|||||||
Reference in New Issue
Block a user