Skip to content

Commit 99df33b

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 99df33b

3 files changed

Lines changed: 318 additions & 88 deletions

File tree

src/relay_listener.py

Lines changed: 128 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,90 @@ 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, refresh_at = self.relay_uri.connection_details()
7383

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

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

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

87-
# Send acknowledgment
88119
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}")
120+
except asyncio.TimeoutError:
121+
logger.error("Timeout waiting for message")
122+
except Exception as e:
123+
logger.error("Error: %s", e)
94124

95125
def process_action(self, payload: dict):
96126
"""Process incoming action payload."""
97127
action_name = payload.get("action_type", "no-op")
98128

99129
if action_name == "echo":
100130
return {"status": "echo", "payload": payload}
101-
elif action_name == "worklist.create_item":
131+
if action_name == "worklist.create_item":
102132
return CreateWorklistItem(self.storage).call(payload)
103-
elif action_name == "worklist.create_test_item":
133+
if action_name == "worklist.create_test_item":
104134
result = CreateWorklistItem(self.storage).call(payload)
105-
patient_name = payload.get("parameters", {}).get("worklist_item", {}).get("participant", {}).get("name")
135+
136+
worklist_item = payload.get("parameters", {}).get("worklist_item", {})
137+
participant = worklist_item.get("participant", {})
138+
patient_name = participant.get("name")
106139

107140
if not patient_name:
108141
logger.warning("No patient name provided for ModalityEmulator test item processing")
109142
return {
110143
"status": "error",
111-
"message": "No patient name provided for ModalityEmulator test item processing",
144+
"message": ("No patient name provided for ModalityEmulator test item processing"),
112145
}
113146

114147
self.process_with_modality_emulator(patient_name=patient_name)
115148

116149
return result
117-
elif action_name == "worklist.update_status":
150+
if action_name == "worklist.update_status":
118151
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}"}
122152

123-
def _connect(self):
153+
logger.error("Unsupported action: %s", action_name)
154+
return {"status": "error", "message": f"Unsupported action: {action_name}"}
155+
156+
def _connect(self, connection_url: str):
124157
"""Connect to Azure Relay."""
125158
return connect(
126-
self.relay_uri.connection_url(),
159+
connection_url,
127160
compression=None,
128161
)
129162

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

136169
def _run_emulator():
137170
try:
138-
ModalityEmulator(self.storage).process_worklist_items(ae, patient_name=patient_name)
171+
ModalityEmulator(self.storage).process_worklist_items(
172+
ae,
173+
patient_name=patient_name,
174+
)
139175
except Exception:
140176
logger.exception("Modality emulator processing failed")
141177

@@ -150,12 +186,25 @@ def _run_emulator():
150186

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

160209
def _use_sas(self) -> bool:
161210
return not self._env.production and bool(self.shared_access_key)
@@ -165,41 +214,61 @@ def _build_credential(self):
165214
return ManagedIdentityCredential()
166215
return DefaultAzureCredential()
167216

168-
def connection_url(self) -> str:
217+
def connection_details(self) -> tuple[str, int]:
169218
base = f"wss://{self.relay_namespace}/$hc/{self.hybrid_connection_name}?sb-hc-action=listen"
219+
170220
if self._use_sas():
171-
token = self._create_sas_token()
221+
token, expires_on = self._create_sas_token()
172222
else:
173-
token = self._create_bearer_token()
174-
return f"{base}&sb-hc-token={urllib.parse.quote_plus(token)}"
223+
token, expires_on = self._create_bearer_token()
224+
225+
connection_url = f"{base}&sb-hc-token={urllib.parse.quote_plus(token)}"
226+
return connection_url, expires_on
175227

176-
def _create_bearer_token(self) -> str:
228+
def connection_url(self) -> str:
229+
connection_url, _ = self.connection_details()
230+
return connection_url
231+
232+
def _create_bearer_token(self) -> tuple[str, int]:
177233
if self._credential is None:
178234
raise CredentialNotAvailableError(
179235
"No credential available — _credential should never be None when not using SAS"
180236
)
181-
return f"Bearer {self._credential.get_token(AZURE_RELAY_SCOPE).token}"
182237

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

196263

197264
def verify_credentials():
198265
"""
199266
Verify relay credentials are available at startup.
200267
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.
268+
In production, raises ClientAuthenticationError if managed identity is
269+
not configured.
270+
In non-production with a SAS key present, logs the auth method and
271+
returns immediately.
203272
"""
204273
uri = RelayURI()
205274
if uri._use_sas():
@@ -211,13 +280,16 @@ def verify_credentials():
211280
)
212281
uri._credential.get_token(AZURE_RELAY_SCOPE)
213282
credential_type = "ManagedIdentityCredential" if uri._env.production else "DefaultAzureCredential"
214-
logger.info(f"Azure Relay credentials verified ({credential_type}).")
283+
logger.info("Azure Relay credentials verified (%s).", credential_type)
215284

216285

217286
async def main():
218287
logging.basicConfig(
219288
level=os.getenv("LOG_LEVEL", "INFO").upper(),
220-
format=os.getenv("LOG_FORMAT", "%(asctime)s - %(name)s - %(levelname)s - %(message)s"),
289+
format=os.getenv(
290+
"LOG_FORMAT",
291+
"%(asctime)s - %(name)s - %(levelname)s - %(message)s",
292+
),
221293
)
222294
configure_telemetry(service_name="relay-listener")
223295

@@ -234,11 +306,11 @@ async def main():
234306
except ConnectionClosedError as e:
235307
code = e.rcvd.code if e.rcvd else "N/A"
236308
reason = e.rcvd.reason if e.rcvd else "N/A"
237-
logger.warning(f"Connection closed with code {code}: {reason}")
309+
logger.warning("Connection closed with code %s: %s", code, reason)
238310
logger.warning("Retrying in 5 seconds...")
239311
await asyncio.sleep(5)
240312
except Exception as e:
241-
logger.warning(f"Connection error: {e}")
313+
logger.warning("Connection error: %s", e)
242314
logger.warning("Retrying in 5 seconds...")
243315
await asyncio.sleep(5)
244316

0 commit comments

Comments
 (0)