diff --git a/src/relay_listener.py b/src/relay_listener.py index bef3f9dc..0f0d59c3 100644 --- a/src/relay_listener.py +++ b/src/relay_listener.py @@ -29,7 +29,7 @@ logger = logging.getLogger(__name__) DB_PATH = os.getenv("MWL_DB_PATH", "/var/lib/pacs/worklist.db") -AZURE_RELAY_SCOPE = "https://relay.azure.com/.default" +AZURE_RELAY_SCOPE = "https://relay.azure.net/.default" SAS_TOKEN_EXPIRY_SECONDS = 3600 @@ -101,7 +101,6 @@ def _connect(self): return connect( self.relay_uri.connection_url(), compression=None, - additional_headers=self.relay_uri.auth_headers(), ) @@ -126,18 +125,16 @@ def connection_url(self) -> str: 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 + else: + token = self._create_bearer_token() + return f"{base}&sb-hc-token={urllib.parse.quote_plus(token)}" - def auth_headers(self) -> dict: - if self._use_sas(): - return {} + def _create_bearer_token(self) -> str: if self._credential is None: 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}"} + return f"Bearer {self._credential.get_token(AZURE_RELAY_SCOPE).token}" def _create_sas_token(self, expiry_seconds: int = SAS_TOKEN_EXPIRY_SECONDS) -> str: uri = f"http://{self.relay_namespace}/{self.hybrid_connection_name}" diff --git a/tests/test_relay_listener.py b/tests/test_relay_listener.py index 548bfced..5fb6ac9e 100644 --- a/tests/test_relay_listener.py +++ b/tests/test_relay_listener.py @@ -119,14 +119,11 @@ def setup(self, monkeypatch): monkeypatch.delenv("AZURE_RELAY_SHARED_ACCESS_KEY", raising=False) yield - def test_connection_url(self): + def test_connection_url(self, mock_azure_credential): subject = RelayURI() - assert subject.connection_url() == "wss://test-namespace/$hc/test-connection?sb-hc-action=listen" - - def test_auth_headers(self, mock_azure_credential): - subject = RelayURI() - assert subject.auth_headers() == {"Authorization": "Bearer test-token"} - mock_azure_credential.get_token.assert_called_once_with("https://relay.azure.com/.default") + url = subject.connection_url() + assert url.startswith("wss://test-namespace/$hc/test-connection?sb-hc-action=listen") + assert "sb-hc-token=Bearer+test-token" in url def test_uses_default_azure_credential(self, mock_azure_credential): with patch("relay_listener.DefaultAzureCredential") as mock_dac: @@ -152,10 +149,6 @@ 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_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: @@ -180,15 +173,11 @@ def test_uses_managed_identity_credential(self, mock_azure_credential): subject = RelayURI() assert subject._credential is mock_mic.return_value - def test_sas_key_is_ignored(self, monkeypatch): + def test_sas_key_is_ignored(self, mock_azure_credential, 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"} + assert "sb-hc-token=Bearer+test-token" in subject.connection_url() class TestVerifyCredentials: @@ -201,12 +190,12 @@ def test_logs_sas_when_key_present(self, monkeypatch): 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") + mock_azure_credential.get_token.assert_called_with("https://relay.azure.net/.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") + mock_azure_credential.get_token.assert_called_with("https://relay.azure.net/.default") def test_raises_client_authentication_error_on_credential_failure(self, monkeypatch): monkeypatch.delenv("AZURE_RELAY_SHARED_ACCESS_KEY", raising=False)