diff --git a/src/relay_listener.py b/src/relay_listener.py index f52e96f6..5852a1c0 100644 --- a/src/relay_listener.py +++ b/src/relay_listener.py @@ -16,6 +16,8 @@ from dotenv import load_dotenv from websockets.asyncio.client import connect +from websockets.frames import CloseCode +from websockets.exceptions import ConnectionClosedError from services.mwl.create_worklist_item import CreateWorklistItem from services.storage import MWLStorage @@ -30,6 +32,7 @@ ACTIONS = { "worklist.create_item": CreateWorklistItem, } +EXPIRED_TOKEN = "ExpiredToken" class RelayListener: @@ -54,7 +57,7 @@ async def listen(self): logger.info(f"Connecting to Azure Relay: {self.relay_uri.hybrid_connection_name}...") - async with connect(self.relay_uri.connection_url(), compression=None) as websocket: + async with self._connect() as websocket: logger.info("Connected - waiting for worklist actions...") async for message in websocket: @@ -88,6 +91,10 @@ def process_action(self, payload: dict): return action_class(self.storage).call(payload) + def _connect(self): + """Connect to Azure Relay.""" + return connect(self.relay_uri.connection_url(), compression=None) + class RelayURI: def __init__(self): @@ -133,6 +140,16 @@ async def main(): except KeyboardInterrupt: logger.warning("\nShutting down...") break + except ConnectionClosedError as e: + code = e.rcvd.code if e.rcvd else "N/A" + reason = e.rcvd.reason if e.rcvd else "N/A" + + if code == CloseCode.INTERNAL_ERROR.value and EXPIRED_TOKEN in reason: + logger.info("SAS token expired, refreshing...") + else: + logger.warning(f"Connection closed with code {code}: {reason}") + logger.warning("Retrying in 5 seconds...") + await asyncio.sleep(5) except Exception as e: logger.warning(f"Connection error: {e}") logger.warning("Retrying in 5 seconds...") diff --git a/tests/test_relay_listener.py b/tests/test_relay_listener.py index 99d2636a..38917cd7 100644 --- a/tests/test_relay_listener.py +++ b/tests/test_relay_listener.py @@ -1,9 +1,11 @@ import json -from unittest.mock import patch +from unittest.mock import AsyncMock, MagicMock, patch import pytest +from websockets.exceptions import ConnectionClosedError +from websockets.frames import Close, CloseCode -from relay_listener import RelayListener, RelayURI +from relay_listener import RelayListener, RelayURI, main from services.storage import WorklistItem @@ -121,3 +123,33 @@ def test_relay_uri_connection_url(self): "&sb-hc-token=SharedAccessSignature+sr%3Dhttp%253A%252F%252Ftest-namespace" "%252Ftest-connection%26sig%3DPMcelSnwGlYX2xFo9Y2aGCg%252BvJ6LsHujiRrA1L6VnP0%253D%26se%3D1003600%26skn%3Dtest-key-name" ) + + +@patch("relay_listener.logger", new_callable=MagicMock) +@patch("relay_listener.MWLStorage", new_callable=MagicMock) +@patch("relay_listener.RelayListener") +@patch("asyncio.sleep", new_callable=AsyncMock) +@pytest.mark.asyncio +async def test_main_handles_connection_closed_and_keyboard_interrupt( + mock_sleep, mock_relay_listener, mock_mwl_storage, mock_logger +): + relay_listener_instance = mock_relay_listener.return_value + relay_listener_instance.listen = AsyncMock() + + relay_listener_instance.listen.side_effect = [ + ConnectionClosedError(Close(CloseCode.INTERNAL_ERROR, "ExpiredToken"), None), + ConnectionClosedError(Close(CloseCode.INTERNAL_ERROR, "Something else"), None), + ConnectionClosedError(Close(CloseCode.BAD_GATEWAY, "Bad gateway"), None), + KeyboardInterrupt(), + ] + + await main() + + assert relay_listener_instance.listen.call_count == 4 + mock_logger.info.assert_any_call("Socket Listener Starting...") + mock_logger.info.assert_any_call("SAS token expired, refreshing...") + mock_logger.warning.assert_any_call("Connection closed with code 1011: Something else") + mock_logger.warning.assert_any_call("Retrying in 5 seconds...") + mock_logger.warning.assert_any_call("Connection closed with code 1014: Bad gateway") + mock_logger.warning.assert_any_call("Retrying in 5 seconds...") + mock_logger.warning.assert_any_call("\nShutting down...")