diff --git a/azure-iot-device/azure/iot/device/iothub/edge_hsm.py b/azure-iot-device/azure/iot/device/iothub/edge_hsm.py index 4c23beabb..8c01cbb0d 100644 --- a/azure-iot-device/azure/iot/device/iothub/edge_hsm.py +++ b/azure-iot-device/azure/iot/device/iothub/edge_hsm.py @@ -13,7 +13,6 @@ from azure.iot.device.common.auth.signing_mechanism import SigningMechanism from azure.iot.device import user_agent -requests_unixsocket.monkeypatch() logger = logging.getLogger(__name__) @@ -58,27 +57,27 @@ def get_certificate(self): :raises: IoTEdgeError if unable to retrieve the certificate. """ - r = requests.get( - self.workload_uri + "trust-bundle", - params={"api-version": self.api_version}, - headers={"User-Agent": urllib.parse.quote_plus(user_agent.get_iothub_user_agent())}, - ) - # Validate that the request was successful - try: - r.raise_for_status() - except requests.exceptions.HTTPError as e: - raise IoTEdgeError("Unable to get trust bundle from Edge") from e - # Decode the trust bundle - try: - bundle = r.json() - except ValueError as e: - raise IoTEdgeError("Unable to decode trust bundle") from e - # Retrieve the certificate - try: - cert = bundle["certificate"] - except KeyError as e: - raise IoTEdgeError("No certificate in trust bundle") from e - return cert + with requests_unixsocket.Session() as session: + r = session.get( + self.workload_uri + "trust-bundle", + params={"api-version": self.api_version}, + headers={"User-Agent": urllib.parse.quote_plus(user_agent.get_iothub_user_agent())}, + ) + # Validate that the request was successful + try: + r.raise_for_status() + except requests.exceptions.HTTPError as e: + raise IoTEdgeError("Unable to get trust bundle from Edge") from e + # Decode the trust bundle + try: + bundle = r.json() + except ValueError as e: + raise IoTEdgeError("Unable to decode trust bundle") from e + # Retrieve the certificate + try: + return bundle["certificate"] + except KeyError as e: + raise IoTEdgeError("No certificate in trust bundle") from e def sign(self, data_str): """ @@ -99,26 +98,27 @@ def sign(self, data_str): ) sign_request = {"keyId": "primary", "algo": "HMACSHA256", "data": encoded_data_str} - r = requests.post( # can we use json field instead of data? - url=path, - params={"api-version": self.api_version}, - headers={"User-Agent": urllib.parse.quote(user_agent.get_iothub_user_agent(), safe="")}, - data=json.dumps(sign_request), - ) - try: - r.raise_for_status() - except requests.exceptions.HTTPError as e: - raise IoTEdgeError("Unable to sign data") from e - try: - sign_response = r.json() - except ValueError as e: - raise IoTEdgeError("Unable to decode signed data") from e - try: - signed_data_str = sign_response["digest"] - except KeyError as e: - raise IoTEdgeError("No signed data received") from e - - return signed_data_str # what format is this? string? bytes? + with requests_unixsocket.Session() as session: + r = session.post( # can we use json field instead of data? + url=path, + params={"api-version": self.api_version}, + headers={ + "User-Agent": urllib.parse.quote(user_agent.get_iothub_user_agent(), safe="") + }, + data=json.dumps(sign_request), + ) + try: + r.raise_for_status() + except requests.exceptions.HTTPError as e: + raise IoTEdgeError("Unable to sign data") from e + try: + sign_response = r.json() + except ValueError as e: + raise IoTEdgeError("Unable to decode signed data") from e + try: + return sign_response["digest"] + except KeyError as e: + raise IoTEdgeError("No signed data received") from e def _format_socket_uri(old_uri): diff --git a/tests/unit/iothub/test_edge_hsm.py b/tests/unit/iothub/test_edge_hsm.py index 8472ad11f..d0c97606e 100644 --- a/tests/unit/iothub/test_edge_hsm.py +++ b/tests/unit/iothub/test_edge_hsm.py @@ -4,22 +4,26 @@ # license information. # -------------------------------------------------------------------------- +import importlib import pytest import logging import requests +import requests_unixsocket import json import base64 import urllib -from azure.iot.device.iothub.edge_hsm import IoTEdgeHsm, IoTEdgeError -from azure.iot.device import user_agent +# Keep module-qualified references because the import regression test reloads this module. +# Directly imported classes would retain their pre-reload identities. +from azure.iot.device.iothub import edge_hsm as edge_hsm_module +from azure.iot.device import user_agent logging.basicConfig(level=logging.DEBUG) @pytest.fixture def edge_hsm(): - return IoTEdgeHsm( + return edge_hsm_module.IoTEdgeHsm( module_id="my_module_id", generation_id="module_generation_id", workload_uri="unix:///var/run/iotedge/workload.sock", @@ -27,8 +31,24 @@ def edge_hsm(): ) +@pytest.fixture +def mock_unix_session(mocker): + mock_session_constructor = mocker.patch.object(requests_unixsocket, "Session") + mock_session = mock_session_constructor.return_value + mock_session.__enter__.return_value = mock_session + return mock_session + + @pytest.mark.describe("IoTEdgeHsm - Instantiation") class TestIoTEdgeHsmInstantiation(object): + @pytest.mark.it("Does not monkeypatch the global requests API when imported") + def test_does_not_monkeypatch_requests(self, mocker): + mock_monkeypatch = mocker.patch.object(requests_unixsocket, "monkeypatch") + + importlib.reload(edge_hsm_module) + + assert mock_monkeypatch.call_count == 0 + @pytest.mark.it("URL encodes the provided module_id parameter and sets it as an attribute") def test_encode_and_set_module_id(self): module_id = "my_module_id" @@ -36,7 +56,7 @@ def test_encode_and_set_module_id(self): api_version = "my_api_version" workload_uri = "unix:///var/run/iotedge/workload.sock" - edge_hsm = IoTEdgeHsm( + edge_hsm = edge_hsm_module.IoTEdgeHsm( module_id=module_id, generation_id=generation_id, workload_uri=workload_uri, @@ -64,7 +84,7 @@ def test_workload_uri_formatting(self, workload_uri, expected_formatted_uri): generation_id = "my_generation_id" api_version = "my_api_version" - edge_hsm = IoTEdgeHsm( + edge_hsm = edge_hsm_module.IoTEdgeHsm( module_id=module_id, generation_id=generation_id, workload_uri=workload_uri, @@ -80,7 +100,7 @@ def test_set_generation_id(self): api_version = "my_api_version" workload_uri = "unix:///var/run/iotedge/workload.sock" - edge_hsm = IoTEdgeHsm( + edge_hsm = edge_hsm_module.IoTEdgeHsm( module_id=module_id, generation_id=generation_id, workload_uri=workload_uri, @@ -96,7 +116,7 @@ def test_set_api_version(self): api_version = "my_api_version" workload_uri = "unix:///var/run/iotedge/workload.sock" - edge_hsm = IoTEdgeHsm( + edge_hsm = edge_hsm_module.IoTEdgeHsm( module_id=module_id, generation_id=generation_id, workload_uri=workload_uri, @@ -108,9 +128,19 @@ def test_set_api_version(self): @pytest.mark.describe("IoTEdgeHsm - .get_certificate()") class TestIoTEdgeHsmGetCertificate(object): + @pytest.mark.it("Closes the Unix socket session after the request") + def test_closes_session(self, mocker, edge_hsm, mock_unix_session): + mock_unix_session.get.return_value.json.return_value = {"certificate": "my certificate"} + + edge_hsm.get_certificate() + + assert requests_unixsocket.Session.call_args == mocker.call() + assert mock_unix_session.__enter__.call_args == mocker.call() + assert mock_unix_session.__exit__.call_count == 1 + @pytest.mark.it("Sends an HTTP GET request to retrieve the trust bundle from Edge") - def test_requests_trust_bundle(self, mocker, edge_hsm): - mock_request_get = mocker.patch.object(requests, "get") + def test_requests_trust_bundle(self, mocker, edge_hsm, mock_unix_session): + mock_request_get = mock_unix_session.get expected_url = edge_hsm.workload_uri + "trust-bundle" expected_params = {"api-version": edge_hsm.api_version} expected_headers = { @@ -125,8 +155,8 @@ def test_requests_trust_bundle(self, mocker, edge_hsm): ) @pytest.mark.it("Returns the certificate from the trust bundle received from Edge") - def test_returns_certificate(self, mocker, edge_hsm): - mock_request_get = mocker.patch.object(requests, "get") + def test_returns_certificate(self, edge_hsm, mock_unix_session): + mock_request_get = mock_unix_session.get mock_response = mock_request_get.return_value certificate = "my certificate" mock_response.json.return_value = {"certificate": certificate} @@ -136,47 +166,57 @@ def test_returns_certificate(self, mocker, edge_hsm): assert returned_cert is certificate @pytest.mark.it("Raises IoTEdgeError if a bad request is made to Edge") - def test_bad_request(self, mocker, edge_hsm): - mock_request_get = mocker.patch.object(requests, "get") + def test_bad_request(self, edge_hsm, mock_unix_session): + mock_request_get = mock_unix_session.get mock_response = mock_request_get.return_value error = requests.exceptions.HTTPError() mock_response.raise_for_status.side_effect = error - with pytest.raises(IoTEdgeError) as e_info: + with pytest.raises(edge_hsm_module.IoTEdgeError) as e_info: edge_hsm.get_certificate() assert e_info.value.__cause__ is error @pytest.mark.it("Raises IoTEdgeError if there is an error in json decoding the trust bundle") - def test_bad_json(self, mocker, edge_hsm): - mock_request_get = mocker.patch.object(requests, "get") + def test_bad_json(self, edge_hsm, mock_unix_session): + mock_request_get = mock_unix_session.get mock_response = mock_request_get.return_value error = ValueError() mock_response.json.side_effect = error - with pytest.raises(IoTEdgeError) as e_info: + with pytest.raises(edge_hsm_module.IoTEdgeError) as e_info: edge_hsm.get_certificate() assert e_info.value.__cause__ is error @pytest.mark.it("Raises IoTEdgeError if the certificate is missing from the trust bundle") - def test_bad_trust_bundle(self, mocker, edge_hsm): - mock_request_get = mocker.patch.object(requests, "get") + def test_bad_trust_bundle(self, edge_hsm, mock_unix_session): + mock_request_get = mock_unix_session.get mock_response = mock_request_get.return_value # Return an empty json dict with no 'certificate' key mock_response.json.return_value = {} - with pytest.raises(IoTEdgeError): + with pytest.raises(edge_hsm_module.IoTEdgeError): edge_hsm.get_certificate() @pytest.mark.describe("IoTEdgeHsm - .sign()") class TestIoTEdgeHsmSign(object): + @pytest.mark.it("Closes the Unix socket session after the request") + def test_closes_session(self, mocker, edge_hsm, mock_unix_session): + mock_unix_session.post.return_value.json.return_value = {"digest": "somedigest"} + + edge_hsm.sign("somedata") + + assert requests_unixsocket.Session.call_args == mocker.call() + assert mock_unix_session.__enter__.call_args == mocker.call() + assert mock_unix_session.__exit__.call_count == 1 + @pytest.mark.it( "Makes an HTTP request to Edge to sign a piece of string data using the HMAC-SHA256 algorithm" ) - def test_requests_data_signing(self, mocker, edge_hsm): + def test_requests_data_signing(self, mocker, edge_hsm, mock_unix_session): data_str = "somedata" data_str_b64 = "c29tZWRhdGE=" - mock_request_post = mocker.patch.object(requests, "post") + mock_request_post = mock_unix_session.post mock_request_post.return_value.json.return_value = {"digest": "somedigest"} expected_url = "{workload_uri}modules/{module_id}/genid/{generation_id}/sign".format( workload_uri=edge_hsm.workload_uri, @@ -197,12 +237,12 @@ def test_requests_data_signing(self, mocker, edge_hsm): ) @pytest.mark.it("Base64 encodes the string data in the request") - def test_b64_encodes_data(self, mocker, edge_hsm): + def test_b64_encodes_data(self, edge_hsm, mock_unix_session): # This test is actually implicitly tested in the first test, but it's # important to have an explicit test for it since it's a requirement data_str = "somedata" data_str_b64 = base64.b64encode(data_str.encode("utf-8")).decode() - mock_request_post = mocker.patch.object(requests, "post") + mock_request_post = mock_unix_session.post mock_request_post.return_value.json.return_value = {"digest": "somedigest"} edge_hsm.sign(data_str) @@ -213,9 +253,9 @@ def test_b64_encodes_data(self, mocker, edge_hsm): assert sent_data == data_str_b64 @pytest.mark.it("Returns the signed data received from Edge") - def test_returns_signed_data(self, mocker, edge_hsm): + def test_returns_signed_data(self, edge_hsm, mock_unix_session): expected_digest = "somedigest" - mock_request_post = mocker.patch.object(requests, "post") + mock_request_post = mock_unix_session.post mock_request_post.return_value.json.return_value = {"digest": expected_digest} signed_data = edge_hsm.sign("somedata") @@ -223,31 +263,31 @@ def test_returns_signed_data(self, mocker, edge_hsm): assert signed_data == expected_digest @pytest.mark.it("Raises IoTEdgeError if a bad request is made to EdgeHub") - def test_bad_request(self, mocker, edge_hsm): - mock_request_post = mocker.patch.object(requests, "post") + def test_bad_request(self, edge_hsm, mock_unix_session): + mock_request_post = mock_unix_session.post mock_response = mock_request_post.return_value error = requests.exceptions.HTTPError() mock_response.raise_for_status.side_effect = error - with pytest.raises(IoTEdgeError) as e_info: + with pytest.raises(edge_hsm_module.IoTEdgeError) as e_info: edge_hsm.sign("somedata") assert e_info.value.__cause__ is error @pytest.mark.it("Raises IoTEdgeError if there is an error in json decoding the signed response") - def test_bad_json(self, mocker, edge_hsm): - mock_request_post = mocker.patch.object(requests, "post") + def test_bad_json(self, edge_hsm, mock_unix_session): + mock_request_post = mock_unix_session.post mock_response = mock_request_post.return_value error = ValueError() mock_response.json.side_effect = error - with pytest.raises(IoTEdgeError) as e_info: + with pytest.raises(edge_hsm_module.IoTEdgeError) as e_info: edge_hsm.sign("somedata") assert e_info.value.__cause__ is error @pytest.mark.it("Raises IoTEdgeError if the signed data is missing from the response") - def test_bad_response(self, mocker, edge_hsm): - mock_request_post = mocker.patch.object(requests, "post") + def test_bad_response(self, edge_hsm, mock_unix_session): + mock_request_post = mock_unix_session.post mock_response = mock_request_post.return_value mock_response.json.return_value = {} - with pytest.raises(IoTEdgeError): + with pytest.raises(edge_hsm_module.IoTEdgeError): edge_hsm.sign("somedata")