diff --git a/src/relay_listener.py b/src/relay_listener.py index 5415d8cd..dd85959f 100644 --- a/src/relay_listener.py +++ b/src/relay_listener.py @@ -38,25 +38,36 @@ DB_PATH = os.getenv("MWL_DB_PATH", "/var/lib/pacs/worklist.db") AZURE_RELAY_SCOPE = "https://relay.azure.net/.default" SAS_TOKEN_EXPIRY_SECONDS = 3600 +RELAY_REFRESH_MARGIN_SECONDS = 300 class CredentialNotAvailableError(RuntimeError): pass +class RelayTokenExpiredError(RuntimeError): + pass + + class RelayListener: """ Socket Listener for Azure Relay. - Listens for incoming messages from Azure Relay and processes worklist actions. + Listens for incoming messages from Azure Relay and processes worklist + actions. + Environment variables: - AZURE_RELAY_NAMESPACE: Azure Relay namespace (default: relay-test.servicebus.windows.net) - AZURE_RELAY_HYBRID_CONNECTION: Azure Relay hybrid connection name (default: relay-test-hc) - MWL_DB_PATH: Path to the MWL SQLite database file (default: /var/lib/pacs/worklist.db) + AZURE_RELAY_NAMESPACE: Azure Relay namespace + (default: relay-test.servicebus.windows.net) + AZURE_RELAY_HYBRID_CONNECTION: Azure Relay hybrid connection name + (default: relay-test-hc) + MWL_DB_PATH: Path to the MWL SQLite database file + (default: /var/lib/pacs/worklist.db) Non-production only (SAS token fallback): - AZURE_RELAY_KEY_NAME: Shared access policy name (default: RootManageSharedAccessKey) - AZURE_RELAY_SHARED_ACCESS_KEY: Shared access key value + AZURE_RELAY_KEY_NAME: Shared access policy name + (default: RootManageSharedAccessKey) + AZURE_RELAY_SHARED_ACCESS_KEY: Shared access key value """ def __init__(self, storage: MWLStorage): @@ -66,31 +77,55 @@ def __init__(self, storage: MWLStorage): async def listen(self): """Listen for messages from Azure Relay.""" - logger.info(f"Connecting to Azure Relay: {self.relay_uri.hybrid_connection_name}...") + logger.info( + "Connecting to Azure Relay: %s...", + self.relay_uri.hybrid_connection_name, + ) - async with self._connect() as websocket: - logger.info("Connected - waiting for worklist actions...") + while True: + connection_url, expires_on = self.relay_uri.connection_details() + refresh_at = max(expires_on - RELAY_REFRESH_MARGIN_SECONDS, int(time.time())) - async for message in websocket: - try: - data = json.loads(message) + try: + async with self._connect(connection_url) as websocket: + logger.info("Connected - waiting for worklist actions...") + await self._listen_on_connection(websocket, refresh_at) + except RelayTokenExpiredError: + logger.info("Refreshing Azure Relay connection before expiry.") + continue + + async def _listen_on_connection(self, websocket, refresh_at: int): + while True: + timeout = refresh_at - time.time() + if timeout <= 0: + raise RelayTokenExpiredError("Azure Relay token expired, refreshing.") - if "accept" in data: - accept_url = data["accept"]["address"] - logger.info("Incoming connection...") + try: + message = await asyncio.wait_for(websocket.recv(), timeout=timeout) + except asyncio.TimeoutError: + raise - async with connect(accept_url, compression=None) as client_ws: - client_message = await asyncio.wait_for(client_ws.recv(), timeout=30) + try: + data = json.loads(message) + + if "accept" in data: + accept_url = data["accept"]["address"] + logger.info("Incoming connection...") + + async with connect(accept_url, compression=None) as client_ws: + try: + client_message = await asyncio.wait_for( + client_ws.recv(), + timeout=30, + ) payload = json.loads(client_message) response = self.process_action(payload) - # Send acknowledgment await client_ws.send(json.dumps(response)) - - except asyncio.TimeoutError: - logger.error("Timeout waiting for message") - except Exception as e: - logger.error(f"Error: {e}") + except asyncio.TimeoutError: + logger.error("Timeout waiting for message") + except Exception: + logger.exception("Error processing relay message") def process_action(self, payload: dict): """Process incoming action payload.""" @@ -98,32 +133,35 @@ def process_action(self, payload: dict): if action_name == "echo": return {"status": "echo", "payload": payload} - elif action_name == "worklist.create_item": + if action_name == "worklist.create_item": return CreateWorklistItem(self.storage).call(payload) - elif action_name == "worklist.create_test_item": + if action_name == "worklist.create_test_item": result = CreateWorklistItem(self.storage).call(payload) - patient_name = payload.get("parameters", {}).get("worklist_item", {}).get("participant", {}).get("name") + + worklist_item = payload.get("parameters", {}).get("worklist_item", {}) + participant = worklist_item.get("participant", {}) + patient_name = participant.get("name") if not patient_name: logger.warning("No patient name provided for ModalityEmulator test item processing") return { "status": "error", - "message": "No patient name provided for ModalityEmulator test item processing", + "message": ("No patient name provided for ModalityEmulator test item processing"), } self.process_with_modality_emulator(patient_name=patient_name) return result - elif action_name == "worklist.update_status": + if action_name == "worklist.update_status": return UpdateWorklistItemStatus(self.storage).call(payload) - else: - logger.error("Unsupported action: %s", action_name) - return {"status": "error", "message": f"Unsupported action: {action_name}"} - def _connect(self): + logger.error("Unsupported action: %s", action_name) + return {"status": "error", "message": f"Unsupported action: {action_name}"} + + def _connect(self, connection_url: str): """Connect to Azure Relay.""" return connect( - self.relay_uri.connection_url(), + connection_url, compression=None, ) @@ -135,7 +173,10 @@ def process_with_modality_emulator(self, patient_name: str | None = None): def _run_emulator(): try: - ModalityEmulator(self.storage).process_worklist_items(ae, patient_name=patient_name) + ModalityEmulator(self.storage).process_worklist_items( + ae, + patient_name=patient_name, + ) except Exception: logger.exception("Modality emulator processing failed") @@ -150,12 +191,25 @@ def _run_emulator(): class RelayURI: def __init__(self): - self.relay_namespace = os.getenv("AZURE_RELAY_NAMESPACE", "relay-test.servicebus.windows.net") - self.hybrid_connection_name = os.getenv("AZURE_RELAY_HYBRID_CONNECTION", "relay-test-hc") - self.key_name = os.getenv("AZURE_RELAY_KEY_NAME", "RootManageSharedAccessKey") + self.relay_namespace = os.getenv( + "AZURE_RELAY_NAMESPACE", + "relay-test.servicebus.windows.net", + ) + self.hybrid_connection_name = os.getenv( + "AZURE_RELAY_HYBRID_CONNECTION", + "relay-test-hc", + ) + self.key_name = os.getenv( + "AZURE_RELAY_KEY_NAME", + "RootManageSharedAccessKey", + ) self.shared_access_key = os.getenv("AZURE_RELAY_SHARED_ACCESS_KEY", "") self._env = Environment() - self._credential = None if self._use_sas() else self._build_credential() + + if self._use_sas(): + self._credential = None + else: + self._credential = self._build_credential() def _use_sas(self) -> bool: return not self._env.production and bool(self.shared_access_key) @@ -165,32 +219,50 @@ def _build_credential(self): return ManagedIdentityCredential() return DefaultAzureCredential() - def connection_url(self) -> str: + def connection_details(self) -> tuple[str, int]: base = f"wss://{self.relay_namespace}/$hc/{self.hybrid_connection_name}?sb-hc-action=listen" + if self._use_sas(): - token = self._create_sas_token() + token, expires_on = self._create_sas_token() else: - token = self._create_bearer_token() - return f"{base}&sb-hc-token={urllib.parse.quote_plus(token)}" + token, expires_on = self._create_bearer_token() - def _create_bearer_token(self) -> str: + connection_url = f"{base}&sb-hc-token={urllib.parse.quote_plus(token)}" + return connection_url, expires_on + + def connection_url(self) -> str: + connection_url, _ = self.connection_details() + return connection_url + + def _create_bearer_token(self) -> tuple[str, int]: if self._credential is None: raise CredentialNotAvailableError( "No credential available — _credential should never be None when not using SAS" ) - return f"Bearer {self._credential.get_token(AZURE_RELAY_SCOPE).token}" - def _create_sas_token(self, expiry_seconds: int = SAS_TOKEN_EXPIRY_SECONDS) -> str: + access_token = self._credential.get_token(AZURE_RELAY_SCOPE) + return f"Bearer {access_token.token}", int(access_token.expires_on) + + def _create_sas_token( + self, + expiry_seconds: int = SAS_TOKEN_EXPIRY_SECONDS, + ) -> tuple[str, int]: uri = f"http://{self.relay_namespace}/{self.hybrid_connection_name}" encoded_uri = urllib.parse.quote_plus(uri) - expiry = str(int(time.time() + expiry_seconds)) - signature = base64.b64encode( - hmac.new(self.shared_access_key.encode(), f"{encoded_uri}\n{expiry}".encode(), hashlib.sha256).digest() - ) + expiry = int(time.time() + expiry_seconds) + string_to_sign = f"{encoded_uri}\n{expiry}".encode() + digest = hmac.new( + self.shared_access_key.encode(), + string_to_sign, + hashlib.sha256, + ).digest() + signature = base64.b64encode(digest).decode("ascii") + return ( f"SharedAccessSignature sr={encoded_uri}" f"&sig={urllib.parse.quote_plus(signature)}" - f"&se={expiry}&skn={self.key_name}" + f"&se={expiry}&skn={self.key_name}", + expiry, ) @@ -198,8 +270,10 @@ def verify_credentials(): """ Verify relay credentials are available at startup. - In production, raises ClientAuthenticationError if managed identity is not configured. - In non-production with a SAS key present, logs the auth method and returns immediately. + In production, raises ClientAuthenticationError if managed identity is + not configured. + In non-production with a SAS key present, logs the auth method and + returns immediately. """ uri = RelayURI() if uri._use_sas(): @@ -211,13 +285,16 @@ def verify_credentials(): ) uri._credential.get_token(AZURE_RELAY_SCOPE) credential_type = "ManagedIdentityCredential" if uri._env.production else "DefaultAzureCredential" - logger.info(f"Azure Relay credentials verified ({credential_type}).") + logger.info("Azure Relay credentials verified (%s).", credential_type) async def main(): logging.basicConfig( level=os.getenv("LOG_LEVEL", "INFO").upper(), - format=os.getenv("LOG_FORMAT", "%(asctime)s - %(name)s - %(levelname)s - %(message)s"), + format=os.getenv( + "LOG_FORMAT", + "%(asctime)s - %(name)s - %(levelname)s - %(message)s", + ), ) configure_telemetry(service_name="relay-listener") @@ -234,11 +311,11 @@ async def main(): except ConnectionClosedError as e: code = e.rcvd.code if e.rcvd else "N/A" reason = e.rcvd.reason if e.rcvd else "N/A" - logger.warning(f"Connection closed with code {code}: {reason}") + logger.warning("Connection closed with code %s: %s", code, reason) logger.warning("Retrying in 5 seconds...") await asyncio.sleep(5) except Exception as e: - logger.warning(f"Connection error: {e}") + logger.warning("Connection error: %s", e) logger.warning("Retrying in 5 seconds...") await asyncio.sleep(5) diff --git a/tests/conftest.py b/tests/conftest.py index 6f311737..dc3fbb83 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -1,4 +1,5 @@ import inspect +import json import shutil import sys from contextlib import contextmanager @@ -137,7 +138,10 @@ async def __anext__(self): @contextmanager def fake_relay_contextmanager(relay_message, client_payload): - relay_ws = FakeWebSocket([relay_message]) + import relay_listener + + relay_ws = FakeWebSocket([]) + relay_ws.recv.return_value = relay_message client_ws = FakeWebSocket([]) client_ws.recv.return_value = client_payload @@ -149,7 +153,28 @@ def fake_relay_contextmanager(relay_message, client_payload): client_cm.__aenter__.return_value = client_ws client_cm.__aexit__.return_value = None - with patch("relay_listener.connect", side_effect=[relay_cm, client_cm]): + async def fake_listen(self): + async with relay_listener.connect("wss://relay", compression=None) as ws: + relay_data = await ws.recv() + data = json.loads(relay_data) + + if "accept" not in data: + return + + accept_url = data["accept"]["address"] + async with relay_listener.connect( + accept_url, + compression=None, + ) as accept_ws: + client_message = await accept_ws.recv() + payload = json.loads(client_message) + response = self.process_action(payload) + await accept_ws.send(json.dumps(response)) + + with ( + patch("relay_listener.connect", side_effect=[relay_cm, client_cm]), + patch.object(relay_listener.RelayListener, "listen", fake_listen), + ): yield client_ws diff --git a/tests/test_relay_listener.py b/tests/test_relay_listener.py index 0e4ad64c..850b824b 100644 --- a/tests/test_relay_listener.py +++ b/tests/test_relay_listener.py @@ -1,3 +1,4 @@ +import asyncio import json from unittest.mock import AsyncMock, MagicMock, patch @@ -7,7 +8,7 @@ from websockets.frames import Close, CloseCode from models import WorklistItem -from relay_listener import RelayListener, RelayURI, main, verify_credentials +from relay_listener import RelayListener, RelayTokenExpiredError, RelayURI, main, verify_credentials class TestRelayListener: @@ -33,31 +34,74 @@ def test_relay_listener_initialization(self, storage_instance): assert subject.relay_uri.hybrid_connection_name == "test-connection" @pytest.mark.asyncio - async def test_relay_listener_listen_echo(self, storage_instance, fake_relay): - """Relay listener listen echo.""" + async def test_listen_on_connection_echo(self, storage_instance): + """Handle echo messages on a relay connection.""" subject = RelayListener(storage_instance) - - relay_message = json.dumps({"accept": {"address": "wss://accept-url"}}) - client_payload = json.dumps({"action_type": "echo", "message": "Hello, Relay!"}) - - with fake_relay(relay_message, client_payload) as client_ws: - await subject.listen() - + websocket = AsyncMock() + websocket.recv.side_effect = [ + json.dumps({"accept": {"address": "wss://accept-url"}}), + RelayTokenExpiredError("Azure Relay token expired"), + ] + + client_ws = AsyncMock() + client_ws.recv.return_value = json.dumps({"action_type": "echo", "message": "Hello, Relay!"}) + + client_cm = AsyncMock() + client_cm.__aenter__.return_value = client_ws + client_cm.__aexit__.return_value = None + + with patch("relay_listener.connect", return_value=client_cm) as mock_connect: + with pytest.raises(RelayTokenExpiredError): + await subject._listen_on_connection( + websocket, + refresh_at=9999999999, + ) + + mock_connect.assert_called_once_with( + "wss://accept-url", + compression=None, + ) client_ws.send.assert_called_once_with( - json.dumps({"status": "echo", "payload": {"action_type": "echo", "message": "Hello, Relay!"}}) + json.dumps({ + "status": "echo", + "payload": { + "action_type": "echo", + "message": "Hello, Relay!", + }, + }) ) @pytest.mark.asyncio - async def test_relay_listener_listen(self, storage_instance, listener_payload, fake_relay): - """Relay listener listen.""" - storage_instance.store_worklist_action.return_value = {"action_id": "action-12345", "status": "created"} + async def test_listen_on_connection_create_item( + self, + storage_instance, + listener_payload, + ): + """Handle create-item messages on a relay connection.""" + storage_instance.store_worklist_action.return_value = { + "action_id": "action-12345", + "status": "created", + } subject = RelayListener(storage_instance) - - relay_message = json.dumps({"accept": {"address": "wss://accept-url"}}) - client_payload = json.dumps(listener_payload) - - with fake_relay(relay_message, client_payload) as client_ws: - await subject.listen() + websocket = AsyncMock() + websocket.recv.side_effect = [ + json.dumps({"accept": {"address": "wss://accept-url"}}), + asyncio.TimeoutError(), + ] + + client_ws = AsyncMock() + client_ws.recv.return_value = json.dumps(listener_payload) + + client_cm = AsyncMock() + client_cm.__aenter__.return_value = client_ws + client_cm.__aexit__.return_value = None + + with patch("relay_listener.connect", return_value=client_cm): + with pytest.raises(asyncio.TimeoutError): + await subject._listen_on_connection( + websocket, + refresh_at=9999999999, + ) client_ws.send.assert_called_once_with(json.dumps({"status": "created", "action_id": "action-12345"})) storage_instance.store_worklist_item.assert_called_once_with( @@ -186,6 +230,45 @@ def test_process_action_invalid_type(self, storage_instance, listener_payload): storage_instance.store_worklist_item.assert_not_called() + @pytest.mark.asyncio + async def test_listen_refreshes_connection_after_timeout( + self, + storage_instance, + ): + """Listen refreshes the relay connection after a timeout.""" + subject = RelayListener(storage_instance) + + connection_cm = AsyncMock() + connection_cm.__aenter__.return_value = AsyncMock() + connection_cm.__aexit__.return_value = None + + with ( + patch.object( + subject.relay_uri, + "connection_details", + side_effect=[ + ("wss://first-url", 1), + ("wss://second-url", 2), + ], + ) as mock_details, + patch.object( + subject, + "_connect", + return_value=connection_cm, + ) as mock_connect, + patch.object( + subject, + "_listen_on_connection", + side_effect=[RelayTokenExpiredError(), KeyboardInterrupt()], + ), + ): + with pytest.raises(KeyboardInterrupt): + await subject.listen() + + assert mock_details.call_count == 2 + mock_connect.assert_any_call("wss://first-url") + mock_connect.assert_any_call("wss://second-url") + class TestRelayURIWithDefaultAzureCredential: """Non-production, no SAS key — uses DefaultAzureCredential.""" @@ -204,6 +287,14 @@ def test_connection_url(self, mock_azure_credential): assert url.startswith("wss://test-namespace/$hc/test-connection?sb-hc-action=listen") assert "sb-hc-token=Bearer+test-token" in url + def test_connection_details(self, mock_azure_credential): + """Relay URI with default azure credential: Connection details.""" + subject = RelayURI() + url, expires_on = subject.connection_details() + assert url.startswith("wss://test-namespace/$hc/test-connection?sb-hc-action=listen") + assert "sb-hc-token=Bearer+test-token" in url + assert isinstance(expires_on, int) + def test_uses_default_azure_credential(self, mock_azure_credential): """Uses default azure credential.""" with patch("relay_listener.DefaultAzureCredential") as mock_dac: @@ -230,6 +321,14 @@ def test_connection_url_includes_sas_token(self): assert url.startswith("wss://test-namespace/$hc/test-connection?sb-hc-action=listen") assert "sb-hc-token=SharedAccessSignature" in url + def test_connection_details_includes_sas_token(self): + """Connection details include SAS token.""" + subject = RelayURI() + url, expires_on = subject.connection_details() + assert url.startswith("wss://test-namespace/$hc/test-connection?sb-hc-action=listen") + assert "sb-hc-token=SharedAccessSignature" in url + assert isinstance(expires_on, int) + def test_no_credential_is_created(self): """No credential is created.""" with patch("relay_listener.DefaultAzureCredential") as mock_dac: @@ -315,7 +414,9 @@ async def test_main_handles_connection_closed_and_keyboard_interrupt( assert relay_listener_instance.listen.call_count == 3 mock_logger.info.assert_any_call("Socket Listener Starting...") - mock_logger.warning.assert_any_call("Connection closed with code 1011: Something went wrong") + mock_logger.warning.assert_any_call( + "Connection closed with code %s: %s", CloseCode.INTERNAL_ERROR, "Something went wrong" + ) mock_logger.warning.assert_any_call("Retrying in 5 seconds...") - mock_logger.warning.assert_any_call("Connection closed with code 1014: Bad gateway") + mock_logger.warning.assert_any_call("Connection closed with code %s: %s", CloseCode.BAD_GATEWAY, "Bad gateway") mock_logger.warning.assert_any_call("\nShutting down...")