Skip to content

Commit 2aeea01

Browse files
committed
Refresh relay connection token before it expires
The relay connection token expires every 24 hours for MI tokens and has a predefined expiry for SAS tokens. By monitoring the time the connection has been open in an outer loop we can force the refresh of the token before it expires.
1 parent 0fb0417 commit 2aeea01

3 files changed

Lines changed: 278 additions & 79 deletions

File tree

src/relay_listener.py

Lines changed: 129 additions & 56 deletions
Original file line numberDiff line numberDiff line change
@@ -38,6 +38,7 @@
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):
@@ -48,15 +49,21 @@ class RelayListener:
4849
"""
4950
Socket Listener for Azure Relay.
5051
51-
Listens for incoming messages from Azure Relay and processes worklist actions.
52+
Listens for incoming messages from Azure Relay and processes worklist
53+
actions.
54+
5255
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)
56+
AZURE_RELAY_NAMESPACE: Azure Relay namespace
57+
(default: relay-test.servicebus.windows.net)
58+
AZURE_RELAY_HYBRID_CONNECTION: Azure Relay hybrid connection name
59+
(default: relay-test-hc)
60+
MWL_DB_PATH: Path to the MWL SQLite database file
61+
(default: /var/lib/pacs/worklist.db)
5662
5763
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
64+
AZURE_RELAY_KEY_NAME: Shared access policy name
65+
(default: RootManageSharedAccessKey)
66+
AZURE_RELAY_SHARED_ACCESS_KEY: Shared access key value
6067
"""
6168

6269
def __init__(self, storage: MWLStorage):
@@ -66,64 +73,91 @@ def __init__(self, storage: MWLStorage):
6673
async def listen(self):
6774
"""Listen for messages from Azure Relay."""
6875

69-
logger.info(f"Connecting to Azure Relay: {self.relay_uri.hybrid_connection_name}...")
76+
logger.info(
77+
"Connecting to Azure Relay: %s...",
78+
self.relay_uri.hybrid_connection_name,
79+
)
7080

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

74-
async for message in websocket:
75-
try:
76-
data = json.loads(message)
85+
try:
86+
async with self._connect(connection_url) as websocket:
87+
logger.info("Connected - waiting for worklist actions...")
88+
await self._listen_on_connection(websocket, refresh_at)
89+
except asyncio.TimeoutError:
90+
logger.info("Refreshing Azure Relay connection before expiry.")
91+
continue
92+
93+
async def _listen_on_connection(self, websocket, refresh_at: int):
94+
while True:
95+
timeout = refresh_at - time.time()
96+
if timeout <= 0:
97+
raise asyncio.TimeoutError
7798

78-
if "accept" in data:
79-
accept_url = data["accept"]["address"]
80-
logger.info("Incoming connection...")
99+
try:
100+
message = await asyncio.wait_for(websocket.recv(), timeout=timeout)
101+
except asyncio.TimeoutError:
102+
raise
81103

82-
async with connect(accept_url, compression=None) as client_ws:
83-
client_message = await asyncio.wait_for(client_ws.recv(), timeout=30)
104+
try:
105+
data = json.loads(message)
106+
107+
if "accept" in data:
108+
accept_url = data["accept"]["address"]
109+
logger.info("Incoming connection...")
110+
111+
async with connect(accept_url, compression=None) as client_ws:
112+
try:
113+
client_message = await asyncio.wait_for(
114+
client_ws.recv(),
115+
timeout=30,
116+
)
84117
payload = json.loads(client_message)
85118
response = self.process_action(payload)
86119

87-
# Send acknowledgment
88120
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}")
121+
except asyncio.TimeoutError:
122+
logger.error("Timeout waiting for message")
123+
except Exception:
124+
logger.exception("Error processing relay message")
94125

95126
def process_action(self, payload: dict):
96127
"""Process incoming action payload."""
97128
action_name = payload.get("action_type", "no-op")
98129

99130
if action_name == "echo":
100131
return {"status": "echo", "payload": payload}
101-
elif action_name == "worklist.create_item":
132+
if action_name == "worklist.create_item":
102133
return CreateWorklistItem(self.storage).call(payload)
103-
elif action_name == "worklist.create_test_item":
134+
if action_name == "worklist.create_test_item":
104135
result = CreateWorklistItem(self.storage).call(payload)
105-
patient_name = payload.get("parameters", {}).get("worklist_item", {}).get("participant", {}).get("name")
136+
137+
worklist_item = payload.get("parameters", {}).get("worklist_item", {})
138+
participant = worklist_item.get("participant", {})
139+
patient_name = participant.get("name")
106140

107141
if not patient_name:
108142
logger.warning("No patient name provided for ModalityEmulator test item processing")
109143
return {
110144
"status": "error",
111-
"message": "No patient name provided for ModalityEmulator test item processing",
145+
"message": ("No patient name provided for ModalityEmulator test item processing"),
112146
}
113147

114148
self.process_with_modality_emulator(patient_name=patient_name)
115149

116150
return result
117-
elif action_name == "worklist.update_status":
151+
if action_name == "worklist.update_status":
118152
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}"}
122153

123-
def _connect(self):
154+
logger.error("Unsupported action: %s", action_name)
155+
return {"status": "error", "message": f"Unsupported action: {action_name}"}
156+
157+
def _connect(self, connection_url: str):
124158
"""Connect to Azure Relay."""
125159
return connect(
126-
self.relay_uri.connection_url(),
160+
connection_url,
127161
compression=None,
128162
)
129163

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

