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
18 changes: 11 additions & 7 deletions msal/managed_identity.py
Original file line number Diff line number Diff line change
Expand Up @@ -198,9 +198,12 @@ def __init__(
client = msal.ManagedIdentityClient(managed_identity, http_client=s)

For Service Fabric managed identity, ``http_client`` must be a
``requests.Session`` using the standard ``requests.adapters.HTTPAdapter``.
``requests.Session`` using ``requests.adapters.HTTPAdapter`` or a
subclass for the Service Fabric endpoint.
MSAL derives a separate session for the Service Fabric endpoint so that
its certificate thumbprint can be validated before the Secret header is sent.
Standard session, retry, and connection-pool settings are preserved,
but custom adapter behavior is not used.

:param token_cache:
Optional. It accepts a :class:`msal.TokenCache` instance to store tokens.
Expand Down Expand Up @@ -726,21 +729,22 @@ def cert_verify(self, conn, url, verify, cert):


def _create_service_fabric_http_client(http_client, endpoint, server_thumbprint):
"""Clone a standard Requests session and attach a pinning-only HTTPS transport.
"""Derive a Requests session with a pinning-only HTTPS transport.

Custom HTTP clients and adapters are rejected because MSAL cannot prove that
they will validate the certificate before transmitting the Secret header.
HTTPAdapter subclasses are accepted as sources of retry and pool settings,
but the derived session always uses MSAL's certificate-pinning adapter.
Other HTTP clients and adapters are unsupported.
"""
if isinstance(http_client, ThrottledHttpClientBase):
http_client = http_client.http_client
if not isinstance(http_client, requests.Session):
raise ManagedIdentityError(
"Service Fabric managed identity requires a requests.Session "
"with the standard HTTPAdapter.")
"with an HTTPAdapter or subclass.")
source_adapter = http_client.get_adapter(endpoint)
if type(source_adapter) is not HTTPAdapter:
if not isinstance(source_adapter, HTTPAdapter):
raise ManagedIdentityError(
"Service Fabric managed identity does not support custom HTTP adapters.")
"Service Fabric managed identity requires an HTTPAdapter or subclass.")

service_fabric_client = requests.Session()
service_fabric_client.headers = http_client.headers.copy()
Expand Down
70 changes: 56 additions & 14 deletions tests/test_mi.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,7 @@
except:
from mock import patch, ANY, mock_open, Mock
import requests
from requests.adapters import HTTPAdapter
from requests.adapters import BaseAdapter, HTTPAdapter
from requests.exceptions import SSLError
from urllib3.util.retry import Retry
from cryptography import x509
Expand Down Expand Up @@ -468,6 +468,8 @@ def log_message(self, format, *args):


class ServiceFabricTlsValidationTestCase(unittest.TestCase):
_adapter_class = HTTPAdapter

def setUp(self):
self._temporary_directory = tempfile.TemporaryDirectory()
private_key = rsa.generate_private_key(public_exponent=65537, key_size=2048)
Expand Down Expand Up @@ -518,14 +520,16 @@ def tearDown(self):
self.server_thread.join()
self._temporary_directory.cleanup()

def _new_session(self):
def _new_session(self, **adapter_kwargs):
session = requests.Session()
self.addCleanup(session.close)
session.trust_env = False
session.mount("https://", self._adapter_class(**adapter_kwargs))
return session

def test_matching_thumbprint_sends_secret_after_validating_certificate(self):
result = _obtain_token_on_service_fabric(
self._new_session(),
_ThrottledHttpClient(self._new_session()),
self.endpoint,
"service-fabric-secret",
":".join(
Expand All @@ -541,9 +545,11 @@ def test_matching_thumbprint_sends_secret_after_validating_certificate(self):
self.server.requests[0]["headers"]["Secret"])

def test_mismatching_thumbprint_prevents_the_secret_from_being_sent(self):
source = self._new_session()
source.verify = False
with self.assertRaises(SSLError):
_obtain_token_on_service_fabric(
self._new_session(),
source,
self.endpoint,
"service-fabric-secret",
"00" * 20,
Expand All @@ -564,19 +570,39 @@ def test_non_https_endpoint_is_rejected_before_a_request_can_send_the_secret(sel

self.assertEqual([], self.server.requests)

def test_malformed_thumbprint_is_rejected_before_a_request_can_send_the_secret(self):
with self.assertRaises(ManagedIdentityError):
_obtain_token_on_service_fabric(
self._new_session(),
self.endpoint,
"service-fabric-secret",
"not-a-thumbprint",
"R",
)

self.assertEqual([], self.server.requests)

def test_derived_client_preserves_standard_session_settings_without_mutation(self):
source = self._new_session()
source = self._new_session(
max_retries=Retry(total=2),
pool_connections=3,
pool_maxsize=7,
pool_block=True,
)
source.verify = False
source.headers["X-Caller-Header"] = "caller-header"
source.cookies.set("caller-cookie", "cookie-value")
source.auth = ("caller", "password")
source.params = {"caller-param": "caller-value"}
source.proxies = {}
source.max_redirects = 7
source.mount("https://", HTTPAdapter(max_retries=Retry(total=2)))
source_adapter = source.get_adapter(self.endpoint)
source_pool_classes = source_adapter.poolmanager.pool_classes_by_scheme.copy()

derived = _create_service_fabric_http_client(
source, self.endpoint, self.thumbprint)
self.addCleanup(derived.close)
derived_adapter = derived.get_adapter(self.endpoint)

self.assertFalse(source.verify)
self.assertTrue(derived.verify)
Expand All @@ -586,8 +612,17 @@ def test_derived_client_preserves_standard_session_settings_without_mutation(sel
self.assertEqual(source.params, derived.params)
self.assertEqual(source.proxies, derived.proxies)
self.assertEqual(source.max_redirects, derived.max_redirects)
self.assertEqual(2, derived.get_adapter(self.endpoint).max_retries.total)
self.assertIsNot(source.get_adapter(self.endpoint), derived.get_adapter(self.endpoint))
self.assertEqual(2, derived_adapter.max_retries.total)
self.assertIsNot(source_adapter.max_retries, derived_adapter.max_retries)
self.assertEqual(3, derived_adapter._pool_connections)
self.assertEqual(7, derived_adapter._pool_maxsize)
self.assertTrue(derived_adapter._pool_block)
self.assertIs(source_adapter, source.get_adapter(self.endpoint))
self.assertIsNot(source_adapter, derived_adapter)
self.assertIsNot(
source_adapter.poolmanager.pool_classes_by_scheme,
derived_adapter.poolmanager.pool_classes_by_scheme)
self.assertEqual(source_pool_classes, source_adapter.poolmanager.pool_classes_by_scheme)
response = derived.get(
self.endpoint,
params={"request-param": "request-value"},
Expand All @@ -601,13 +636,10 @@ def test_derived_client_preserves_standard_session_settings_without_mutation(sel
self.assertIn("caller-cookie=cookie-value", request["headers"]["Cookie"])
self.assertTrue(request["headers"]["Authorization"].startswith("Basic "))

def test_custom_adapter_is_rejected_before_a_request_can_send_the_secret(self):
def test_non_http_adapter_is_rejected_before_a_request_can_send_the_secret(self):
source = self._new_session()

class CustomAdapter(HTTPAdapter):
pass

source.mount("https://", CustomAdapter())
adapter = Mock(spec=BaseAdapter)
source.mount("https://", adapter)
with self.assertRaises(ManagedIdentityError):
_obtain_token_on_service_fabric(
source,
Expand All @@ -618,6 +650,16 @@ class CustomAdapter(HTTPAdapter):
)

self.assertEqual([], self.server.requests)
adapter.send.assert_not_called()


class _ServiceFabricSourceHTTPAdapter(HTTPAdapter):
def send(self, *args, **kwargs):
raise AssertionError("Service Fabric requests must use MSAL's pinned adapter")


class ServiceFabricHTTPAdapterSubclassTestCase(ServiceFabricTlsValidationTestCase):
_adapter_class = _ServiceFabricSourceHTTPAdapter


@patch.dict(os.environ, {
Expand Down
Loading