Skip to content

Commit f6daaf5

Browse files
committed
Reconnect immediately if connection fails on token expiry
1 parent 12742c7 commit f6daaf5

2 files changed

Lines changed: 38 additions & 2 deletions

File tree

src/relay_listener.py

Lines changed: 12 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -16,6 +16,7 @@
1616

1717
from dotenv import load_dotenv
1818
from websockets.asyncio.client import connect
19+
from websockets.exceptions import ConnectionClosedError
1920

2021
from services.mwl.create_worklist_item import CreateWorklistItem
2122
from services.storage import MWLStorage
@@ -54,7 +55,7 @@ async def listen(self):
5455

5556
logger.info(f"Connecting to Azure Relay: {self.relay_uri.hybrid_connection_name}...")
5657

57-
async with connect(self.relay_uri.connection_url(), compression=None) as websocket:
58+
async with self._connect() as websocket:
5859
logger.info("Connected - waiting for worklist actions...")
5960

6061
async for message in websocket:
@@ -88,6 +89,10 @@ def process_action(self, payload: dict):
8889

8990
return action_class(self.storage).call(payload)
9091

92+
def _connect(self):
93+
"""Connect to Azure Relay."""
94+
return connect(self.relay_uri.connection_url(), compression=None)
95+
9196

9297
class RelayURI:
9398
def __init__(self):
@@ -133,6 +138,12 @@ async def main():
133138
except KeyboardInterrupt:
134139
logger.warning("\nShutting down...")
135140
break
141+
except ConnectionClosedError as e:
142+
if "ExpiredToken" in str(e) and e.code == 1011:
143+
logger.info("SAS token expired, refreshing...")
144+
else:
145+
logger.warning(f"Connection closed with code {e.code}: {e.reason}")
146+
raise e
136147
except Exception as e:
137148
logger.warning(f"Connection error: {e}")
138149
logger.warning("Retrying in 5 seconds...")

tests/test_relay_listener.py

Lines changed: 26 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -3,8 +3,11 @@
33

44
import pytest
55

6-
from relay_listener import RelayListener, RelayURI
6+
from relay_listener import main, RelayListener, RelayURI
77
from services.storage import WorklistItem
8+
from unittest.mock import AsyncMock, MagicMock
9+
from websockets.exceptions import ConnectionClosedError
10+
from websockets.frames import Close, CloseCode
811

912

1013
class TestRelayListener:
@@ -121,3 +124,25 @@ def test_relay_uri_connection_url(self):
121124
"&sb-hc-token=SharedAccessSignature+sr%3Dhttp%253A%252F%252Ftest-namespace"
122125
"%252Ftest-connection%26sig%3DPMcelSnwGlYX2xFo9Y2aGCg%252BvJ6LsHujiRrA1L6VnP0%253D%26se%3D1003600%26skn%3Dtest-key-name"
123126
)
127+
128+
129+
@patch("relay_listener.logger", new_callable=MagicMock)
130+
@patch("relay_listener.MWLStorage", new_callable=MagicMock)
131+
@patch("relay_listener.RelayListener")
132+
@patch("asyncio.sleep", new_callable=AsyncMock)
133+
@pytest.mark.asyncio
134+
async def test_main_handles_connection_closed_and_keyboard_interrupt(mock_sleep, mock_relay_listener, mock_mwl_storage, mock_logger):
135+
relay_listener_instance = mock_relay_listener.return_value
136+
relay_listener_instance.listen = AsyncMock()
137+
138+
relay_listener_instance.listen.side_effect = [
139+
ConnectionClosedError(Close(CloseCode.INTERNAL_ERROR, "ExpiredToken"), None),
140+
KeyboardInterrupt()
141+
]
142+
143+
await main()
144+
145+
assert relay_listener_instance.listen.call_count == 2
146+
mock_logger.info.assert_any_call("Socket Listener Starting...")
147+
mock_logger.info.assert_any_call("SAS token expired, refreshing...")
148+
mock_logger.warning.assert_any_call("\nShutting down...")

0 commit comments

Comments
 (0)