Skip to content

Commit 444a620

Browse files
authored
Merge pull request #44 from NHSDigital/fix/reconnect-immediately-on-token-expiry
Reconnect immediately if connection fails on token expiry
2 parents 12742c7 + 9dc49ee commit 444a620

2 files changed

Lines changed: 52 additions & 3 deletions

File tree

src/relay_listener.py

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

1717
from dotenv import load_dotenv
1818
from websockets.asyncio.client import connect
19+
from websockets.frames import CloseCode
20+
from websockets.exceptions import ConnectionClosedError
1921

2022
from services.mwl.create_worklist_item import CreateWorklistItem
2123
from services.storage import MWLStorage
@@ -30,6 +32,7 @@
3032
ACTIONS = {
3133
"worklist.create_item": CreateWorklistItem,
3234
}
35+
EXPIRED_TOKEN = "ExpiredToken"
3336

3437

3538
class RelayListener:
@@ -54,7 +57,7 @@ async def listen(self):
5457

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

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

6063
async for message in websocket:
@@ -88,6 +91,10 @@ def process_action(self, payload: dict):
8891

8992
return action_class(self.storage).call(payload)
9093

94+
def _connect(self):
95+
"""Connect to Azure Relay."""
96+
return connect(self.relay_uri.connection_url(), compression=None)
97+
9198

9299
class RelayURI:
93100
def __init__(self):
@@ -133,6 +140,16 @@ async def main():
133140
except KeyboardInterrupt:
134141
logger.warning("\nShutting down...")
135142
break
143+
except ConnectionClosedError as e:
144+
code = e.rcvd.code if e.rcvd else "N/A"
145+
reason = e.rcvd.reason if e.rcvd else "N/A"
146+
147+
if code == CloseCode.INTERNAL_ERROR.value and EXPIRED_TOKEN in reason:
148+
logger.info("SAS token expired, refreshing...")
149+
else:
150+
logger.warning(f"Connection closed with code {code}: {reason}")
151+
logger.warning("Retrying in 5 seconds...")
152+
await asyncio.sleep(5)
136153
except Exception as e:
137154
logger.warning(f"Connection error: {e}")
138155
logger.warning("Retrying in 5 seconds...")

tests/test_relay_listener.py

Lines changed: 34 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,9 +1,11 @@
11
import json
2-
from unittest.mock import patch
2+
from unittest.mock import AsyncMock, MagicMock, patch
33

44
import pytest
5+
from websockets.exceptions import ConnectionClosedError
6+
from websockets.frames import Close, CloseCode
57

6-
from relay_listener import RelayListener, RelayURI
8+
from relay_listener import RelayListener, RelayURI, main
79
from services.storage import WorklistItem
810

911

@@ -121,3 +123,33 @@ def test_relay_uri_connection_url(self):
121123
"&sb-hc-token=SharedAccessSignature+sr%3Dhttp%253A%252F%252Ftest-namespace"
122124
"%252Ftest-connection%26sig%3DPMcelSnwGlYX2xFo9Y2aGCg%252BvJ6LsHujiRrA1L6VnP0%253D%26se%3D1003600%26skn%3Dtest-key-name"
123125
)
126+
127+
128+
@patch("relay_listener.logger", new_callable=MagicMock)
129+
@patch("relay_listener.MWLStorage", new_callable=MagicMock)
130+
@patch("relay_listener.RelayListener")
131+
@patch("asyncio.sleep", new_callable=AsyncMock)
132+
@pytest.mark.asyncio
133+
async def test_main_handles_connection_closed_and_keyboard_interrupt(
134+
mock_sleep, mock_relay_listener, mock_mwl_storage, mock_logger
135+
):
136+
relay_listener_instance = mock_relay_listener.return_value
137+
relay_listener_instance.listen = AsyncMock()
138+
139+
relay_listener_instance.listen.side_effect = [
140+
ConnectionClosedError(Close(CloseCode.INTERNAL_ERROR, "ExpiredToken"), None),
141+
ConnectionClosedError(Close(CloseCode.INTERNAL_ERROR, "Something else"), None),
142+
ConnectionClosedError(Close(CloseCode.BAD_GATEWAY, "Bad gateway"), None),
143+
KeyboardInterrupt(),
144+
]
145+
146+
await main()
147+
148+
assert relay_listener_instance.listen.call_count == 4
149+
mock_logger.info.assert_any_call("Socket Listener Starting...")
150+
mock_logger.info.assert_any_call("SAS token expired, refreshing...")
151+
mock_logger.warning.assert_any_call("Connection closed with code 1011: Something else")
152+
mock_logger.warning.assert_any_call("Retrying in 5 seconds...")
153+
mock_logger.warning.assert_any_call("Connection closed with code 1014: Bad gateway")
154+
mock_logger.warning.assert_any_call("Retrying in 5 seconds...")
155+
mock_logger.warning.assert_any_call("\nShutting down...")

0 commit comments

Comments
 (0)