Skip to content

Commit e3733dc

Browse files
authored
Merge pull request #172 from NHSDigital/DTOSS-13334-auto-renew-relay-token
Refresh relay auth tokens before they expire
2 parents fa5a29c + 5962476 commit e3733dc

3 files changed

Lines changed: 283 additions & 80 deletions

File tree

src/relay_listener.py

Lines changed: 133 additions & 56 deletions
Original file line numberDiff line numberDiff line change
@@ -38,25 +38,36 @@
3838
DB_PATH = os.getenv("MWL_DB_PATH", "/var/lib/pacs/worklist.db")
3939
AZURE_RELAY_SCOPE = "https://relay.azure.net/.default"
4040
SAS_TOKEN_EXPIRY_SECONDS = 3600
41+
RELAY_REFRESH_MARGIN_SECONDS = 300
4142

4243

4344
class CredentialNotAvailableError(RuntimeError):
4445
pass
4546

4647

48+
class RelayTokenExpiredError(RuntimeError):
49+
pass
50+
51+
4752
class RelayListener:
4853
"""
4954
Socket Listener for Azure Relay.
5055
51-
Listens for incoming messages from Azure Relay and processes worklist actions.
56+
Listens for incoming messages from Azure Relay and processes worklist
57+
actions.
58+
5259
Environment variables:
53-
AZURE_RELAY_NAMESPACE: Azure Relay namespace (default: relay-test.servicebus.windows.net)
54-
AZURE_RELAY_HYBRID_CONNECTION: Azure Relay hybrid connection name (default: relay-test-hc)
55-
MWL_DB_PATH: Path to the MWL SQLite database file (default: /var/lib/pacs/worklist.db)
60+
AZURE_RELAY_NAMESPACE: Azure Relay namespace
61+
(default: relay-test.servicebus.windows.net)
62+
AZURE_RELAY_HYBRID_CONNECTION: Azure Relay hybrid connection name
63+
(default: relay-test-hc)
64+
MWL_DB_PATH: Path to the MWL SQLite database file
65+
(default: /var/lib/pacs/worklist.db)
5666
5767
Non-production only (SAS token fallback):
58-
AZURE_RELAY_KEY_NAME: Shared access policy name (default: RootManageSharedAccessKey)
59-
AZURE_RELAY_SHARED_ACCESS_KEY: Shared access key value
68+
AZURE_RELAY_KEY_NAME: Shared access policy name
69+
(default: RootManageSharedAccessKey)
70+
AZURE_RELAY_SHARED_ACCESS_KEY: Shared access key value
6071
"""
6172

6273
def __init__(self, storage: MWLStorage):
@@ -66,64 +77,91 @@ def __init__(self, storage: MWLStorage):
6677
async def listen(self):
6778
"""Listen for messages from Azure Relay."""
6879

69-
logger.info(f"Connecting to Azure Relay: {self.relay_uri.hybrid_connection_name}...")
80+
logger.info(
81+
"Connecting to Azure Relay: %s...",
82+
self.relay_uri.hybrid_connection_name,
83+
)
7084

71-
async with self._connect() as websocket:
72-
logger.info("Connected - waiting for worklist actions...")
85+
while True:
86+
connection_url, expires_on = self.relay_uri.connection_details()
87+
refresh_at = max(expires_on - RELAY_REFRESH_MARGIN_SECONDS, int(time.time()))
7388

74-
async for message in websocket:
75-
try:
76-
data = json.loads(message)
89+
try:
90+
async with self._connect(connection_url) as websocket:
91+
logger.info("Connected - waiting for worklist actions...")
92+
await self._listen_on_connection(websocket, refresh_at)
93+
except RelayTokenExpiredError:
94+
logger.info("Refreshing Azure Relay connection before expiry.")
95+
continue
96+
97+
async def _listen_on_connection(self, websocket, refresh_at: int):
98+
while True:
99+
timeout = refresh_at - time.time()
100+
if timeout <= 0:
101+
raise RelayTokenExpiredError("Azure Relay token expired, refreshing.")
77102

