From 76a1f54e00b8fb748e2fb82847acca7585633464 Mon Sep 17 00:00:00 2001 From: Carlos Martinez Date: Tue, 28 Apr 2026 16:15:19 +0100 Subject: [PATCH 1/4] Authenticate Relay with Managed Identity token Replace locally-stored SAS token in websockets URL with with Azure-generated Managed Identity JWT token in HTTP header. --- .env.development | 2 - ...DR-005_Managed_Identity_For_Azure_Relay.md | 41 ++++++++++++ pyproject.toml | 1 + src/relay_listener.py | 62 +++++++------------ tests/conftest.py | 10 ++- tests/test_relay_listener.py | 53 +++++++--------- uv.lock | 4 +- 7 files changed, 99 insertions(+), 74 deletions(-) create mode 100644 docs/adr/ADR-005_Managed_Identity_For_Azure_Relay.md diff --git a/.env.development b/.env.development index 99569602..8a790578 100644 --- a/.env.development +++ b/.env.development @@ -1,8 +1,6 @@ # Azure Relay Configuration AZURE_RELAY_NAMESPACE=manbrs-gateway-dev.servicebus.windows.net AZURE_RELAY_HYBRID_CONNECTION=name-of-your-choice-relay-test-hc -AZURE_RELAY_KEY_NAME=RootManageSharedAccessKey -AZURE_RELAY_SHARED_ACCESS_KEY=YOUR_SHARED_ACCESS_KEY_HERE CLOUD_API_ENDPOINT=https://localhost:8000/api/v1/dicom CLOUD_API_TOKEN=testtoken diff --git a/docs/adr/ADR-005_Managed_Identity_For_Azure_Relay.md b/docs/adr/ADR-005_Managed_Identity_For_Azure_Relay.md new file mode 100644 index 00000000..617abbc5 --- /dev/null +++ b/docs/adr/ADR-005_Managed_Identity_For_Azure_Relay.md @@ -0,0 +1,41 @@ +# ADR-005: Use Managed Identity for Azure Relay authentication + +Date: 2026-04-28 + +Status: Accepted + +## Context + +The gateway connects to Azure Relay Hybrid Connections to receive worklist actions from Manage Breast Screening. A connection must be authenticated to Azure Relay. + +The initial implementation used Shared Access Signature (SAS) tokens. These are HMAC-SHA256 signatures computed from a shared secret key, embedded in the WebSocket connection URL as a query parameter (`sb-hc-token`). This required: + +- A shared access key to be provisioned and stored as an environment variable (`AZURE_RELAY_SHARED_ACCESS_KEY`) +- The key name to be configured separately (`AZURE_RELAY_KEY_NAME`) +- Manual key rotation when keys needed to change + +As the gateway runs inside the hospital network but is provisioned via Azure Arc, it can be assigned a managed identity through Arc-enabled infrastructure. Storing a long-lived shared secret in the environment is therefore unnecessary operational overhead and a potential security risk. + +## Decision + +We will use **Azure Managed Identity** to authenticate to Azure Relay, replacing SAS token generation. + +At runtime, `DefaultAzureCredential` from the `azure-identity` SDK obtains a short-lived JWT from Azure AD, scoped to `https://relay.azure.com/.default`. This token is passed as an `Authorization: Bearer` HTTP header on the WebSocket upgrade request, which Azure Relay validates against Azure AD. + +The `AZURE_RELAY_KEY_NAME` and `AZURE_RELAY_SHARED_ACCESS_KEY` environment variables are removed. + +The gateway's managed identity must be assigned the **Azure Relay Listener** role on the hybrid connection resource in Azure. + +`DefaultAzureCredential` is used (rather than `ManagedIdentityCredential` directly) so that the credential chain works in all environments: managed identity in Azure deployments, and Azure CLI credentials on developer machines. + +## Consequences + +### Positive Consequences + +- **No secrets to manage:** No shared key to store, rotate, or accidentally leak +- **Fail-fast on misconfiguration:** A startup credential check raises `ClientAuthenticationError` immediately if the managed identity is not correctly configured, rather than failing silently in the reconnect loop +- **Consistent with platform direction:** Aligns with how the gateway already authenticates to the DICOM API (also via managed identity) + +### Negative Consequences + +- **Azure infrastructure dependency:** The managed identity and its role assignment must exist before the service can start; there is no fallback diff --git a/pyproject.toml b/pyproject.toml index 6cfcd827..e90609d9 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -16,6 +16,7 @@ dependencies = [ "requests>=2.33.1", "python-dotenv>=1.2.1", "azure-monitor-opentelemetry>=1.8.7", + "azure-identity>=1.23.0,<2.0.0", ] [dependency-groups] diff --git a/src/relay_listener.py b/src/relay_listener.py index 88e65012..45f506f0 100644 --- a/src/relay_listener.py +++ b/src/relay_listener.py @@ -5,19 +5,14 @@ """ import asyncio -import base64 -import hashlib -import hmac import json import logging import os -import time -import urllib.parse +from azure.identity import DefaultAzureCredential from dotenv import load_dotenv from websockets.asyncio.client import connect from websockets.exceptions import ConnectionClosedError -from websockets.frames import CloseCode from services.mwl.create_worklist_item import CreateWorklistItem from services.storage import MWLStorage @@ -28,8 +23,7 @@ logger = logging.getLogger(__name__) DB_PATH = os.getenv("MWL_DB_PATH", "/var/lib/pacs/worklist.db") -EXPIRED_TOKEN = "ExpiredToken" -SAS_TOKEN_EXPIRY_SECONDS = 3600 +AZURE_RELAY_SCOPE = "https://relay.azure.com/.default" class RelayListener: @@ -40,8 +34,6 @@ class RelayListener: 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) - AZURE_RELAY_KEY_NAME: Azure Relay shared access key name (default: RootManageSharedAccessKey) - AZURE_RELAY_SHARED_ACCESS_KEY: Azure Relay shared access key (default: none) MWL_DB_PATH: Path to the MWL SQLite database file (default: /var/lib/pacs/worklist.db) """ @@ -91,36 +83,31 @@ def process_action(self, payload: dict): def _connect(self): """Connect to Azure Relay.""" - return connect(self.relay_uri.connection_url(), compression=None) + return connect( + self.relay_uri.connection_url(), + compression=None, + additional_headers=self.relay_uri.auth_headers(), + ) 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.shared_access_key = os.getenv("AZURE_RELAY_SHARED_ACCESS_KEY", "") - - def create_sas_token(self, expiry_seconds: int = SAS_TOKEN_EXPIRY_SECONDS) -> str: - """Create SAS token for Azure Relay authentication.""" - 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() - ) - return ( - f"SharedAccessSignature sr={encoded_uri}" - f"&sig={urllib.parse.quote_plus(signature)}" - f"&se={expiry}&skn={self.key_name}" - ) + self._credential = DefaultAzureCredential() def connection_url(self) -> str: - token = self.create_sas_token() - return ( - f"wss://{self.relay_namespace}/$hc/{self.hybrid_connection_name}" - f"?sb-hc-action=listen&sb-hc-token={urllib.parse.quote_plus(token)}" - ) + return f"wss://{self.relay_namespace}/$hc/{self.hybrid_connection_name}?sb-hc-action=listen" + + def auth_headers(self) -> dict: + token = self._credential.get_token(AZURE_RELAY_SCOPE).token + return {"Authorization": f"Bearer {token}"} + + +def verify_credentials(): + """Verify managed identity credentials are available. Raises ClientAuthenticationError if not.""" + DefaultAzureCredential().get_token(AZURE_RELAY_SCOPE) + logger.info("Managed identity credentials verified.") async def main(): @@ -131,6 +118,7 @@ async def main(): configure_telemetry(service_name="relay-listener") logger.info("Socket Listener Starting...") + verify_credentials() storage = MWLStorage(db_path=DB_PATH) while True: @@ -142,13 +130,9 @@ 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" - - if code == CloseCode.INTERNAL_ERROR.value and EXPIRED_TOKEN in reason: - logger.info("SAS token expired, refreshing...") - else: - logger.warning(f"Connection closed with code {code}: {reason}") - logger.warning("Retrying in 5 seconds...") - await asyncio.sleep(5) + logger.warning(f"Connection closed with code {code}: {reason}") + logger.warning("Retrying in 5 seconds...") + await asyncio.sleep(5) except Exception as e: logger.warning(f"Connection error: {e}") logger.warning("Retrying in 5 seconds...") diff --git a/tests/conftest.py b/tests/conftest.py index 7bdfdf6e..e88642b6 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -2,7 +2,7 @@ import sys from contextlib import contextmanager from pathlib import Path -from unittest.mock import AsyncMock, patch +from unittest.mock import AsyncMock, MagicMock, patch import numpy as np import pytest @@ -15,6 +15,14 @@ sys.path.append(f"{Path(__file__).parent.parent}/src") +@pytest.fixture(autouse=True) +def mock_azure_credential(): + mock = MagicMock() + mock.get_token.return_value.token = "test-token" + with patch("relay_listener.DefaultAzureCredential", return_value=mock): + yield mock + + @pytest.fixture def tmp_dir(): return f"{Path(__file__).parent}/tmp" diff --git a/tests/test_relay_listener.py b/tests/test_relay_listener.py index c384f9d8..f6ec2a5b 100644 --- a/tests/test_relay_listener.py +++ b/tests/test_relay_listener.py @@ -2,11 +2,12 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest +from azure.core.exceptions import ClientAuthenticationError from websockets.exceptions import ConnectionClosedError from websockets.frames import Close, CloseCode from models import WorklistItem -from relay_listener import RelayListener, RelayURI, main +from relay_listener import RelayListener, RelayURI, main, verify_credentials class TestRelayListener: @@ -15,8 +16,6 @@ def setup(self, monkeypatch): monkeypatch.setenv("MWL_DB_PATH", "/tmp/test_worklist.db") monkeypatch.setenv("AZURE_RELAY_NAMESPACE", "test-namespace") monkeypatch.setenv("AZURE_RELAY_HYBRID_CONNECTION", "test-connection") - monkeypatch.setenv("AZURE_RELAY_KEY_NAME", "test-key-name") - monkeypatch.setenv("AZURE_RELAY_SHARED_ACCESS_KEY", "test-key-value") yield @pytest.fixture @@ -31,8 +30,6 @@ def test_relay_listener_initialization(self, storage_instance): assert isinstance(subject.relay_uri, RelayURI) assert subject.relay_uri.relay_namespace == "test-namespace" assert subject.relay_uri.hybrid_connection_name == "test-connection" - assert subject.relay_uri.key_name == "test-key-name" - assert subject.relay_uri.shared_access_key == "test-key-value" @pytest.mark.asyncio async def test_relay_listener_listen_echo(self, storage_instance, fake_relay): @@ -53,11 +50,9 @@ async def test_relay_listener_listen(self, storage_instance, listener_payload, f storage_instance.store_worklist_action.return_value = {"action_id": "action-12345", "status": "created"} subject = RelayListener(storage_instance) url = subject.relay_uri.connection_url() - assert url.startswith("wss://test-namespace/$hc/test-connection") - assert "sb-hc-token=" in url + assert url == "wss://test-namespace/$hc/test-connection?sb-hc-action=listen" 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: @@ -115,28 +110,27 @@ def test_process_action_invalid_type(self, storage_instance, listener_payload): storage_instance.store_worklist_item.assert_not_called() - def test_relay_uri_create_sas_token(self): + def test_relay_uri_connection_url(self): subject = RelayURI() - token = subject.create_sas_token(expiry_seconds=3600) + assert subject.connection_url() == "wss://test-namespace/$hc/test-connection?sb-hc-action=listen" - with patch("time.time", return_value=1000000): - token = subject.create_sas_token(expiry_seconds=3600) + def test_relay_uri_auth_headers(self, mock_azure_credential): + subject = RelayURI() + headers = subject.auth_headers() + assert headers == {"Authorization": "Bearer test-token"} + mock_azure_credential.get_token.assert_called_once_with("https://relay.azure.com/.default") - assert token == ( - "SharedAccessSignature sr=http%3A%2F%2Ftest-namespace%2Ftest-connection" - "&sig=PMcelSnwGlYX2xFo9Y2aGCg%2BvJ6LsHujiRrA1L6VnP0%3D&se=1003600&skn=test-key-name" - ) - def test_relay_uri_connection_url(self): - subject = RelayURI() - with patch("time.time", return_value=1000000): - url = subject.connection_url() +def test_verify_credentials_succeeds(mock_azure_credential): + verify_credentials() + mock_azure_credential.get_token.assert_called_once_with("https://relay.azure.com/.default") - assert url == ( - "wss://test-namespace/$hc/test-connection?sb-hc-action=listen" - "&sb-hc-token=SharedAccessSignature+sr%3Dhttp%253A%252F%252Ftest-namespace" - "%252Ftest-connection%26sig%3DPMcelSnwGlYX2xFo9Y2aGCg%252BvJ6LsHujiRrA1L6VnP0%253D%26se%3D1003600%26skn%3Dtest-key-name" - ) + +def test_verify_credentials_raises_on_failure(): + with patch("relay_listener.DefaultAzureCredential") as mock: + mock.return_value.get_token.side_effect = ClientAuthenticationError("no credentials") + with pytest.raises(ClientAuthenticationError): + verify_credentials() @patch("relay_listener.logger", new_callable=MagicMock) @@ -151,19 +145,16 @@ async def test_main_handles_connection_closed_and_keyboard_interrupt( relay_listener_instance.listen = AsyncMock() relay_listener_instance.listen.side_effect = [ - ConnectionClosedError(Close(CloseCode.INTERNAL_ERROR, "ExpiredToken"), None), - ConnectionClosedError(Close(CloseCode.INTERNAL_ERROR, "Something else"), None), + ConnectionClosedError(Close(CloseCode.INTERNAL_ERROR, "Something went wrong"), None), ConnectionClosedError(Close(CloseCode.BAD_GATEWAY, "Bad gateway"), None), KeyboardInterrupt(), ] await main() - assert relay_listener_instance.listen.call_count == 4 + assert relay_listener_instance.listen.call_count == 3 mock_logger.info.assert_any_call("Socket Listener Starting...") - mock_logger.info.assert_any_call("SAS token expired, refreshing...") - mock_logger.warning.assert_any_call("Connection closed with code 1011: Something else") + mock_logger.warning.assert_any_call("Connection closed with code 1011: 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("Retrying in 5 seconds...") mock_logger.warning.assert_any_call("\nShutting down...") diff --git a/uv.lock b/uv.lock index 13975b33..98c63e95 100644 --- a/uv.lock +++ b/uv.lock @@ -486,6 +486,7 @@ name = "manage-breast-screening-gateway" version = "0.1.0" source = { virtual = "." } dependencies = [ + { name = "azure-identity" }, { name = "azure-monitor-opentelemetry" }, { name = "numpy" }, { name = "pillow" }, @@ -511,10 +512,11 @@ dev = [ [package.metadata] requires-dist = [ + { name = "azure-identity", specifier = ">=1.23.0,<2.0.0" }, { name = "azure-monitor-opentelemetry", specifier = ">=1.8.7" }, { name = "numpy", specifier = ">=2.4.0,<3" }, { name = "pillow", specifier = ">=12.2.0" }, - { name = "pydicom", specifier = ">=3.0.1" }, + { name = "pydicom", specifier = ">=3.0.2" }, { name = "pylibjpeg", specifier = ">=2.0.0" }, { name = "pylibjpeg-openjpeg", specifier = ">=2.5.0" }, { name = "pynetdicom", specifier = ">=3.0.4" }, From 1c59daa5c09b5768715e627216fcce55959f61e4 Mon Sep 17 00:00:00 2001 From: Carlos Martinez Date: Thu, 30 Apr 2026 10:55:52 +0100 Subject: [PATCH 2/4] Use ManagedIdentityCredential in production only. In non-production environments, use a SAS token if set in the environment, otherwise default to DefaultAzureCredential. --- .env.development | 3 + ...DR-005_Managed_Identity_For_Azure_Relay.md | 24 ++-- src/relay_listener.py | 63 +++++++++- tests/conftest.py | 5 +- tests/test_relay_listener.py | 109 +++++++++++++++--- 5 files changed, 175 insertions(+), 29 deletions(-) diff --git a/.env.development b/.env.development index 8a790578..00c8a20e 100644 --- a/.env.development +++ b/.env.development @@ -1,6 +1,9 @@ # Azure Relay Configuration AZURE_RELAY_NAMESPACE=manbrs-gateway-dev.servicebus.windows.net AZURE_RELAY_HYBRID_CONNECTION=name-of-your-choice-relay-test-hc +# Optional: set these to use SAS token auth locally instead of managed identity / az login +AZURE_RELAY_KEY_NAME=RootManageSharedAccessKey +AZURE_RELAY_SHARED_ACCESS_KEY=YOUR_SHARED_ACCESS_KEY_HERE CLOUD_API_ENDPOINT=https://localhost:8000/api/v1/dicom CLOUD_API_TOKEN=testtoken diff --git a/docs/adr/ADR-005_Managed_Identity_For_Azure_Relay.md b/docs/adr/ADR-005_Managed_Identity_For_Azure_Relay.md index 617abbc5..8fd1181a 100644 --- a/docs/adr/ADR-005_Managed_Identity_For_Azure_Relay.md +++ b/docs/adr/ADR-005_Managed_Identity_For_Azure_Relay.md @@ -16,26 +16,32 @@ The initial implementation used Shared Access Signature (SAS) tokens. These are As the gateway runs inside the hospital network but is provisioned via Azure Arc, it can be assigned a managed identity through Arc-enabled infrastructure. Storing a long-lived shared secret in the environment is therefore unnecessary operational overhead and a potential security risk. +That said, setting up a working Relay connection locally is already complex. Mandating managed identity for all environments would add further friction for developers, who would need Azure CLI credentials with a Relay Listener role assignment before they could run the service. + ## Decision -We will use **Azure Managed Identity** to authenticate to Azure Relay, replacing SAS token generation. +In **production** (`ENVIRONMENT=prod`), the gateway uses `ManagedIdentityCredential` exclusively. The SAS token path is unavailable regardless of what environment variables are set. The gateway's managed identity must be assigned the **Azure Relay Listener** role on the hybrid connection resource in Azure. + +In **non-production** environments, the auth method is determined by whether `AZURE_RELAY_SHARED_ACCESS_KEY` is set: -At runtime, `DefaultAzureCredential` from the `azure-identity` SDK obtains a short-lived JWT from Azure AD, scoped to `https://relay.azure.com/.default`. This token is passed as an `Authorization: Bearer` HTTP header on the WebSocket upgrade request, which Azure Relay validates against Azure AD. +- If set, a SAS token is generated and embedded in the WebSocket URL (`sb-hc-token`), preserving the simpler local development setup. +- If absent, `DefaultAzureCredential` is used, which works with Azure CLI credentials (`az login`) for developers who have the Listener role assigned to their identity. -The `AZURE_RELAY_KEY_NAME` and `AZURE_RELAY_SHARED_ACCESS_KEY` environment variables are removed. +The token is passed as an `Authorization: Bearer` HTTP header on the WebSocket upgrade request for managed identity paths. Azure Relay validates it against Azure AD. -The gateway's managed identity must be assigned the **Azure Relay Listener** role on the hybrid connection resource in Azure. +A startup credential check (`verify_credentials()`) runs before the listen loop. In production this will raise `ClientAuthenticationError` immediately if the managed identity is not correctly configured. In non-production with a SAS key it logs the auth method and continues. -`DefaultAzureCredential` is used (rather than `ManagedIdentityCredential` directly) so that the credential chain works in all environments: managed identity in Azure deployments, and Azure CLI credentials on developer machines. +`ManagedIdentityCredential` is preferred over `DefaultAzureCredential` in production because it is predictable — it only attempts the IMDS endpoint and fails clearly, rather than traversing a credential chain that could succeed unexpectedly via another mechanism. ## Consequences ### Positive Consequences -- **No secrets to manage:** No shared key to store, rotate, or accidentally leak -- **Fail-fast on misconfiguration:** A startup credential check raises `ClientAuthenticationError` immediately if the managed identity is not correctly configured, rather than failing silently in the reconnect loop -- **Consistent with platform direction:** Aligns with how the gateway already authenticates to the DICOM API (also via managed identity) +- **No secrets in production:** No shared key to store, rotate, or accidentally leak in deployed environments +- **Fail-fast on misconfiguration:** Startup validation raises `ClientAuthenticationError` immediately rather than failing silently in the reconnect loop +- **Preserved local developer experience:** SAS tokens continue to work locally when `AZURE_RELAY_SHARED_ACCESS_KEY` is set +- **Consistent with platform direction:** Aligns with how the gateway already authenticates to the DICOM API ### Negative Consequences -- **Azure infrastructure dependency:** The managed identity and its role assignment must exist before the service can start; there is no fallback +- **Azure infrastructure dependency in production:** The managed identity and its role assignment must exist before the service can start diff --git a/src/relay_listener.py b/src/relay_listener.py index 45f506f0..ef823334 100644 --- a/src/relay_listener.py +++ b/src/relay_listener.py @@ -5,15 +5,21 @@ """ import asyncio +import base64 +import hashlib +import hmac import json import logging import os +import time +import urllib.parse -from azure.identity import DefaultAzureCredential +from azure.identity import DefaultAzureCredential, ManagedIdentityCredential from dotenv import load_dotenv from websockets.asyncio.client import connect from websockets.exceptions import ConnectionClosedError +from environment import Environment from services.mwl.create_worklist_item import CreateWorklistItem from services.storage import MWLStorage from telemetry import configure_telemetry @@ -24,6 +30,7 @@ DB_PATH = os.getenv("MWL_DB_PATH", "/var/lib/pacs/worklist.db") AZURE_RELAY_SCOPE = "https://relay.azure.com/.default" +SAS_TOKEN_EXPIRY_SECONDS = 3600 class RelayListener: @@ -35,6 +42,10 @@ class RelayListener: 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 """ def __init__(self, storage: MWLStorage): @@ -94,20 +105,60 @@ 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._credential = DefaultAzureCredential() + 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() + + def _use_sas(self) -> bool: + return not self._env.production and bool(self.shared_access_key) + + def _build_credential(self): + if self._env.production: + return ManagedIdentityCredential() + return DefaultAzureCredential() def connection_url(self) -> str: - return f"wss://{self.relay_namespace}/$hc/{self.hybrid_connection_name}?sb-hc-action=listen" + base = f"wss://{self.relay_namespace}/$hc/{self.hybrid_connection_name}?sb-hc-action=listen" + if self._use_sas(): + token = self._create_sas_token() + return f"{base}&sb-hc-token={urllib.parse.quote_plus(token)}" + return base def auth_headers(self) -> dict: + if self._use_sas(): + return {} token = self._credential.get_token(AZURE_RELAY_SCOPE).token return {"Authorization": f"Bearer {token}"} + def _create_sas_token(self, expiry_seconds: int = SAS_TOKEN_EXPIRY_SECONDS) -> str: + 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() + ) + return ( + f"SharedAccessSignature sr={encoded_uri}" + f"&sig={urllib.parse.quote_plus(signature)}" + f"&se={expiry}&skn={self.key_name}" + ) + def verify_credentials(): - """Verify managed identity credentials are available. Raises ClientAuthenticationError if not.""" - DefaultAzureCredential().get_token(AZURE_RELAY_SCOPE) - logger.info("Managed identity credentials verified.") + """ + 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. + """ + uri = RelayURI() + if uri._use_sas(): + logger.info("Using SAS token authentication for Azure Relay.") + else: + 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}).") async def main(): diff --git a/tests/conftest.py b/tests/conftest.py index e88642b6..e0c4a96c 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -19,7 +19,10 @@ def mock_azure_credential(): mock = MagicMock() mock.get_token.return_value.token = "test-token" - with patch("relay_listener.DefaultAzureCredential", return_value=mock): + with ( + patch("relay_listener.DefaultAzureCredential", return_value=mock), + patch("relay_listener.ManagedIdentityCredential", return_value=mock), + ): yield mock diff --git a/tests/test_relay_listener.py b/tests/test_relay_listener.py index f6ec2a5b..236378cd 100644 --- a/tests/test_relay_listener.py +++ b/tests/test_relay_listener.py @@ -49,8 +49,6 @@ async def test_relay_listener_listen_echo(self, storage_instance, fake_relay): async def test_relay_listener_listen(self, storage_instance, listener_payload, fake_relay): storage_instance.store_worklist_action.return_value = {"action_id": "action-12345", "status": "created"} subject = RelayListener(storage_instance) - url = subject.relay_uri.connection_url() - assert url == "wss://test-namespace/$hc/test-connection?sb-hc-action=listen" relay_message = json.dumps({"accept": {"address": "wss://accept-url"}}) client_payload = json.dumps(listener_payload) @@ -110,27 +108,112 @@ def test_process_action_invalid_type(self, storage_instance, listener_payload): storage_instance.store_worklist_item.assert_not_called() - def test_relay_uri_connection_url(self): + +class TestRelayURIWithDefaultAzureCredential: + """Non-production, no SAS key — uses DefaultAzureCredential.""" + + @pytest.fixture(autouse=True) + def setup(self, monkeypatch): + monkeypatch.setenv("AZURE_RELAY_NAMESPACE", "test-namespace") + monkeypatch.setenv("AZURE_RELAY_HYBRID_CONNECTION", "test-connection") + monkeypatch.delenv("AZURE_RELAY_SHARED_ACCESS_KEY", raising=False) + yield + + def test_connection_url(self): subject = RelayURI() assert subject.connection_url() == "wss://test-namespace/$hc/test-connection?sb-hc-action=listen" - def test_relay_uri_auth_headers(self, mock_azure_credential): + def test_auth_headers(self, mock_azure_credential): subject = RelayURI() - headers = subject.auth_headers() - assert headers == {"Authorization": "Bearer test-token"} + assert subject.auth_headers() == {"Authorization": "Bearer test-token"} mock_azure_credential.get_token.assert_called_once_with("https://relay.azure.com/.default") + def test_uses_default_azure_credential(self, mock_azure_credential): + with patch("relay_listener.DefaultAzureCredential") as mock_dac: + mock_dac.return_value = mock_azure_credential + subject = RelayURI() + assert subject._credential is mock_dac.return_value + -def test_verify_credentials_succeeds(mock_azure_credential): - verify_credentials() - mock_azure_credential.get_token.assert_called_once_with("https://relay.azure.com/.default") +class TestRelayURIWithSasToken: + """Non-production with SAS key present — uses SAS token.""" + + @pytest.fixture(autouse=True) + def setup(self, monkeypatch): + monkeypatch.setenv("AZURE_RELAY_NAMESPACE", "test-namespace") + monkeypatch.setenv("AZURE_RELAY_HYBRID_CONNECTION", "test-connection") + monkeypatch.setenv("AZURE_RELAY_KEY_NAME", "test-key-name") + monkeypatch.setenv("AZURE_RELAY_SHARED_ACCESS_KEY", "test-key-value") + yield + + def test_connection_url_includes_sas_token(self): + subject = RelayURI() + url = subject.connection_url() + assert url.startswith("wss://test-namespace/$hc/test-connection?sb-hc-action=listen") + assert "sb-hc-token=SharedAccessSignature" in url + + def test_auth_headers_are_empty(self): + subject = RelayURI() + assert subject.auth_headers() == {} + + def test_no_credential_is_created(self): + with patch("relay_listener.DefaultAzureCredential") as mock_dac: + with patch("relay_listener.ManagedIdentityCredential") as mock_mic: + RelayURI() + mock_dac.assert_not_called() + mock_mic.assert_not_called() + + +class TestRelayURIInProduction: + """Production environment — always uses ManagedIdentityCredential, never SAS.""" + + @pytest.fixture(autouse=True) + def setup(self, monkeypatch): + monkeypatch.setenv("AZURE_RELAY_NAMESPACE", "test-namespace") + monkeypatch.setenv("AZURE_RELAY_HYBRID_CONNECTION", "test-connection") + monkeypatch.setenv("ENVIRONMENT", "prod") + yield + + def test_uses_managed_identity_credential(self, mock_azure_credential): + with patch("relay_listener.ManagedIdentityCredential") as mock_mic: + mock_mic.return_value = mock_azure_credential + subject = RelayURI() + assert subject._credential is mock_mic.return_value + + def test_sas_key_is_ignored(self, monkeypatch): + monkeypatch.setenv("AZURE_RELAY_SHARED_ACCESS_KEY", "some-key") + subject = RelayURI() + assert not subject._use_sas() + assert "sb-hc-token" not in subject.connection_url() + + def test_auth_headers_use_bearer_token(self, mock_azure_credential): + subject = RelayURI() + assert subject.auth_headers() == {"Authorization": "Bearer test-token"} -def test_verify_credentials_raises_on_failure(): - with patch("relay_listener.DefaultAzureCredential") as mock: - mock.return_value.get_token.side_effect = ClientAuthenticationError("no credentials") - with pytest.raises(ClientAuthenticationError): +class TestVerifyCredentials: + def test_logs_sas_when_key_present(self, monkeypatch): + monkeypatch.setenv("AZURE_RELAY_SHARED_ACCESS_KEY", "test-key") + with patch("relay_listener.logger") as mock_logger: verify_credentials() + mock_logger.info.assert_called_with("Using SAS token authentication for Azure Relay.") + + def test_verifies_default_azure_credential_when_no_key(self, mock_azure_credential, monkeypatch): + monkeypatch.delenv("AZURE_RELAY_SHARED_ACCESS_KEY", raising=False) + verify_credentials() + mock_azure_credential.get_token.assert_called_with("https://relay.azure.com/.default") + + def test_verifies_managed_identity_in_production(self, mock_azure_credential, monkeypatch): + monkeypatch.setenv("ENVIRONMENT", "prod") + verify_credentials() + mock_azure_credential.get_token.assert_called_with("https://relay.azure.com/.default") + + def test_raises_on_credential_failure(self, monkeypatch): + monkeypatch.delenv("AZURE_RELAY_SHARED_ACCESS_KEY", raising=False) + with patch("relay_listener.DefaultAzureCredential") as mock: + mock.return_value.get_token.side_effect = ClientAuthenticationError("no credentials") + with pytest.raises(ClientAuthenticationError): + verify_credentials() @patch("relay_listener.logger", new_callable=MagicMock) From fc74dd0b970d603692fa8ae44048d808967002a8 Mon Sep 17 00:00:00 2001 From: Carlos Martinez Date: Thu, 30 Apr 2026 11:09:06 +0100 Subject: [PATCH 3/4] Appease pyright gods --- src/relay_listener.py | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/src/relay_listener.py b/src/relay_listener.py index ef823334..ef1e8c5e 100644 --- a/src/relay_listener.py +++ b/src/relay_listener.py @@ -128,6 +128,8 @@ def connection_url(self) -> str: def auth_headers(self) -> dict: if self._use_sas(): return {} + if self._credential is None: + raise RuntimeError("No credential available — _credential should never be None when not using SAS") token = self._credential.get_token(AZURE_RELAY_SCOPE).token return {"Authorization": f"Bearer {token}"} @@ -156,6 +158,8 @@ def verify_credentials(): if uri._use_sas(): logger.info("Using SAS token authentication for Azure Relay.") else: + if uri._credential is None: + raise RuntimeError("No credential available — _credential should never be None when not using SAS") 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}).") From bd3d7187a4411ed64e507b5904e326136d5f6e33 Mon Sep 17 00:00:00 2001 From: Carlos Martinez Date: Thu, 30 Apr 2026 14:48:23 +0100 Subject: [PATCH 4/4] Define CredentialNotAvailableError --- src/relay_listener.py | 12 ++++++++++-- tests/test_relay_listener.py | 2 +- 2 files changed, 11 insertions(+), 3 deletions(-) diff --git a/src/relay_listener.py b/src/relay_listener.py index ef1e8c5e..bef3f9dc 100644 --- a/src/relay_listener.py +++ b/src/relay_listener.py @@ -33,6 +33,10 @@ SAS_TOKEN_EXPIRY_SECONDS = 3600 +class CredentialNotAvailableError(RuntimeError): + pass + + class RelayListener: """ Socket Listener for Azure Relay. @@ -129,7 +133,9 @@ def auth_headers(self) -> dict: if self._use_sas(): return {} if self._credential is None: - raise RuntimeError("No credential available — _credential should never be None when not using SAS") + raise CredentialNotAvailableError( + "No credential available — _credential should never be None when not using SAS" + ) token = self._credential.get_token(AZURE_RELAY_SCOPE).token return {"Authorization": f"Bearer {token}"} @@ -159,7 +165,9 @@ def verify_credentials(): logger.info("Using SAS token authentication for Azure Relay.") else: if uri._credential is None: - raise RuntimeError("No credential available — _credential should never be None when not using SAS") + raise CredentialNotAvailableError( + "No credential available — _credential should never be None when not using SAS" + ) 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}).") diff --git a/tests/test_relay_listener.py b/tests/test_relay_listener.py index 236378cd..548bfced 100644 --- a/tests/test_relay_listener.py +++ b/tests/test_relay_listener.py @@ -208,7 +208,7 @@ def test_verifies_managed_identity_in_production(self, mock_azure_credential, mo verify_credentials() mock_azure_credential.get_token.assert_called_with("https://relay.azure.com/.default") - def test_raises_on_credential_failure(self, monkeypatch): + def test_raises_client_authentication_error_on_credential_failure(self, monkeypatch): monkeypatch.delenv("AZURE_RELAY_SHARED_ACCESS_KEY", raising=False) with patch("relay_listener.DefaultAzureCredential") as mock: mock.return_value.get_token.side_effect = ClientAuthenticationError("no credentials")