Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
84 changes: 42 additions & 42 deletions azure-iot-device/azure/iot/device/iothub/edge_hsm.py
Original file line number Diff line number Diff line change
Expand Up @@ -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__)


Expand Down Expand Up @@ -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):
"""
Expand All @@ -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):
Expand Down
110 changes: 75 additions & 35 deletions tests/unit/iothub/test_edge_hsm.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,39 +4,59 @@
# 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",
api_version="my_api_version",
)


@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"
generation_id = "my_generation_id"
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,
Expand Down Expand Up @@ -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,
Expand All @@ -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,
Expand All @@ -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,
Expand All @@ -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 = {
Expand All @@ -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}
Expand All @@ -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,
Expand All @@ -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)
Expand All @@ -213,41 +253,41 @@ 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")

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")