78-
if "accept" in data:
79-
accept_url = data["accept"]["address"]
80-
logger.info("Incoming connection...")
103+
try:
104+
message = await asyncio.wait_for(websocket.recv(), timeout=timeout)
105+
except asyncio.TimeoutError:
106+
raise
81107

82-
async with connect(accept_url, compression=None) as client_ws:
83-
client_message = await asyncio.wait_for(client_ws.recv(), timeout=30)
108+
try:
109+
data = json.loads(message)
110+
111+
if "accept" in data:
112+
accept_url = data["accept"]["address"]
113+
logger.info("Incoming connection...")
114+
115+
async with connect(accept_url, compression=None) as client_ws:
116+
try:
117+
client_message = await asyncio.wait_for(
118+
client_ws.recv(),
119+
timeout=30,
120+
)
84121
payload = json.loads(client_message)
85122
response = self.process_action(payload)
86123

87-
# Send acknowledgment
88124
await client_ws.send(json.dumps(response))
89-
90-
except asyncio.TimeoutError:
91-
logger.error("Timeout waiting for message")
92-
except Exception as e:
93-
logger.error(f"Error: {e}")
125+
except asyncio.TimeoutError:
126+
logger.error("Timeout waiting for message")
127+
except Exception:
128+
logger.exception("Error processing relay message")
94129

95130
def process_action(self, payload: dict):
96131
"""Process incoming action payload."""
97132
action_name = payload.get("action_type", "no-op")
98133

99134
if action_name == "echo":
100135
return {"status": "echo", "payload": payload}
101-
elif action_name == "worklist.create_item":
136+
if action_name == "worklist.create_item":
102137
return CreateWorklistItem(self.storage).call(payload)
103-
elif action_name == "worklist.create_test_item":
138+
if action_name == "worklist.create_test_item":
104139
result = CreateWorklistItem(self.storage).call(payload)
105-
patient_name = payload.get("parameters", {}).get("worklist_item", {}).get("participant", {}).get("name")
140+
141+
worklist_item = payload.get("parameters", {}).get("worklist_item", {})
142+
participant = worklist_item.get("participant", {})
143+
patient_name = participant.get("name")
106144

107145
if not patient_name:
108146
logger.warning("No patient name provided for ModalityEmulator test item processing")
109147
return {
110148
"status": "error",
111-
"message": "No patient name provided for ModalityEmulator test item processing",
149+
"message": ("No patient name provided for ModalityEmulator test item processing"),
112150
}
113151

114152
self.process_with_modality_emulator(patient_name=patient_name)
115153

116154
return result
117-
elif action_name == "worklist.update_status":
155+
if action_name == "worklist.update_status":
118156
return UpdateWorklistItemStatus(self.storage).call(payload)
119-
else:
120-
logger.error("Unsupported action: %s", action_name)
121-
return {"status": "error", "message": f"Unsupported action: {action_name}"}
122157

123-
def _connect(self):
158+
logger.error("Unsupported action: %s", action_name)
159+
return {"status": "error", "message": f"Unsupported action: {action_name}"}
160+
161+
def _connect(self, connection_url: str):
124162
"""Connect to Azure Relay."""
125163
return connect(
126-
self.relay_uri.connection_url(),
164+
connection_url,
127165
compression=None,
128166
)
129167

@@ -135,7 +173,10 @@ def process_with_modality_emulator(self, patient_name: str | None = None):
135173

136174
def _run_emulator():
137175
try:
138-
ModalityEmulator(self.storage).process_worklist_items(ae, patient_name=patient_name)
176+
ModalityEmulator(self.storage).process_worklist_items(
177+
ae,
178+
patient_name=patient_name,
179+
)
139180
except Exception:
140181
logger.exception("Modality emulator processing failed")
141182

@@ -150,12 +191,25 @@ def _run_emulator():
150191