136170
def _run_emulator():
137171
try:
138-
ModalityEmulator(self.storage).process_worklist_items(ae, patient_name=patient_name)
172+
ModalityEmulator(self.storage).process_worklist_items(
173+
ae,
174+
patient_name=patient_name,
175+
)
139176
except Exception:
140177
logger.exception("Modality emulator processing failed")
141178

@@ -150,12 +187,25 @@ def _run_emulator():
150187

151188
class RelayURI:
152189
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")
190+
self.relay_namespace = os.getenv(
191+
"AZURE_RELAY_NAMESPACE",
192+
"relay-test.servicebus.windows.net",
193+
)
194+
self.hybrid_connection_name = os.getenv(
195+
"AZURE_RELAY_HYBRID_CONNECTION",
196+
"relay-test-hc",
197+
)
198+
self.key_name = os.getenv(
199+
"AZURE_RELAY_KEY_NAME",
200+
"RootManageSharedAccessKey",
201+
)
156202
self.shared_access_key = os.getenv("AZURE_RELAY_SHARED_ACCESS_KEY", "")
157203
self._env = Environment()
158-
self._credential = None if self._use_sas() else self._build_credential()
204+
205+
if self._use_sas():
206+
self._credential = None
207+
else:
208+
self._credential = self._build_credential()
159209

160210
def _use_sas(self) -> bool:
161211
return not self._env.production and bool(self.shared_access_key)
@@ -165,41 +215,61 @@ def _build_credential(self):
165215
return ManagedIdentityCredential()
166216
return DefaultAzureCredential()
167217

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

176-
def _create_bearer_token(self) -> str:
226+
connection_url = f"{base}&sb-hc-token={urllib.parse.quote_plus(token)}"
227+
return connection_url, expires_on
228+
229+
def connection_url(self) -> str:
230+
connection_url, _ = self.connection_details()
231+
return connection_url
232+
233+
def _create_bearer_token(self) -> tuple[str, int]:
177234
if self._credential is None:
178235
raise CredentialNotAvailableError(
179236
"No credential available — _credential should never be None when not using SAS"
180237
)
181-
return f"Bearer {self._credential.get_token(AZURE_RELAY_SCOPE).token}"
182238

183-
def _create_sas_token(self, expiry_seconds: int = SAS_TOKEN_EXPIRY_SECONDS) -> str:
239+
access_token = self._credential.get_token(AZURE_RELAY_SCOPE)
240+
return f"Bearer {access_token.token}", int(access_token.expires_on)
241+
242+
def _create_sas_token(
243+
self,
244+
expiry_seconds: int = SAS_TOKEN_EXPIRY_SECONDS,
245+
) -> tuple[str, int]:
184246
uri = f"http://{self.relay_namespace}/{self.hybrid_connection_name}"
185247
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-
)
248+
expiry = int(time.time() + expiry_seconds)
249+
string_to_sign = f"{encoded_uri}\n{expiry}".encode()
250+
digest = hmac.new(
251+
self.shared_access_key.encode(),
252+
string_to_sign,
253+
hashlib.sha256,
254+
).digest()
255+
signature = base64.b64encode(digest).decode("ascii")
256+
190257
return (
191258
f"SharedAccessSignature sr={encoded_uri}"
192259
f"&sig={urllib.parse.quote_plus(signature)}"
193-
f"&se={expiry}&skn={self.key_name}"
260+
f"&se={expiry}&skn={self.key_name}",
261+
expiry,
194262
)
195263

196264

197265
def verify_credentials():
198266
"""
199267
Verify relay credentials are available at startup.
200268
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.
269+
In production, raises ClientAuthenticationError if managed identity is
270+
not configured.
271+
In non-production with a SAS key present, logs the auth method and
272+
returns immediately.
203273
"""
204274
uri = RelayURI()
205275
if uri._use_sas():
@@ -211,13 +281,16 @@ def verify_credentials():
211281
)
212282
uri._credential.get_token(AZURE_RELAY_SCOPE)
213283
credential_type = "ManagedIdentityCredential" if uri._env.production else "DefaultAzureCredential"
214-
logger.info(f"Azure Relay credentials verified ({credential_type}).")
284+
logger.info("Azure Relay credentials verified (%s).", credential_type)
215285

216286

217287
async def main():
218288
logging.basicConfig(
219289
level=os.getenv("LOG_LEVEL", "INFO").upper(),
220-
format=os.getenv("LOG_FORMAT", "%(asctime)s - %(name)s - %(levelname)s - %(message)s"),
290+
format=os.getenv(
291+
"LOG_FORMAT",
292+
"%(asctime)s - %(name)s - %(levelname)s - %(message)s",
293+
),
221294
)
222295
configure_telemetry(service_name="relay-listener")
223296

@@ -234,11 +307,11 @@ async def main():
234307
except ConnectionClosedError as e:
235308
code = e.rcvd.code if e.rcvd else "N/A"
236309
reason = e.rcvd.reason if e.rcvd else "N/A"
237-
logger.warning(f"Connection closed with code {code}: {reason}")
310+
logger.warning("Connection closed with code %s: %s", code, reason)
238311
logger.warning("Retrying in 5 seconds...")
239312
await asyncio.sleep(5)
240313
except Exception as e:
241-
logger.warning(f"Connection error: {e}")
314+
logger.warning("Connection error: %s", e)
242315
logger.warning("Retrying in 5 seconds...")
243316
await asyncio.sleep(5)
244317

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)