diff --git a/msal/managed_identity.py b/msal/managed_identity.py index 7fc9e4bd..9db2660a 100644 --- a/msal/managed_identity.py +++ b/msal/managed_identity.py @@ -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. @@ -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() diff --git a/tests/test_mi.py b/tests/test_mi.py index f0d0260c..c7c1fef6 100644 --- a/tests/test_mi.py +++ b/tests/test_mi.py @@ -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 @@ -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) @@ -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( @@ -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, @@ -564,8 +570,25 @@ 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") @@ -573,10 +596,13 @@ def test_derived_client_preserves_standard_session_settings_without_mutation(sel 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) @@ -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"}, @@ -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, @@ -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, {