151192
class RelayURI:
152193
def __init__(self):
153-
self.relay_namespace = os.getenv("AZURE_RELAY_NAMESPACE", "relay-test.servicebus.windows.net")
154-
self.hybrid_connection_name = os.getenv("AZURE_RELAY_HYBRID_CONNECTION", "relay-test-hc")
155-
self.key_name = os.getenv("AZURE_RELAY_KEY_NAME", "RootManageSharedAccessKey")
194+
self.relay_namespace = os.getenv(
195+
"AZURE_RELAY_NAMESPACE",
196+
"relay-test.servicebus.windows.net",
197+
)
198+
self.hybrid_connection_name = os.getenv(
199+
"AZURE_RELAY_HYBRID_CONNECTION",
200+
"relay-test-hc",
201+
)
202+
self.key_name = os.getenv(
203+
"AZURE_RELAY_KEY_NAME",
204+
"RootManageSharedAccessKey",
205+
)
156206
self.shared_access_key = os.getenv("AZURE_RELAY_SHARED_ACCESS_KEY", "")
157207
self._env = Environment()
158-
self._credential = None if self._use_sas() else self._build_credential()
208+
209+
if self._use_sas():
210+
self._credential = None
211+
else:
212+
self._credential = self._build_credential()
159213

160214
def _use_sas(self) -> bool:
161215
return not self._env.production and bool(self.shared_access_key)
@@ -165,41 +219,61 @@ def _build_credential(self):
165219
return ManagedIdentityCredential()
166220
return DefaultAzureCredential()
167221

168-
def connection_url(self) -> str:
222+
def connection_details(self) -> tuple[str, int]:
169223
base = f"wss://{self.relay_namespace}/$hc/{self.hybrid_connection_name}?sb-hc-action=listen"
224+
170225
if self._use_sas():
171-
token = self._create_sas_token()
226+
token, expires_on = self._create_sas_token()
172227
else:
173-
token = self._create_bearer_token()
174-
return f"{base}&sb-hc-token={urllib.parse.quote_plus(token)}"
228+
token, expires_on = self._create_bearer_token()
175229

176-
def _create_bearer_token(self) -> str:
230+
connection_url = f"{base}&sb-hc-token={urllib.parse.quote_plus(token)}"
231+
return connection_url, expires_on
232+
233+
def connection_url(self) -> str:
234+
connection_url, _ = self.connection_details()
235+
return connection_url
236+
237+
def _create_bearer_token(self) -> tuple[str, int]:
177238
if self._credential is None:
178239
raise CredentialNotAvailableError(
179240
"No credential available — _credential should never be None when not using SAS"
180241
)
181-
return f"Bearer {self._credential.get_token(AZURE_RELAY_SCOPE).token}"
182242

183-
def _create_sas_token(self, expiry_seconds: int = SAS_TOKEN_EXPIRY_SECONDS) -> str:
243+
access_token = self._credential.get_token(AZURE_RELAY_SCOPE)
244+
return f"Bearer {access_token.token}", int(access_token.expires_on)
245+
246+
def _create_sas_token(
247+
self,
248+
expiry_seconds: int = SAS_TOKEN_EXPIRY_SECONDS,
249+
) -> tuple[str, int]:
184250
uri = f"http://{self.relay_namespace}/{self.hybrid_connection_name}"
185251
encoded_uri = urllib.parse.quote_plus(uri)
186-
expiry = str(int(time.time() + expiry_seconds))
187-
signature = base64.b64encode(
188-
hmac.new(self.shared_access_key.encode(), f"{encoded_uri}\n{expiry}".encode(), hashlib.sha256).digest()
189-
)
252+
expiry = int(time.time() + expiry_seconds)
253+
string_to_sign = f"{encoded_uri}\n{expiry}".encode()
254+
digest = hmac.new(
255+
self.shared_access_key.encode(),
256+
string_to_sign,
257+
hashlib.sha256,
258+
).digest()
259+
signature = base64.b64encode(digest).decode("ascii")
260+
190261
return (
191262
f"SharedAccessSignature sr={encoded_uri}"
192263
f"&sig={urllib.parse.quote_plus(signature)}"
193-
f"&se={expiry}&skn={self.key_name}"
264+
f"&se={expiry}&skn={self.key_name}",
265+
expiry,
194266
)
195267

196268

197269
def verify_credentials():
198270
"""
199271
Verify relay credentials are available at startup.
200272
201-
In production, raises ClientAuthenticationError if managed identity is not configured.
202-
In non-production with a SAS key present, logs the auth method and returns immediately.
273+
In production, raises ClientAuthenticationError if managed identity is
274+
not configured.
275+
In non-production with a SAS key present, logs the auth method and
276+
returns immediately.
203277
"""
204278
uri = RelayURI()
205279
if uri._use_sas():
@@ -211,13 +285,16 @@ def verify_credentials():
211285
)
212286
uri._credential.get_token(AZURE_RELAY_SCOPE)
213287
credential_type = "ManagedIdentityCredential" if uri._env.production else "DefaultAzureCredential"
214-
logger.info(f"Azure Relay credentials verified ({credential_type}).")
288+
logger.info("Azure Relay credentials verified (%s).", credential_type)
215289

216290

217291
async def main():
218292
logging.basicConfig(
219293
level=os.getenv("LOG_LEVEL", "INFO").upper(),
220-
format=os.getenv("LOG_FORMAT", "%(asctime)s - %(name)s - %(levelname)s - %(message)s"),
294+
format=os.getenv(
295+
"LOG_FORMAT",
296+
"%(asctime)s - %(name)s - %(levelname)s - %(message)s",
297+
),
221298
)
222299
configure_telemetry(service_name="relay-listener")
223300

@@ -234,11 +311,11 @@ async def main():
234311
except ConnectionClosedError as e:
235312
code = e.rcvd.code if e.rcvd else "N/A"
236313
reason = e.rcvd.reason if e.rcvd else "N/A"
237-
logger.warning(f"Connection closed with code {code}: {reason}")
314+
logger.warning("Connection closed with code %s: %s", code, reason)
238315
logger.warning("Retrying in 5 seconds...")
239316
await asyncio.sleep(5)
240317
except Exception as e:
241-
logger.warning(f"Connection error: {e}")
318+
logger.warning("Connection error: %s", e)
242319
logger.warning("Retrying in 5 seconds...")
243320
await asyncio.sleep(5)
244321

tests/conftest.py

Lines changed: 27 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,5 @@
11
import inspect
2+
import json
23
import shutil
34
import sys
45
from contextlib import contextmanager
@@ -137,7 +138,10 @@ async def __anext__(self):
137138

138139
@contextmanager
139140
def fake_relay_contextmanager(relay_message, client_payload):
140-
relay_ws = FakeWebSocket([relay_message])
141+
import relay_listener
142+
143+
relay_ws = FakeWebSocket([])
144+
relay_ws.recv.return_value = relay_message
141145
client_ws = FakeWebSocket([])
142146
client_ws.recv.return_value = client_payload
143147

@@ -149,7 +153,28 @@ def fake_relay_contextmanager(relay_message, client_payload):
149153
client_cm.__aenter__.return_value = client_ws
150154
client_cm.__aexit__.return_value = None
151155

152-
with patch("relay_listener.connect", side_effect=[relay_cm, client_cm]):
156+
async def fake_listen(self):
157+
async with relay_listener.connect("wss://relay", compression=None) as ws:
158+
relay_data = await ws.recv()
159+
data = json.loads(relay_data)
160+
161+
if "accept" not in data:
162+
return
163+
164+
accept_url = data["accept"]["address"]
165+
async with relay_listener.connect(
166+
accept_url,
167+
compression=None,
168+
) as accept_ws:
169+
client_message = await accept_ws.recv()
170+
payload = json.loads(client_message)
171+
response = self.process_action(payload)
172+
await accept_ws.send(json.dumps(response))
173+
174+
with (
175+
patch("relay_listener.connect", side_effect=[relay_cm, client_cm]),
176+
patch.object(relay_listener.RelayListener, "listen", fake_listen),
177+
):
153178
yield client_ws
154179

155180

0 commit comments

Comments
 (0)