From 3835390cd95e1bafb2b98ab05178eaa1f201cb3e Mon Sep 17 00:00:00 2001 From: Nilesh Choudhary Date: Thu, 24 Sep 2026 10:25:18 +0100 Subject: [PATCH 1/4] feat: add opt-in Service Fabric HTTP options MSAL-owned secure transport accepts unopened custom clients via explicit options. None preserves legacy behavior. HTTPS proxies are authenticated separately. Per-call options are out of scope. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- docs/index.rst | 138 +++++ msal/__init__.py | 1 + msal/managed_identity.py | 275 +++++++++- tests/test_mi.py | 1109 +++++++++++++++++++++++++++++++++++++- 4 files changed, 1497 insertions(+), 26 deletions(-) diff --git a/docs/index.rst b/docs/index.rst index 1c5d79bc..024e725d 100644 --- a/docs/index.rst +++ b/docs/index.rst @@ -188,3 +188,141 @@ And then feed the configuration object into a :class:`ManagedIdentityClient` obj :members: .. automethod:: __init__ + +.. _service-fabric-http-options: + +Service Fabric HTTP options +--------------------------- + +.. autoclass:: msal.ServiceFabricHttpOptions + +Pass ``service_fabric_http_options={}`` to :class:`msal.ManagedIdentityClient` +to use MSAL's isolated, certificate-pinned Service Fabric transport. +The required ``http_client`` remains the transport for other managed identity +providers. On this new Service Fabric path, MSAL does not inspect, open, invoke, +mutate, or close that client. It need not be a Requests session or be opened:: + + import msal + + class UnopenedConsumerClient: + def get(self, url, **kwargs): + raise RuntimeError("Open the consumer transport before non-Service-Fabric use") + + client = msal.ManagedIdentityClient( + msal.SystemAssignedManagedIdentity(), + http_client=UnopenedConsumerClient(), + service_fabric_http_options={}, + ) + # In a Service Fabric environment: + result = client.acquire_token_for_client(resource="https://management.azure.com/") + +All five keys are optional. Unknown keys and explicit ``None`` field values +are invalid; omission is different from ``None``. + +.. list-table:: + :header-rows: 1 + :widths: 15 20 65 + + * - Key + - Default + - Accepted values and behavior + * - ``headers`` + - ``{}`` + - A dict of string HTTP header names and values, overlaid on standard + session headers case-insensitively (last supplied value wins). + ``Secret`` and ``Host`` overrides, in any case, are forbidden. Names + must be ASCII HTTP tokens; values must be transport-encodable, without + CR/LF, invalid control characters, or leading whitespace. Empty values + are allowed. + * - ``proxies`` + - ``{}`` + - A dict of Requests-style selection keys (``http``, ``https``, ``all``, + or scheme plus ``://hostname``) to HTTP/HTTPS proxy URLs. Host-specific + keys cannot contain credentials, ports, paths, queries, or fragments. + URLs require a transport-valid hostname (including IDNA validation) + and valid explicit port, if supplied, and + cannot contain queries, fragments, or non-root paths. Proxy credentials + are allowed. SOCKS and forwarding without endpoint authentication are + unsupported. HTTPS endpoints use CONNECT tunnels and retain pinning. + * - ``trust_env`` + - ``False`` + - A bool. Only ``True`` opts into Requests environment-derived settings, + including proxy selection, ``NO_PROXY``, and netrc authentication. + For HTTPS proxy authentication, ``REQUESTS_CA_BUNDLE`` (or, if unset, + ``CURL_CA_BUNDLE``) selects a CA bundle file or directory. + Requests precedence rules apply; environment proxies can take + precedence over session proxy settings. Neither environment settings + nor explicit proxies can bypass the endpoint pin. + * - ``timeout`` + - ``(5, 30)`` + - Positive finite int/float seconds, or a two-item tuple + ``(connect, read)``. A scalar applies to both waits. Booleans, lists, + zero, negative values, NaN, infinity, and disabled timeouts are invalid. + Limits apply to each wait on every attempt, not total wall-clock time + (DNS resolution and multiple addresses may add elapsed time). + * - ``max_retries`` + - ``0`` + - Non-negative int, excluding bool. Counts additional attempts only for + connection failures before request transmission: at most ``1 + N`` + attempts. No retry for TLS/certificate failures, reads, redirects, + HTTP statuses, or ``Retry-After``; no backoff. + +MSAL snapshots the dictionary and its supported nested dictionaries at +construction. Later mutations have no effect; create a new client to +reconfigure. Validation occurs only when Service Fabric needs network I/O: +not at construction, on a token-cache-only call, or in another environment. +Invalid explicit options raise :class:`msal.ManagedIdentityError` before session +allocation and cannot be hidden by proactive-refresh cached-token fallback. +Selected environment proxies are checked before Requests parses them for +transmission. Malformed or unsupported selected proxies raise +:class:`msal.ManagedIdentityError` without exposing proxy credentials; these +environment failures retain the existing eligible cached-token fallback. + +MSAL owns and closes a separate session and response for each network +acquisition, including its retries. There is no new close API or persistent +pool. Every connection authenticates the actual endpoint certificate against +``IDENTITY_SERVER_THUMBPRINT`` before sending the environment-sourced +``Secret``. Self-signed certificates remain supported by this exact pin; +there is no caller TLS override. Endpoints must use HTTPS. Redirects are never +followed, even on the same origin; all HTTP 300-399 responses raise +:class:`msal.ManagedIdentityError` without exposing the redirect target. +Transport/TLS and existing endpoint/response errors retain their exception +behavior and existing eligible cached-token fallback. + +HTTPS proxies are authenticated independently of the Service Fabric endpoint. +Before sending CONNECT or proxy credentials, MSAL validates the proxy's +certificate chain and hostname against Requests' default CA bundle. +Only ``trust_env=True`` enables the environment CA selection described above, +including for explicitly configured proxies. This trust never replaces the +inner endpoint pin or requires its self-signed certificate to be CA-trusted. +There is no option to disable proxy verification. HTTP CONNECT proxies remain +supported; their proxy credentials travel over plain HTTP, so use them only +on trusted networks. Neither proxy type receives ``Secret`` outside the +endpoint's pinned TLS connection. + +HTTPS proxy tunneling requires urllib3's pre-CONNECT TLS and TLS-in-TLS +capabilities (available in supported urllib3 1.26 and 2.x configurations). +Older stacks such as urllib3 1.25 cannot provide this authenticated route; +MSAL raises :class:`msal.ManagedIdentityError` before connection or credential +transmission, rather than treating an HTTPS proxy as HTTP. This does not raise +the global dependency minimum or disable direct/HTTP-proxy acquisition. +Capability failures retain eligible cached-token fallback. Proxy TLS +validation failures are not retried, just like endpoint pin failures. + +Compatibility and migration +~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +Omitting ``service_fabric_http_options`` or passing top-level ``None`` preserves +the legacy Service Fabric contract: a Requests session with an HTTPAdapter +or subclass is required, and its settings are inherited by MSAL's pinned +transport (not its custom sending behavior). ``None`` does **not** mean ``{}``. +Other providers, token protocol, token cache, and HTTP cache behavior are +unchanged. + +Consumers of the new keyword or public type must require the first MSAL release +that provides it. Older MSAL versions reject the keyword/type with +``TypeError``/``ImportError``. Do not catch these and silently fall back. +Rollback by removing the argument is valid only when the consumer can satisfy +the legacy Requests-session contract. Downgrades must update consumer code +and dependency requirements together; an unopened arbitrary transport is not +a legacy Service Fabric fallback. diff --git a/msal/__init__.py b/msal/__init__.py index ea681317..dd010cfc 100644 --- a/msal/__init__.py +++ b/msal/__init__.py @@ -38,6 +38,7 @@ from .managed_identity import ( SystemAssignedManagedIdentity, UserAssignedManagedIdentity, ManagedIdentityClient, + ServiceFabricHttpOptions, ManagedIdentityError, ArcPlatformNotSupportedError, ) diff --git a/msal/managed_identity.py b/msal/managed_identity.py index 9db2660a..deae0bac 100644 --- a/msal/managed_identity.py +++ b/msal/managed_identity.py @@ -5,20 +5,27 @@ import copy import hashlib import hmac +import ipaddress import json import logging +import math import os +import re import ssl import sys import time import uuid -from urllib.parse import urlparse # Python 3+ +from urllib.parse import urlparse, urlsplit # Python 3+ from collections import UserDict # Python 3+ -from typing import List, Optional, Union # Needed in Python 3.7 & 3.8 +from typing import Dict, List, Optional, Tuple, TypedDict, Union import requests from requests.adapters import HTTPAdapter from urllib3.connection import HTTPSConnection from urllib3.connectionpool import HTTPSConnectionPool +from urllib3.exceptions import MaxRetryError +from urllib3.util import parse_url +from urllib3.util import ssl_ as urllib3_ssl +from urllib3.util.retry import Retry from .token_cache import TokenCache from .individual_cache import _IndividualCache as IndividualCache from .throttled_http_client import ThrottledHttpClientBase, RetryAfterParser @@ -33,6 +40,119 @@ class ManagedIdentityError(ValueError): pass +class ServiceFabricHttpOptions(TypedDict, total=False): + """Optional settings for MSAL's owned, certificate-pinned Service Fabric client. + + See :class:`ManagedIdentityClient` for validation, defaults, and ownership. + No setting can disable endpoint certificate pinning or enable redirects. + """ + headers: Dict[str, str] + proxies: Dict[str, str] + trust_env: bool + timeout: Union[int, float, Tuple[Union[int, float], Union[int, float]]] + max_retries: int + + +class _ServiceFabricOptionsError(ManagedIdentityError): + """A configuration error must not be hidden by proactive-refresh fallback.""" + + +def _snapshot_service_fabric_options(options): + if not isinstance(options, dict): + return options + snapshot = dict(options) + for field in ("headers", "proxies"): + if isinstance(snapshot.get(field), dict): + snapshot[field] = dict(snapshot[field]) + return snapshot + + +def _valid_proxy_host(host): + if not host: + return False + try: + ipaddress.ip_address(host) + return True + except ValueError: + return all(re.fullmatch( + r"[a-zA-Z0-9](?:[a-zA-Z0-9-]{0,61}[a-zA-Z0-9])?", label) + for label in host.encode("idna").decode("ascii").rstrip(".").split(".")) + + +def _valid_service_fabric_proxy(value, *, selection_key=False, allow_default_scheme=False): + if not isinstance(value, str) or not value or any( + character.isspace() or ord(character) < 32 for character in value): + return False + try: + if selection_key: + if value in ("http", "https", "all"): + return True + scheme, separator, host = value.partition("://") + # Requests selection uses urlparse(endpoint).hostname, including + # unbracketed IPv6 literals, rather than a proxy URL's authority. + return bool(separator and scheme in ("http", "https", "all") + and _valid_proxy_host(host)) + if allow_default_scheme: + value = requests.utils.prepend_scheme_if_needed(value, "http") + parsed = urlsplit(value) + if (parsed.scheme not in ("http", "https") + or not _valid_proxy_host(parsed.hostname) + or parsed.query or parsed.fragment or "?" in value or "#" in value + or parsed.path not in ("", "/")): + return False + # The transport's IDNA parser is stricter than Python's builtin codec. + # Consume its potentially credential-bearing errors only in this helper. + parse_url(value) + return (not parsed.netloc.endswith(":") + and (parsed.port is None or 0 < parsed.port <= 65535)) + except (ValueError, UnicodeError): + return False + + +def _validate_service_fabric_options(options): + if not isinstance(options, dict): + raise _ServiceFabricOptionsError("Service Fabric HTTP options must be a dict.") + if options.keys() - ServiceFabricHttpOptions.__annotations__.keys(): + raise _ServiceFabricOptionsError("Unknown Service Fabric HTTP option.") + result = dict(headers={}, proxies={}, trust_env=False, timeout=(5, 30), max_retries=0) + result.update(options) + for field in ("headers", "proxies"): + if not isinstance(result[field], dict) or any( + not isinstance(key, str) or not isinstance(value, str) + for key, value in result[field].items()): + raise _ServiceFabricOptionsError( + "Service Fabric {} must be a dict of strings.".format(field)) + for name, value in result["headers"].items(): + if (not re.fullmatch(r"[!#$%&'*+\-.^_`|~0-9A-Za-z]+", name) + or name.lower() in ("secret", "host") + or (value and value[0].isspace()) + or any(ord(character) < 32 and character != "\t" + or ord(character) == 127 or ord(character) > 255 + for character in value)): + raise _ServiceFabricOptionsError("Invalid or reserved Service Fabric header.") + for key, value in result["proxies"].items(): + if (not _valid_service_fabric_proxy(key, selection_key=True) + or not _valid_service_fabric_proxy(value)): + raise _ServiceFabricOptionsError("Invalid Service Fabric proxies configuration.") + if not isinstance(result["trust_env"], bool): + raise _ServiceFabricOptionsError("Service Fabric trust_env must be a bool.") + retries = result["max_retries"] + if isinstance(retries, bool) or not isinstance(retries, int) or retries < 0: + raise _ServiceFabricOptionsError("Service Fabric max_retries must be a non-negative int.") + timeout = result["timeout"] + waits = timeout if isinstance(timeout, tuple) else (timeout, timeout) + try: + valid_timeout = len(waits) == 2 and all( + isinstance(wait, (int, float)) and not isinstance(wait, bool) + and wait > 0 and math.isfinite(wait) for wait in waits) + except OverflowError: + valid_timeout = False + if not valid_timeout: + raise _ServiceFabricOptionsError( + "Service Fabric timeout must contain positive finite numbers.") + return result + + class ManagedIdentity(UserDict): """Feed an instance of this class to :class:`msal.ManagedIdentityClient` to acquire token for the specified managed identity. @@ -175,6 +295,7 @@ def __init__( token_cache=None, http_cache=None, client_capabilities: Optional[List[str]] = None, + service_fabric_http_options: Optional[ServiceFabricHttpOptions] = None, ): """Create a managed identity client. @@ -197,7 +318,8 @@ def __init__( managed_identity = ... client = msal.ManagedIdentityClient(managed_identity, http_client=s) - For Service Fabric managed identity, ``http_client`` must be a + Unless ``service_fabric_http_options`` is supplied, for Service Fabric + managed identity, ``http_client`` must be a ``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 @@ -205,6 +327,20 @@ def __init__( Standard session, retry, and connection-pool settings are preserved, but custom adapter behavior is not used. + :param service_fabric_http_options: + Optional :class:`msal.ServiceFabricHttpOptions` dictionary. Any dict, + including ``{}``, selects an isolated MSAL-owned Service Fabric + session without inspecting or using ``http_client``. Omission or + ``None`` retains the legacy behavior described above. Other managed + identity providers still use the required ``http_client``. + + Supported dictionaries are snapshotted at construction. Validation + is deferred until Service Fabric needs a network acquisition, not + performed on token-cache hits or in other environments. Invalid + options raise :class:`ManagedIdentityError`, even when a proactive + refresh could otherwise fall back to a cached token. See + :ref:`service-fabric-http-options` for all settings and defaults. + :param token_cache: Optional. It accepts a :class:`msal.TokenCache` instance to store tokens. It will use an in-memory token cache by default. @@ -270,6 +406,8 @@ def __init__( ) self._token_cache = token_cache or TokenCache() self._client_capabilities = client_capabilities + self._service_fabric_http_options = _snapshot_service_fabric_options( + service_fabric_http_options) def acquire_token_for_client( self, @@ -349,6 +487,7 @@ def acquire_token_for_client( access_token_to_refresh.encode("utf-8")).hexdigest() if access_token_to_refresh else None, client_capabilities=self._client_capabilities, + service_fabric_http_options=self._service_fabric_http_options, ) if "access_token" in result: expires_in = result.get("expires_in", 3600) @@ -368,6 +507,8 @@ def acquire_token_for_client( result[self._TOKEN_SOURCE] = self._TOKEN_SOURCE_IDP if (result and "error" not in result) or (not access_token_from_cache): return result + except _ServiceFabricOptionsError: + raise except: # The exact HTTP exception is transportation-layer dependent # Typically network error. Potential AAD outage? if not access_token_from_cache: # It means there is no fall back option @@ -429,6 +570,7 @@ def _obtain_token( *, access_token_sha256_to_refresh: Optional[str] = None, client_capabilities: Optional[List[str]] = None, + service_fabric_http_options=None, ): if ("IDENTITY_ENDPOINT" in os.environ and "IDENTITY_HEADER" in os.environ and "IDENTITY_SERVER_THUMBPRINT" in os.environ @@ -447,6 +589,7 @@ def _obtain_token( resource, access_token_sha256_to_refresh=access_token_sha256_to_refresh, client_capabilities=client_capabilities, + service_fabric_http_options=service_fabric_http_options, ) if "IDENTITY_ENDPOINT" in os.environ and "IDENTITY_HEADER" in os.environ: return _obtain_token_on_app_service( @@ -601,6 +744,7 @@ def _obtain_token_on_service_fabric( *, access_token_sha256_to_refresh: str = None, client_capabilities: Optional[List[str]] = None, + service_fabric_http_options=None, ): """Obtains token for `Service Fabric `_ @@ -609,22 +753,51 @@ def _obtain_token_on_service_fabric( # See also https://github.com/Azure/azure-sdk-for-python/blob/main/sdk/identity/azure-identity/tests/managed-identity-live/service-fabric/service_fabric.md # Protocol https://learn.microsoft.com/en-us/azure/service-fabric/how-to-managed-identity-service-fabric-app-code#acquiring-an-access-token-using-rest-api logger.debug("Obtaining token via managed identity on Azure Service Fabric") + options = (_validate_service_fabric_options(service_fabric_http_options) + if service_fabric_http_options is not None else None) parsed_endpoint = urlparse(endpoint) if parsed_endpoint.scheme.lower() != "https" or not parsed_endpoint.hostname: raise ManagedIdentityError( "Service Fabric managed identity endpoint must use HTTPS.") - service_fabric_http_client = _create_service_fabric_http_client( - http_client, endpoint, _normalize_service_fabric_thumbprint(server_thumbprint)) - resp = service_fabric_http_client.get( - endpoint, - params={k: v for k, v in { + thumbprint = _normalize_service_fabric_thumbprint(server_thumbprint) + params = {k: v for k, v in { "api-version": "2019-07-01-preview", "resource": resource, "token_sha256_to_refresh": access_token_sha256_to_refresh, "xms_cc": ",".join(client_capabilities) if client_capabilities else None, - }.items() if v is not None}, - headers={"Secret": identity_header}, - ) + }.items() if v is not None} + if options is None: + service_fabric_http_client = _create_service_fabric_http_client( + http_client, endpoint, thumbprint) + return _parse_service_fabric_response(service_fabric_http_client.get( + endpoint, params=params, headers={"Secret": identity_header})) + + service_fabric_http_client = _create_owned_service_fabric_http_client(options, thumbprint) + response = None + + def receive_response(raw_response, **ignored): + nonlocal response + # Capture ownership before Requests processes redirects or reads content. + response = raw_response + if 300 <= response.status_code < 400: + raise ManagedIdentityError( + "Service Fabric returned a redirect (HTTP {}).".format(response.status_code)) + + try: + service_fabric_http_client.get( + endpoint, params=params, headers={"Secret": identity_header}, + timeout=options["timeout"], allow_redirects=False, stream=True, + hooks={"response": receive_response}) + return _parse_service_fabric_response(response) + finally: + try: + if response is not None: + response.close() + finally: + service_fabric_http_client.close() + + +def _parse_service_fabric_response(resp): try: payload = json.loads(resp.text) if payload.get("access_token") and payload.get("expires_on"): @@ -688,13 +861,40 @@ class _ServiceFabricHTTPSConnectionPool(HTTPSConnectionPool): ConnectionCls = _ServiceFabricHTTPSConnection +class _ServiceFabricAuthenticatedProxyConnection(_ServiceFabricHTTPSConnection): + _proxy_ssl_context = None + + def _connect_tls_proxy(self, hostname, sock): + # This hook runs before CONNECT. The endpoint's CERT_NONE must never + # configure the outer TLS handshake, which authenticates proxy credentials. + hostname = hostname.strip("[]").split("%", 1)[0].rstrip(".") + proxy_socket = self._proxy_ssl_context.wrap_socket(sock, server_hostname=hostname) + self.proxy_is_verified = True + return proxy_socket + + +class _ServiceFabricConnectionRetry(Retry): + """Exclude non-connection retries even on urllib3 without an 'other' limit.""" + + def increment(self, method=None, url=None, response=None, error=None, + _pool=None, _stacktrace=None): + if (error is not None and not self._is_connection_error(error) + and not self._is_read_error(error)): + raise MaxRetryError(_pool, url, error) from error + return super(_ServiceFabricConnectionRetry, self).increment( + method=method, url=url, response=response, error=error, + _pool=_pool, _stacktrace=_stacktrace) + + class _ServiceFabricHTTPAdapter(HTTPAdapter): """Use certificate-thumbprint authentication for the Service Fabric endpoint.""" - def __init__(self, server_thumbprint, *args, **kwargs): + def __init__(self, server_thumbprint, *args, _restrict_proxies=False, **kwargs): + self._restrict_proxies = _restrict_proxies connection_class = type( "_PinnedServiceFabricHTTPSConnection", - (_ServiceFabricHTTPSConnection,), + (_ServiceFabricAuthenticatedProxyConnection if _restrict_proxies + else _ServiceFabricHTTPSConnection,), {"_server_thumbprint": server_thumbprint}, ) self._connection_pool_class = type( @@ -716,18 +916,65 @@ def init_poolmanager(self, connections, maxsize, block=False, **pool_kwargs): self._configure_pool_manager(self.poolmanager) def proxy_manager_for(self, proxy, **proxy_kwargs): + if self._restrict_proxies and not _valid_service_fabric_proxy(proxy): + raise ManagedIdentityError("Unsupported Service Fabric proxy configuration.") pool_manager = super(_ServiceFabricHTTPAdapter, self).proxy_manager_for( proxy, **proxy_kwargs) self._configure_pool_manager(pool_manager) return pool_manager + def send(self, request, stream=False, timeout=None, verify=True, cert=None, proxies=None): + if self._restrict_proxies: + proxy = requests.utils.select_proxy(request.url, proxies) + # Requests parses selected proxies before calling proxy_manager_for. + if proxy and not _valid_service_fabric_proxy(proxy, allow_default_scheme=True): + raise ManagedIdentityError("Unsupported Service Fabric proxy configuration.") + if proxy and urlsplit(proxy).scheme.lower() == "https": + # Older urllib3 accepts arbitrary pool kwargs without supporting + # HTTPS tunneling. Require the actual pre-CONNECT TLS hook and + # TLS-in-TLS implementation instead of passing ignored options. + if (not callable(getattr(HTTPSConnection, "_connect_tls_proxy", None)) + or not callable(getattr(urllib3_ssl, "SSLTransport", None)) + or not callable(getattr(urllib3_ssl.SSLContext, "wrap_bio", None))): + raise ManagedIdentityError( + "Installed urllib3 does not support authenticated HTTPS proxy tunneling.") + ca_bundle = verify if isinstance(verify, str) else ( + requests.utils.extract_zipped_paths(requests.adapters.DEFAULT_CA_BUNDLE_PATH)) + context = ssl.create_default_context( + capath=ca_bundle if os.path.isdir(ca_bundle) else None, + cafile=None if os.path.isdir(ca_bundle) else ca_bundle) + self._connection_pool_class.ConnectionCls._proxy_ssl_context = context + return super(_ServiceFabricHTTPAdapter, self).send( + request, stream=stream, timeout=timeout, verify=verify, cert=cert, proxies=proxies) + def cert_verify(self, conn, url, verify, cert): # The exact Service Fabric certificate thumbprint is the trust anchor. - # Do not inherit caller-provided verify=False or a custom CA configuration. + # CA configuration applies only to the owned transport's outer proxy TLS. super(_ServiceFabricHTTPAdapter, self).cert_verify( conn, url, verify=False, cert=cert) +def _create_owned_service_fabric_http_client(options, server_thumbprint): + session = requests.Session() + try: + session.headers.update(options["headers"]) + session.proxies.update(options["proxies"]) + session.trust_env = options["trust_env"] + adapter = _ServiceFabricHTTPAdapter( + server_thumbprint, _restrict_proxies=True, + max_retries=_ServiceFabricConnectionRetry( + total=options["max_retries"], connect=options["max_retries"], + read=False, redirect=0, status=0, + backoff_factor=0, respect_retry_after_header=False)) + default_adapter = session.adapters["https://"] + session.mount("https://", adapter) + default_adapter.close() + except BaseException: + session.close() + raise + return session + + def _create_service_fabric_http_client(http_client, endpoint, server_thumbprint): """Derive a Requests session with a pinning-only HTTPS transport. diff --git a/tests/test_mi.py b/tests/test_mi.py index c7c1fef6..e343b8b3 100644 --- a/tests/test_mi.py +++ b/tests/test_mi.py @@ -1,13 +1,19 @@ import hashlib import json import os +import select +import socket import ssl import sys -import tempfile import threading import time +import traceback import uuid from datetime import datetime, timedelta, timezone +from collections import UserDict +from concurrent.futures import ThreadPoolExecutor +from contextlib import nullcontext +from urllib.parse import parse_qs, urlparse from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer from typing import List, Optional import unittest @@ -22,7 +28,7 @@ from cryptography import x509 from cryptography.hazmat.primitives import hashes, serialization from cryptography.hazmat.primitives.asymmetric import rsa -from cryptography.x509.oid import NameOID +from cryptography.x509.oid import ExtendedKeyUsageOID, NameOID from tests.test_throttled_http_client import ( MinimalResponse, ThrottledHttpClientBaseTestCase, DummyHttpClient) @@ -30,6 +36,7 @@ SystemAssignedManagedIdentity, UserAssignedManagedIdentity, ManagedIdentityClient, ManagedIdentityError, + ServiceFabricHttpOptions, ArcPlatformNotSupportedError, ) from msal.managed_identity import ( @@ -44,6 +51,9 @@ DEFAULT_TO_VM, _create_service_fabric_http_client, _obtain_token_on_service_fabric, + _ServiceFabricHTTPAdapter, + _ServiceFabricHTTPSConnection, + _create_owned_service_fabric_http_client, ) from msal.token_cache import is_subdict_of @@ -451,27 +461,33 @@ def do_GET(self): "path": self.path, "headers": dict(self.headers), }) - body = json.dumps({ + body = getattr(self.server, "body", json.dumps({ "access_token": "AT", "expires_on": str(int(time.time()) + 3600), "resource": "R", "token_type": "Bearer", - }).encode("utf-8") - self.send_response(200) + })).encode("utf-8") + self.send_response(getattr(self.server, "status", 200)) + for name, value in getattr(self.server, "response_headers", {}).items(): + self.send_header(name, value) self.send_header("Content-Type", "application/json") self.send_header("Content-Length", str(len(body))) self.end_headers() - self.wfile.write(body) + if getattr(self.server, "stall", None): + self.server.stall.wait(2) + try: + self.wfile.write(body) + except (OSError, ssl.SSLError): + pass def log_message(self, format, *args): pass -class ServiceFabricTlsValidationTestCase(unittest.TestCase): +class _ServiceFabricTlsFixture(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) certificate = ( x509.CertificateBuilder() @@ -493,8 +509,11 @@ def setUp(self): .sign(private_key, hashes.SHA256()) ) self.thumbprint = certificate.fingerprint(hashes.SHA1()).hex() - certificate_path = os.path.join(self._temporary_directory.name, "server.pem") - private_key_path = os.path.join(self._temporary_directory.name, "server.key") + file_prefix = ".service-fabric-test-" + uuid.uuid4().hex + certificate_path = file_prefix + ".pem" + private_key_path = file_prefix + ".key" + self.addCleanup(lambda: os.path.exists(certificate_path) and os.remove(certificate_path)) + self.addCleanup(lambda: os.path.exists(private_key_path) and os.remove(private_key_path)) with open(certificate_path, "wb") as certificate_file: certificate_file.write(certificate.public_bytes(serialization.Encoding.PEM)) with open(private_key_path, "wb") as private_key_file: @@ -509,8 +528,12 @@ def setUp(self): tls_context = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER) tls_context.minimum_version = ssl.TLSVersion.TLSv1_2 tls_context.load_cert_chain(certificate_path, private_key_path) + os.remove(certificate_path) + os.remove(private_key_path) + self.tls_context = tls_context self.server.socket = tls_context.wrap_socket(self.server.socket, server_side=True) - self.server_thread = threading.Thread(target=self.server.serve_forever) + self.server_thread = threading.Thread( + target=lambda: self.server.serve_forever(poll_interval=0.01)) self.server_thread.start() self.endpoint = "https://localhost:{}/token".format(self.server.server_port) @@ -518,7 +541,6 @@ def tearDown(self): self.server.shutdown() self.server.server_close() self.server_thread.join() - self._temporary_directory.cleanup() def _new_session(self, **adapter_kwargs): session = requests.Session() @@ -527,6 +549,8 @@ def _new_session(self, **adapter_kwargs): session.mount("https://", self._adapter_class(**adapter_kwargs)) return session + +class ServiceFabricTlsValidationTestCase(_ServiceFabricTlsFixture): def test_matching_thumbprint_sends_secret_after_validating_certificate(self): result = _obtain_token_on_service_fabric( _ThrottledHttpClient(self._new_session()), @@ -662,6 +686,1067 @@ class ServiceFabricHTTPAdapterSubclassTestCase(ServiceFabricTlsValidationTestCas _adapter_class = _ServiceFabricSourceHTTPAdapter +class _UnopenedHttpClient: + def __getattribute__(self, name): + if name == "__class__": + return type(self) + raise AssertionError("Consumer transport must remain unopened and uninspected") + + def __setattr__(self, name, value): + raise AssertionError("Consumer transport must not be mutated") + + +class _ConnectProxyHandler(BaseHTTPRequestHandler): + def do_GET(self): + self.server.requests.append(dict(self.headers)) + self.send_response(200) + self.send_header("Content-Length", "0") + self.end_headers() + + def do_CONNECT(self): + self.server.connects.append((self.path, dict(self.headers))) + with socket.create_connection(self.server.target, timeout=2) as upstream: + self.send_response(200) + self.end_headers() + self.connection.settimeout(2) + while True: + ready, _, _ = select.select([self.connection, upstream], [], [], 2) + if not ready: + return + for source in ready: + try: + data = source.recv(65536) + if not data: + return + if source is self.connection: + self.server.tunnel_data.append(data) + (upstream if source is self.connection else self.connection).sendall(data) + except OSError: + return + + def log_message(self, format, *args): + pass + + +class ServiceFabricHttpOptionsTestCase(_ServiceFabricTlsFixture): + def setUp(self): + super().setUp() + env = patch.dict(os.environ, { + "IDENTITY_ENDPOINT": self.endpoint, + "IDENTITY_HEADER": "service-fabric-secret", + "IDENTITY_SERVER_THUMBPRINT": self.thumbprint, + }, clear=True) + env.start() + self.addCleanup(env.stop) + + def _app(self, options=None, **kwargs): + return ManagedIdentityClient( + SystemAssignedManagedIdentity(), http_client=_UnopenedHttpClient(), + service_fabric_http_options={} if options is None else options, **kwargs) + + def _seed_cache(self, app, *, refresh=False): + app._token_cache.add({ + "client_id": None, "scope": ["R"], + "token_endpoint": "https://localhost/managed_identity", + "response": {"access_token": "cached", "expires_in": 3600, + "token_type": "Bearer", **({"refresh_in": -1} if refresh else {})}, + }) + + def _proxy(self, *, tls=False, target=None, hostname="localhost"): + server = ThreadingHTTPServer(("localhost", 0), _ConnectProxyHandler) + server.target = target or ("localhost", self.server.server_port) + server.connects, server.tunnel_data, server.requests = [], [], [] + if tls: + ca_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + ca_name = x509.Name([x509.NameAttribute(NameOID.COMMON_NAME, "Test proxy CA")]) + ca = (x509.CertificateBuilder().subject_name(ca_name).issuer_name(ca_name) + .public_key(ca_key.public_key()).serial_number(x509.random_serial_number()) + .not_valid_before(datetime.now(timezone.utc) - timedelta(minutes=1)) + .not_valid_after(datetime.now(timezone.utc) + timedelta(days=1)) + .add_extension(x509.BasicConstraints(ca=True, path_length=0), critical=True) + .add_extension(x509.KeyUsage( + digital_signature=False, content_commitment=False, + key_encipherment=False, data_encipherment=False, + key_agreement=False, key_cert_sign=True, crl_sign=True, + encipher_only=False, decipher_only=False), critical=True) + .add_extension(x509.SubjectKeyIdentifier.from_public_key( + ca_key.public_key()), critical=False) + .add_extension(x509.AuthorityKeyIdentifier.from_issuer_public_key( + ca_key.public_key()), critical=False) + .sign(ca_key, hashes.SHA256())) + key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + certificate = (x509.CertificateBuilder() + .subject_name(x509.Name([x509.NameAttribute(NameOID.COMMON_NAME, hostname)])) + .issuer_name(ca_name).public_key(key.public_key()) + .serial_number(x509.random_serial_number()) + .not_valid_before(datetime.now(timezone.utc) - timedelta(minutes=1)) + .not_valid_after(datetime.now(timezone.utc) + timedelta(days=1)) + .add_extension(x509.SubjectAlternativeName([x509.DNSName(hostname)]), critical=False) + .add_extension(x509.BasicConstraints(ca=False, path_length=None), critical=True) + .add_extension(x509.KeyUsage( + digital_signature=True, content_commitment=False, + key_encipherment=True, data_encipherment=False, + key_agreement=False, key_cert_sign=False, crl_sign=False, + encipher_only=False, decipher_only=False), critical=True) + .add_extension(x509.ExtendedKeyUsage( + [ExtendedKeyUsageOID.SERVER_AUTH]), critical=False) + .add_extension(x509.SubjectKeyIdentifier.from_public_key( + key.public_key()), critical=False) + .add_extension(x509.AuthorityKeyIdentifier.from_issuer_public_key( + ca_key.public_key()), critical=False) + .sign(ca_key, hashes.SHA256())) + prefix = ".service-fabric-proxy-test-" + uuid.uuid4().hex + server.ca_path = os.path.abspath(prefix + "-ca.pem") + server.thumbprint = certificate.fingerprint(hashes.SHA1()).hex() + certificate_path, key_path = prefix + ".pem", prefix + ".key" + for path, data in ( + (server.ca_path, ca.public_bytes(serialization.Encoding.PEM)), + (certificate_path, certificate.public_bytes(serialization.Encoding.PEM)), + (key_path, key.private_bytes(serialization.Encoding.PEM, + serialization.PrivateFormat.TraditionalOpenSSL, serialization.NoEncryption()))): + self.addCleanup(lambda path=path: os.path.exists(path) and os.remove(path)) + with open(path, "wb") as output: + output.write(data) + context = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER) + context.load_cert_chain(certificate_path, key_path) + os.remove(certificate_path) + os.remove(key_path) + server.socket = context.wrap_socket(server.socket, server_side=True) + thread = threading.Thread(target=lambda: server.serve_forever(poll_interval=0.01)) + thread.start() + def close(): + server.shutdown() + server.server_close() + thread.join() + self.addCleanup(close) + return server, "{}://localhost:{}".format("https" if tls else "http", server.server_port) + + def test_public_type_and_unopened_consumer_journey(self): + self.assertEqual( + {"headers", "proxies", "trust_env", "timeout", "max_retries"}, + set(ServiceFabricHttpOptions.__annotations__)) + self.assertFalse(ServiceFabricHttpOptions.__total__) + self.assertEqual(frozenset(), ServiceFabricHttpOptions.__required_keys__) + http_cache = {} + app = self._app(ServiceFabricHttpOptions(), http_cache=http_cache, + client_capabilities=["CP1", "CP2"]) + token = app.acquire_token_for_client(resource="R") + self.assertEqual(("AT", "Bearer", "identity_provider"), + (token["access_token"], token["token_type"], token["token_source"])) + self.assertTrue(3595 <= token["expires_in"] <= 3600) + with patch("msal.managed_identity.requests.Session") as allocation: + self.assertEqual("cache", app.acquire_token_for_client(resource="R")["token_source"]) + allocation.assert_not_called() + app.acquire_token_for_client(resource="R", claims_challenge="challenge") + self.assertEqual(2, len(self.server.requests)) + params = parse_qs(urlparse(self.server.requests[-1]["path"]).query) + self.assertEqual({ + "resource": ["R"], "api-version": ["2019-07-01-preview"], + "xms_cc": ["CP1,CP2"], + "token_sha256_to_refresh": [hashlib.sha256(b"AT").hexdigest()], + }, params) + self.assertEqual({}, http_cache) + self.assertEqual("service-fabric-secret", self.server.requests[0]["headers"]["Secret"]) + + def test_defaults_partial_fields_and_full_configuration(self): + explicit = {"headers": {"User-Agent": "custom"}, "proxies": {}, + "trust_env": True, "timeout": (1.5, 2), "max_retries": 2} + cases = [{}] + [{key: value} for key, value in explicit.items()] + cases += [explicit] + [ + {key: value for key, value in explicit.items() if key != omitted} + for omitted in explicit] + from msal.managed_identity import _validate_service_fabric_options + for options in cases: + with self.subTest(options=options): + normalized = _validate_service_fabric_options(options) + session = _create_owned_service_fabric_http_client(normalized, self.thumbprint) + try: + self.assertEqual(options.get("proxies", {}), session.proxies) + self.assertIs(options.get("trust_env", False), session.trust_env) + self.assertEqual(options.get("timeout", (5, 30)), normalized["timeout"]) + self.assertIn("Accept", session.headers) + self.assertEqual(options.get("headers", {}).get( + "User-Agent", requests.utils.default_user_agent()), session.headers["User-Agent"]) + retry = session.get_adapter(self.endpoint).max_retries + self.assertEqual(options.get("max_retries", 0), retry.total) + self.assertEqual(retry.total, retry.connect) + self.assertFalse(retry.read) + self.assertEqual((0, 0, 0), ( + retry.status, retry.redirect, retry.backoff_factor)) + self.assertFalse(retry.respect_retry_after_header) + finally: + session.close() + + def test_snapshot_and_case_insensitive_header_overlay(self): + proxy, address = self._proxy() + headers = {"X-Label": "first", "x-label": "last", "User-Agent": "custom", "X-Empty": ""} + proxies = {"https": address} + options = {"headers": headers, "proxies": proxies, "timeout": 2} + app = self._app(options) + headers["x-label"], proxies["https"], options["timeout"] = "changed", "invalid", None + options["unknown"] = "changed" + app.acquire_token_for_client(resource="R") + received = requests.structures.CaseInsensitiveDict(self.server.requests[0]["headers"]) + self.assertEqual(("last", "custom", "", "service-fabric-secret"), ( + received["X-Label"], received["User-Agent"], received["X-Empty"], received["Secret"])) + self.assertIn("Accept", received) + self.assertEqual(1, len(proxy.connects)) + self.assertEqual("changed", headers["x-label"]) + self.assertEqual("invalid", proxies["https"]) + + def test_invalid_options_are_lazy_private_and_preallocation(self): + private = "sensitive-option-value" + cases = [False, 0, "", [], (), UserDict(), {private: private}] + cases += [{key: None} for key in ServiceFabricHttpOptions.__annotations__] + cases += [{"headers": value} for value in [ + [], UserDict(), {1: private}, {"X-Test": 1}, {"sEcReT": private}, + {"HOST": private}, {"": private}, {"Bad Name": private}, {"x:bad": private}, + {"é": private}, {"X-Test": "a\r\n" + private}, {"X-Test": "\n"}, + {"X-Test": "\x00"}, {"X-Test": "\x7f"}, {"X-Test": "☃"}, + {"X-Test": " " + private}, {"X-Test": "\t" + private}]] + cases += [{"proxies": value} for value in [ + [], UserDict(), {1: private}, {"https": 123}]] + cases += [{"trust_env": value} for value in [0, 1, "", "true", []]] + cases += [{"max_retries": value} for value in [True, False, -1, 1.5, "1", []]] + cases += [{"timeout": value} for value in [ + True, False, 0, -1, float("nan"), float("inf"), float("-inf"), + "2", [1, 2], (), (1,), (1, 2, 3), (1, None), (True, 1), + (1, float("nan")), 10 ** 1000]] + for options in cases: + with self.subTest(options=options), patch( + "msal.managed_identity.requests.Session") as allocation: + app = self._app(options) + with self.assertRaises(ManagedIdentityError) as error: + app.acquire_token_for_client(resource="R") + self.assertNotIn(private, str(error.exception)) + allocation.assert_not_called() + self.assertEqual([], self.server.requests) + + def test_proxy_validation_matrix(self): + invalid_keys = ["", "ftp", "https://", "https://a:80", "https://user:password@a", + "https://a/path", "https://a/", "https://a?", "https://a#", + "https://a:", "https://a b", "https://a\\b", "https://[invalid]"] + invalid_urls = ["", "localhost:8080", "socks5://localhost:1080", "ftp://a", + "http://", "https://a:0", "https://a:65536", "https://a:invalid", + "https://a:", "https://a/path", "https://a?q", "https://a#f", + "https://a?", "https://a#", "https://a b", "https://[invalid]", + "https://a/;", "https://a/;params"] + for proxies in ([{key: "http://localhost"} for key in invalid_keys] + + [{"https": value} for value in invalid_urls]): + with self.subTest(proxies=proxies), patch( + "msal.managed_identity.requests.Session") as allocation: + with self.assertRaises(ManagedIdentityError): + self._app({"proxies": proxies}).acquire_token_for_client(resource="R") + allocation.assert_not_called() + + def test_transport_invalid_proxy_idna_is_private_and_preallocation(self): + private = "synthetic-proxy-password" + for scheme in ("http", "https"): + for cached in (False, True): + with self.subTest(scheme=scheme, cached=cached), patch( + "msal.managed_identity.requests.Session", wraps=requests.Session + ) as allocation, patch.object( + _ServiceFabricHTTPSConnection, "_new_conn") as connect, self.assertLogs( + level="DEBUG") as logs: + app = self._app({"proxies": { + "https": "{}://user:{}@\u2603.example".format(scheme, private)}}) + if cached: + self._seed_cache(app) + self.assertEqual("cached", app.acquire_token_for_client( + resource="R")["access_token"]) + self._seed_cache(app, refresh=True) + with self.assertRaisesRegex( + ManagedIdentityError, "^Invalid Service Fabric proxies configuration\\.$") as error: + app.acquire_token_for_client(resource="R") + self.assertIsNone(error.exception.__cause__) + self.assertIsNone(error.exception.__context__) + self.assertNotIn(private, "".join(traceback.format_exception( + type(error.exception), error.exception, error.exception.__traceback__))) + allocation.assert_not_called() + connect.assert_not_called() + self.assertNotIn(private, "\n".join(logs.output)) + self.assertEqual([], self.server.requests) + + def test_validation_is_skipped_on_cache_hit_but_not_refresh_fallback(self): + for refresh in (False, True): + app = self._app({"timeout": None}) + self._seed_cache(app, refresh=refresh) + with patch("msal.managed_identity.requests.Session") as allocation: + if refresh: + with self.assertRaises(ManagedIdentityError): + app.acquire_token_for_client(resource="R") + else: + self.assertEqual("cached", app.acquire_token_for_client(resource="R")["access_token"]) + allocation.assert_not_called() + + def test_matching_pin_normalization(self): + for pin in (self.thumbprint, self.thumbprint.upper(), + " \t" + ":".join(self.thumbprint[i:i+2].upper() + for i in range(0, 40, 2)) + "\r\n"): + with self.subTest(pin=pin), patch.dict(os.environ, {"IDENTITY_SERVER_THUMBPRINT": pin}): + self.assertEqual("AT", self._app().acquire_token_for_client(resource="R")["access_token"]) + self.assertEqual(3, len(self.server.requests)) + + def test_bad_pin_and_endpoint_never_send_or_retry_secret(self): + for pin, endpoint, error in [ + ("00" * 20, self.endpoint, SSLError), + ("bad", self.endpoint, ManagedIdentityError), + ("", self.endpoint, ManagedIdentityError), + (self.thumbprint, self.endpoint.replace("https:", "http:"), ManagedIdentityError)]: + with self.subTest(pin=pin), patch.dict(os.environ, { + "IDENTITY_SERVER_THUMBPRINT": pin, "IDENTITY_ENDPOINT": endpoint}), patch.object( + _ServiceFabricHTTPSConnection, "connect", autospec=True, + side_effect=_ServiceFabricHTTPSConnection.connect) as connect: + with self.assertRaises(error): + self._app({"max_retries": 2}).acquire_token_for_client(resource="R") + self.assertEqual(1 if error is SSLError else 0, connect.call_count) + self.assertEqual([], self.server.requests) + + def test_redirect_status_and_destination_matrix_is_not_followed(self): + proxy, target = self._proxy() + other = _ServiceFabricTlsFixture() + other.setUp() + self.addCleanup(other.doCleanups) + self.addCleanup(other.tearDown) + for status in range(300, 400): + for destination in (self.endpoint + "/next", + other.endpoint + "/private-target", + target + "/private-target"): + with self.subTest(status=status, destination=destination): + self.server.status = status + self.server.response_headers = {"Location": destination} + with self.assertRaisesRegex(ManagedIdentityError, str(status)) as error: + self._app({"max_retries": 2}).acquire_token_for_client(resource="R") + self.assertNotIn("private-target", str(error.exception)) + self.assertEqual(300, len(self.server.requests)) + self.assertEqual([], proxy.connects) + self.assertEqual([], proxy.requests) + self.assertEqual([], other.server.requests) + + def test_explicit_http_and_https_connect_proxy_pins_endpoint(self): + for tls in (False, True): + proxy, address = self._proxy(tls=tls) + address = address.replace("://", "://proxy-user:proxy-password@") + with patch("requests.adapters.DEFAULT_CA_BUNDLE_PATH", proxy.ca_path) if tls else nullcontext(): + for key in ("https", "all", "https://localhost", "all://localhost"): + with self.subTest(tls=tls, key=key): + self.assertEqual("AT", self._app({"proxies": {key: address}} + ).acquire_token_for_client(resource="R")["access_token"]) + before = len(self.server.requests) + with patch.dict(os.environ, { + "IDENTITY_SERVER_THUMBPRINT": proxy.thumbprint if tls else "00" * 20}): + with self.assertRaises((SSLError, requests.exceptions.ProxyError)): + self._app({"proxies": {"https": address}, "max_retries": 2} + ).acquire_token_for_client(resource="R") + self.assertEqual(before, len(self.server.requests)) + self.assertEqual(5, len(proxy.connects)) + self.assertTrue(all(headers["Proxy-Authorization"] == requests.auth._basic_auth_str( + "proxy-user", "proxy-password") for _, headers in proxy.connects)) + self.assertNotIn(b"service-fabric-secret", b"".join(proxy.tunnel_data)) + self.assertTrue(all("Secret" not in headers for _, headers in proxy.connects)) + self.assertTrue(all("Proxy-Authorization" not in request["headers"] + for request in self.server.requests)) + + def test_https_proxy_strict_verification_accepts_valid_chain_and_checks_hostname(self): + create_default_context = ssl.create_default_context + contexts = [] + + def strict_context(*args, **kwargs): + context = create_default_context(*args, **kwargs) + context.verify_flags |= ssl.VERIFY_X509_STRICT + contexts.append(context) + return context + + for hostname in ("localhost", "wrong.invalid"): + with self.subTest(hostname=hostname): + proxy, address = self._proxy(tls=True, hostname=hostname) + address = address.replace("://", "://proxy-user:proxy-password@") + context = strict_context(cafile=proxy.ca_path) + # Verify even the hostname-negative fixture's chain with its own + # valid name before testing the actual proxy hostname below. + with socket.create_connection(("localhost", proxy.server_port), timeout=2) as sock: + with context.wrap_socket(sock, server_hostname=hostname) as tls_socket: + self.assertEqual(proxy.thumbprint, + hashlib.sha1(tls_socket.getpeercert(binary_form=True)).hexdigest()) + self.assertEqual([], proxy.connects) + self.assertEqual([], proxy.requests) + before = len(self.server.requests) + with patch("requests.adapters.DEFAULT_CA_BUNDLE_PATH", proxy.ca_path), patch( + "msal.managed_identity.ssl.create_default_context", + side_effect=strict_context) as create_context: + app = self._app({"proxies": {"https": address}, "max_retries": 2}) + if hostname == "localhost": + self.assertEqual("AT", app.acquire_token_for_client( + resource="R")["access_token"]) + else: + with self.assertRaises((SSLError, requests.exceptions.ProxyError)) as error: + app.acquire_token_for_client(resource="R") + message = "".join(traceback.format_exception( + type(error.exception), error.exception, error.exception.__traceback__)) + self.assertIn("CERTIFICATE_VERIFY_FAILED", message) + self.assertRegex(message, "[Hh]ostname mismatch") + self.assertIn("certificate is not valid for 'localhost'", message) + create_context.assert_called_once_with(capath=None, cafile=proxy.ca_path) + self.assertTrue(contexts[-1].verify_flags & ssl.VERIFY_X509_STRICT) + self.assertEqual(ssl.CERT_REQUIRED, contexts[-1].verify_mode) + self.assertTrue(contexts[-1].check_hostname) + expected_requests = int(hostname == "localhost") + self.assertEqual(before + expected_requests, len(self.server.requests)) + self.assertEqual(expected_requests, len(proxy.connects)) + self.assertEqual([], proxy.requests) + if hostname == "localhost": + self.assertEqual(requests.auth._basic_auth_str("proxy-user", "proxy-password"), + proxy.connects[0][1]["Proxy-Authorization"]) + self.assertNotIn("Secret", proxy.connects[0][1]) + self.assertEqual("service-fabric-secret", + self.server.requests[-1]["headers"]["Secret"]) + self.assertNotIn("Proxy-Authorization", self.server.requests[-1]["headers"]) + else: + self.assertEqual([], proxy.tunnel_data) + + def test_https_proxy_rejects_untrusted_chain_and_wrong_hostname_before_connect(self): + factory = requests.Session + private = "synthetic-https-proxy-password" + authorization = requests.auth._basic_auth_str("proxy-user", private) + for failure in ("untrusted", "hostname"): + proxy, address = self._proxy(tls=True, + hostname="wrong.invalid" if failure == "hostname" else "localhost") + address = address.replace("://", "://proxy-user:{}@".format(private)) + for environment in (False, True): + for cached in (False, True): + session = factory() + self.addCleanup(session.close) + trust = (patch("requests.adapters.DEFAULT_CA_BUNDLE_PATH", proxy.ca_path) + if failure == "hostname" else nullcontext()) + with self.subTest(failure=failure, environment=environment, cached=cached), trust, patch.dict( + os.environ, {"HTTPS_PROXY": address} if environment else {}), patch( + "msal.managed_identity.requests.Session", return_value=session), patch.object( + session, "close", wraps=session.close) as close, patch.object( + _ServiceFabricHTTPSConnection, "connect", autospec=True, + side_effect=_ServiceFabricHTTPSConnection.connect) as connect, patch( + "urllib3.util.retry.time.sleep") as sleep, self.assertLogs(level="DEBUG") as logs: + app = self._app({"trust_env": environment, "max_retries": 2, + "proxies": {} if environment else {"https": address}}) + if cached: + self._seed_cache(app, refresh=True) + self.assertEqual("cached", app.acquire_token_for_client(resource="R")["access_token"]) + else: + with self.assertRaises((SSLError, requests.exceptions.ProxyError)) as error: + app.acquire_token_for_client(resource="R") + message = "".join(traceback.format_exception( + type(error.exception), error.exception, error.exception.__traceback__)) + self.assertIn("CERTIFICATE_VERIFY_FAILED", message) + for secret in (private, authorization, "service-fabric-secret"): + self.assertNotIn(secret, message) + self.assertEqual(1, connect.call_count) + sleep.assert_not_called() + close.assert_called_once_with() + adapter = session.get_adapter(self.endpoint) + self.assertTrue(all(len(manager.pools) == 0 + for manager in adapter.proxy_manager.values())) + for secret in (private, authorization, "service-fabric-secret"): + self.assertNotIn(secret, "\n".join(logs.output)) + self.assertEqual([], proxy.connects) + self.assertEqual([], proxy.requests) + self.assertEqual([], proxy.tunnel_data) + self.assertEqual([], self.server.requests) + + def test_https_proxy_ca_environment_selection_does_not_change_endpoint_pin(self): + proxy, address = self._proxy(tls=True) + address = address.replace("://", "://proxy-user:proxy-password@") + for route in ("explicit", "HTTPS_PROXY", "ALL_PROXY"): + for bundle in ("REQUESTS_CA_BUNDLE", "CURL_CA_BUNDLE"): + for trust_env in (False, True): + environment = {bundle: proxy.ca_path} + if bundle == "REQUESTS_CA_BUNDLE": + environment["CURL_CA_BUNDLE"] = "unused-lower-priority-bundle" + if route != "explicit": + environment[route] = address + options = {"trust_env": trust_env, + "proxies": {"https": address} if route == "explicit" else {}} + before = len(proxy.connects) + with self.subTest(route=route, bundle=bundle, trust_env=trust_env), patch.dict( + os.environ, environment), self.assertLogs(level="DEBUG") as logs: + if route == "explicit" and not trust_env: + with self.assertRaises((SSLError, requests.exceptions.ProxyError)): + self._app(options).acquire_token_for_client(resource="R") + self.assertEqual(before, len(proxy.connects)) + else: + self.assertEqual("AT", self._app(options).acquire_token_for_client( + resource="R")["access_token"]) + self.assertEqual(before + int(trust_env), len(proxy.connects)) + if trust_env: + for cached in (False, True): + with patch.dict(os.environ, {"IDENTITY_SERVER_THUMBPRINT": proxy.thumbprint}): + app = self._app(dict(options, max_retries=2)) + if cached: + self._seed_cache(app, refresh=True) + self.assertEqual("cached", app.acquire_token_for_client( + resource="R")["access_token"]) + else: + with self.assertRaises((SSLError, requests.exceptions.ProxyError)): + app.acquire_token_for_client(resource="R") + self.assertEqual(before + 3, len(proxy.connects)) + for secret in ("proxy-password", "service-fabric-secret", + requests.auth._basic_auth_str("proxy-user", "proxy-password")): + self.assertNotIn(secret, "\n".join(logs.output)) + self.assertEqual(10, len(self.server.requests)) + self.assertTrue(all(request["headers"]["Secret"] == "service-fabric-secret" + and "Proxy-Authorization" not in request["headers"] for request in self.server.requests)) + self.assertTrue(all("Secret" not in headers for _, headers in proxy.connects)) + self.assertNotIn(b"service-fabric-secret", b"".join(proxy.tunnel_data)) + + def test_https_proxy_missing_tls_tunneling_capability_fails_before_io(self): + factory = requests.Session + proxy, address = self._proxy() + private = "synthetic-unsupported-proxy-password" + secure_address = address.replace("http://", "https://proxy-user:{}@".format(private)) + for missing in ("urllib3.connection.HTTPSConnection._connect_tls_proxy", + "urllib3.util.ssl_.SSLTransport", "urllib3.util.ssl_.SSLContext.wrap_bio"): + for environment in (False, True): + for cached in (False, True): + session = factory() + self.addCleanup(session.close) + with self.subTest(missing=missing, environment=environment, cached=cached), patch( + missing, None, create=True), patch.dict(os.environ, + {"HTTPS_PROXY": secure_address} if environment else {}), patch( + "msal.managed_identity.requests.Session", return_value=session), patch.object( + session, "close", wraps=session.close) as close, patch.object( + _ServiceFabricHTTPSConnection, "_new_conn") as connect, self.assertLogs( + level="DEBUG") as logs: + app = self._app({"trust_env": environment, "max_retries": 2, + "proxies": {} if environment else {"https": secure_address}}) + if cached: + self._seed_cache(app, refresh=True) + self.assertEqual("cached", app.acquire_token_for_client(resource="R")["access_token"]) + else: + with self.assertRaisesRegex(ManagedIdentityError, "authenticated HTTPS proxy") as error: + app.acquire_token_for_client(resource="R") + self.assertNotIn(private, "".join(traceback.format_exception( + type(error.exception), error.exception, error.exception.__traceback__))) + connect.assert_not_called() + close.assert_called_once_with() + self.assertNotIn(private, "\n".join(logs.output)) + self.assertEqual([], proxy.connects) + self.assertEqual([], self.server.requests) + with patch("urllib3.connection.HTTPSConnection._connect_tls_proxy", None, create=True): + self.assertEqual("AT", self._app({"proxies": {"https": address}} + ).acquire_token_for_client(resource="R")["access_token"]) + self.assertEqual(1, len(proxy.connects)) + + def test_reconnection_reauthenticates_changed_endpoint_certificate(self): + other = _ServiceFabricTlsFixture() + other.setUp() + self.addCleanup(other.doCleanups) + self.addCleanup(other.tearDown) + app = self._app({"max_retries": 2}) + app.acquire_token_for_client(resource="R") + self.tls_context.set_servername_callback( + lambda sock, name, context: setattr(sock, "context", other.tls_context)) + with patch.object(_ServiceFabricHTTPSConnection, "connect", autospec=True, + side_effect=_ServiceFabricHTTPSConnection.connect) as connect: + with self.assertRaises(SSLError): + app.acquire_token_for_client(resource="R", claims_challenge="refresh") + self.assertEqual(1, connect.call_count) + self.assertEqual(1, len(self.server.requests)) + + def test_dict_subclasses_are_accepted_and_snapshotted(self): + class Options(dict): + pass + options = Options(headers=Options({"X-Subclass": "original"})) + app = self._app(options) + options["headers"]["X-Subclass"] = "mutated" + app.acquire_token_for_client(resource="R") + self.assertEqual("original", self.server.requests[-1]["headers"]["X-Subclass"]) + + def test_proxy_selection_precedence_and_credentials(self): + proxy, address = self._proxy() + credentials = address.replace("://", "://proxy-user:proxy-password@") + proxies = {"http": "http://unreachable.invalid", "all": "http://unreachable.invalid", + "https": "http://unreachable.invalid", "https://localhost": credentials} + self._app({"proxies": proxies}).acquire_token_for_client(resource="R") + self.assertEqual(1, len(proxy.connects)) + self.assertTrue(proxy.connects[0][1]["Proxy-Authorization"].startswith("Basic ")) + self.assertNotIn("Proxy-Authorization", self.server.requests[-1]["headers"]) + from msal.managed_identity import _validate_service_fabric_options + options = _validate_service_fabric_options({ + "proxies": {"https://::1": "http://[::1]:8080"}}) + self.assertEqual("http://[::1]:8080", requests.utils.select_proxy( + "https://[::1]/token", options["proxies"])) + + def test_forwarding_mode_fails_closed(self): + connection = _ServiceFabricHTTPSConnection("localhost") + connection._server_thumbprint = self.thumbprint + connection.sock = Mock() + with patch("urllib3.connection.HTTPSConnection.connect"), patch.object( + _ServiceFabricHTTPSConnection, "proxy_is_forwarding", + new_callable=lambda: property(lambda self: True), create=True): + with self.assertRaises(ssl.SSLCertVerificationError): + connection.connect() + self.assertIsNone(connection.sock) + self.assertEqual([], self.server.requests) + + def test_environment_routing_and_netrc_are_explicit_opt_in(self): + proxy, address = self._proxy() + for variable in ("HTTPS_PROXY", "ALL_PROXY"): + with self.subTest(variable=variable), patch.dict(os.environ, { + variable: address, "HTTP_PROXY": address}), patch( + "requests.sessions.get_netrc_auth", return_value=("user", "password")) as netrc: + self._app().acquire_token_for_client(resource="R") + netrc.assert_not_called() + before = len(proxy.connects) + self._app({"trust_env": True}).acquire_token_for_client(resource="R") + self.assertEqual(before + 1, len(proxy.connects)) + self.assertTrue(netrc.called) + self.assertTrue(self.server.requests[-1]["headers"]["Authorization"].startswith("Basic ")) + with patch.dict(os.environ, {"NO_PROXY": "localhost"}): + self._app({"trust_env": True}).acquire_token_for_client(resource="R") + self.assertEqual(before + 1, len(proxy.connects)) + with patch.dict(os.environ, {"HTTPS_PROXY": address, "REQUESTS_CA_BUNDLE": "unused-ca"}): + self._app({"trust_env": True, "proxies": {"https": "http://unused.invalid"}} + ).acquire_token_for_client(resource="R") + before = len(self.server.requests) + with patch.dict(os.environ, {"HTTPS_PROXY": "socks5://user:private@localhost:9"}): + with self.assertRaisesRegex(ManagedIdentityError, "proxy") as error: + self._app({"trust_env": True}).acquire_token_for_client(resource="R") + self.assertNotIn("private", str(error.exception)) + self.assertEqual(before, len(self.server.requests)) + + def test_malformed_selected_environment_proxies_have_private_errors(self): + private = "synthetic-env-proxy-password" + factory = requests.Session + original_send = HTTPAdapter.send + for variable in ("HTTPS_PROXY", "ALL_PROXY"): + for scheme in ("http", "https"): + for authority in ("localhost:invalid", "[::1", "\u2603.example"): + for cached in (False, True): + with self.subTest(variable=variable, scheme=scheme, + authority=authority, cached=cached), patch.dict(os.environ, { + variable: "{}://user:{}@{}".format(scheme, private, authority) + }), patch.object(_ServiceFabricHTTPSConnection, "_new_conn") as connect, patch.object( + HTTPAdapter, "send", autospec=True, side_effect=original_send + ) as send, self.assertLogs( + level="DEBUG") as logs: + session = factory() + self.addCleanup(session.close) + with patch("msal.managed_identity.requests.Session", + return_value=session), patch.object( + session, "close", wraps=session.close) as close: + app = self._app({"trust_env": True}) + if cached: + self._seed_cache(app, refresh=True) + self.assertEqual("cached", app.acquire_token_for_client( + resource="R")["access_token"]) + else: + with self.assertRaisesRegex( + ManagedIdentityError, "^Unsupported Service Fabric proxy configuration\\.$" + ) as error: + app.acquire_token_for_client(resource="R") + self.assertIsNone(error.exception.__cause__) + self.assertIsNone(error.exception.__context__) + self.assertNotIn(private, "".join(traceback.format_exception( + type(error.exception), error.exception, error.exception.__traceback__))) + close.assert_called_once_with() + connect.assert_not_called() + send.assert_not_called() + self.assertNotIn(private, "\n".join(logs.output)) + self.assertEqual([], self.server.requests) + + def test_only_selected_environment_proxy_is_validated(self): + proxy, address = self._proxy() + invalid = "http://user:synthetic-env-proxy-password@localhost:invalid" + for options, environment, proxied in ( + ({}, {"HTTPS_PROXY": invalid}, False), + ({"trust_env": True}, {"HTTPS_PROXY": invalid, "NO_PROXY": "localhost"}, False), + ({"trust_env": True}, {"HTTPS_PROXY": address, "ALL_PROXY": invalid, + "HTTP_PROXY": invalid}, True), + ({"trust_env": True}, {"HTTPS_PROXY": address.replace( + "http://localhost", "127.0.0.1")}, True), + ({"trust_env": True, "proxies": {"https://localhost": address}}, + {"HTTPS_PROXY": invalid}, True)): + with self.subTest(options=options, environment=environment), patch.dict( + os.environ, environment): + before = len(proxy.connects) + self.assertEqual("AT", self._app(options).acquire_token_for_client( + resource="R")["access_token"]) + self.assertEqual(before + int(proxied), len(proxy.connects)) + self.assertEqual(5, len(self.server.requests)) + + def test_environment_proxy_transport_error_retains_retry_and_fallback(self): + from urllib3.exceptions import NewConnectionError + for cached in (False, True): + with self.subTest(cached=cached), patch.dict(os.environ, { + "HTTPS_PROXY": "http://localhost:9"}), patch.object( + _ServiceFabricHTTPSConnection, "_new_conn", + side_effect=NewConnectionError(None, "controlled connection failure")) as connect: + app = self._app({"trust_env": True, "max_retries": 1}) + if cached: + self._seed_cache(app, refresh=True) + self.assertEqual("cached", app.acquire_token_for_client( + resource="R")["access_token"]) + else: + with self.assertRaises(requests.exceptions.ProxyError): + app.acquire_token_for_client(resource="R") + self.assertEqual(2, connect.call_count) + self.assertEqual([], self.server.requests) + + def test_real_netrc_and_environment_proxy_still_require_endpoint_pin(self): + netrc_path = os.path.abspath(".service-fabric-netrc-" + uuid.uuid4().hex) + with open(netrc_path, "w") as netrc: + netrc.write("machine localhost login netrc-user password netrc-password\n") + self.addCleanup(os.remove, netrc_path) + proxy, address = self._proxy(tls=True) + with patch.dict(os.environ, { + "NETRC": netrc_path, "HTTPS_PROXY": address, "REQUESTS_CA_BUNDLE": proxy.ca_path}): + self._app().acquire_token_for_client(resource="R") + self.assertNotIn("Authorization", self.server.requests[-1]["headers"]) + self.assertEqual([], proxy.connects) + self._app({"trust_env": True}).acquire_token_for_client(resource="R") + self.assertEqual(requests.auth._basic_auth_str("netrc-user", "netrc-password"), + self.server.requests[-1]["headers"]["Authorization"]) + with patch.dict(os.environ, {"IDENTITY_SERVER_THUMBPRINT": "00" * 20}): + with self.assertRaises((SSLError, requests.exceptions.ProxyError)): + self._app({"trust_env": True}).acquire_token_for_client(resource="R") + self.assertEqual(2, len(self.server.requests)) + self.assertEqual(2, len(proxy.connects)) + self.assertNotIn(b"service-fabric-secret", b"".join(proxy.tunnel_data)) + + def test_connect_retry_boundaries_timeouts_and_eventual_pin_success(self): + from urllib3.exceptions import NewConnectionError, ConnectTimeoutError + original = _ServiceFabricHTTPSConnection._new_conn + for retries in (0, 1, 2): + for eventual_success in (False, True): + for exception_type in (NewConnectionError, ConnectTimeoutError): + attempts = [] + def connect(connection): + attempts.append(connection.timeout) + if eventual_success and len(attempts) == retries + 1: + return original(connection) + raise exception_type(connection, "controlled pre-send failure") + with self.subTest(retries=retries, success=eventual_success, error=exception_type), patch.object( + _ServiceFabricHTTPSConnection, "_new_conn", connect), patch( + "urllib3.util.retry.time.sleep") as sleep: + app = self._app({"max_retries": retries, "timeout": (1.25, 2)}) + if eventual_success: + self.assertEqual("AT", app.acquire_token_for_client(resource="R")["access_token"]) + else: + with self.assertRaises(requests.exceptions.ConnectionError): + app.acquire_token_for_client(resource="R") + self.assertEqual([1.25] * (1 + retries), attempts) + sleep.assert_not_called() + + def test_non_connection_errors_are_not_retried_or_reclassified(self): + from urllib3.exceptions import ( + NewConnectionError, ProtocolError, ProxyError, ReadTimeoutError, + SSLError as Urllib3SSLError) + for failure, expected in ( + (Urllib3SSLError("controlled TLS failure"), SSLError), + (ProxyError("controlled proxy TLS failure", Urllib3SSLError("TLS")), + requests.exceptions.ProxyError), + (ProxyError("controlled other proxy failure", ValueError("other")), + requests.exceptions.ProxyError), + (ReadTimeoutError(None, None, "controlled read failure"), + requests.exceptions.ReadTimeout), + (ProtocolError("controlled protocol failure"), + requests.exceptions.ConnectionError)): + for first_connection_fails in (False, True): + for cached in (False, True): + failures = ([NewConnectionError(None, "controlled connection failure")] + if first_connection_fails else []) + [failure, failure, failure] + with self.subTest(failure=type(failure), cached=cached, + first_connection_fails=first_connection_fails), patch.object( + _ServiceFabricHTTPSConnection, "_new_conn", side_effect=failures) as connect, patch( + "urllib3.util.retry.time.sleep") as sleep: + app = self._app({"max_retries": 2}) + if cached: + self._seed_cache(app, refresh=True) + self.assertEqual("cached", app.acquire_token_for_client( + resource="R")["access_token"]) + else: + with self.assertRaises(expected): + app.acquire_token_for_client(resource="R") + self.assertEqual(1 + int(first_connection_fails), connect.call_count) + sleep.assert_not_called() + self.assertEqual([], self.server.requests) + + def test_response_failures_and_retry_after_are_not_retried(self): + for status in (404, 429, 500): + self.server.status = status + self.server.body = '{"error":{"code":"ManagedIdentityNotFound"}}' + self.server.response_headers = {"Retry-After": "1"} + before = len(self.server.requests) + result = self._app({"max_retries": 2}).acquire_token_for_client(resource="R") + self.assertEqual("invalid_client", result["error"]) + self.assertEqual(before + 1, len(self.server.requests)) + + def test_other_tls_handshake_failure_is_not_retried(self): + server, address = self._proxy() + with patch.dict(os.environ, {"IDENTITY_ENDPOINT": address.replace("http:", "https:")}), patch.object( + _ServiceFabricHTTPSConnection, "connect", autospec=True, + side_effect=_ServiceFabricHTTPSConnection.connect) as connect: + with self.assertRaises(SSLError): + self._app({"max_retries": 2}).acquire_token_for_client(resource="R") + self.assertEqual(1, connect.call_count) + self.assertEqual([], server.requests) + self.assertEqual([], self.server.requests) + + def test_timeouts_apply_to_each_wait_and_read_failure_is_not_retried(self): + original = HTTPAdapter.send + for timeout in ((5, 30), 0.5, (0.5, 1)): + with self.subTest(timeout=timeout), patch.object( + HTTPAdapter, "send", autospec=True, side_effect=original) as send: + options = {} if timeout == (5, 30) else {"timeout": timeout} + self._app(options).acquire_token_for_client(resource="R") + self.assertEqual(timeout, send.call_args.kwargs["timeout"]) + self.server.stall = threading.Event() + try: + before = len(self.server.requests) + with self.assertRaises(requests.exceptions.ConnectionError): + self._app({"timeout": (1, 0.05), "max_retries": 2} + ).acquire_token_for_client(resource="R") + self.assertEqual(before + 1, len(self.server.requests)) + finally: + self.server.stall.set() + + def test_owned_resources_close_for_success_and_all_failure_paths(self): + factory = requests.Session + for outcome in ("success", "setup", "transport", "tls", "redirect", + "json", "expiry", "interrupt", "response-close"): + sessions, responses = [], [] + def session_factory(): + session = factory() + session.close = Mock(wraps=session.close) + sessions.append(session) + return session + build_response = HTTPAdapter.build_response + def capture(adapter, request, raw): + response = build_response(adapter, request, raw) + response.close = Mock(wraps=response.close) + if outcome == "response-close": + response.close.side_effect = RuntimeError("close failure") + self.addCleanup(requests.Response.close, response) + responses.append(response) + return response + self.server.status, self.server.body = 200, json.dumps({ + "access_token": "AT", "expires_on": int(time.time()) + 3600, + "token_type": "Bearer"}) + if outcome == "redirect": + self.server.status = 302 + elif outcome == "json": + self.server.body = "not json" + elif outcome == "expiry": + self.server.body = '{"access_token":"AT","expires_on":"not an expiry"}' + failure = {"setup": RuntimeError("setup"), "transport": requests.exceptions.ConnectionError(), + "tls": SSLError(), "interrupt": KeyboardInterrupt()}.get(outcome) + target = ("msal.managed_identity._ServiceFabricHTTPAdapter" if outcome == "setup" + else "requests.sessions.Session.get") + with self.subTest(outcome=outcome), patch( + "msal.managed_identity.requests.Session", side_effect=session_factory), patch.object( + HTTPAdapter, "build_response", capture): + if failure: + with patch(target, side_effect=failure), self.assertRaises(type(failure)): + self._app().acquire_token_for_client(resource="R") + elif outcome == "success": + self._app().acquire_token_for_client(resource="R") + else: + with self.assertRaises((ManagedIdentityError, ValueError, RuntimeError)): + self._app().acquire_token_for_client(resource="R") + self.assertEqual(1, len(sessions)) + sessions[0].close.assert_called_once_with() + for response in responses: + response.close.assert_called_once_with() + + def test_fallback_matrix_and_diagnostic_privacy(self): + for cached in (False, True): + for failure in ("transport", "tls", "endpoint", "redirect", "json", "expiry", "options"): + with self.subTest(cached=cached, failure=failure): + app = self._app({"unknown-private-key": "private-value"} if failure == "options" else {}) + if cached: + self._seed_cache(app, refresh=True) + self.server.status = 302 if failure == "redirect" else 200 + self.server.response_headers = {"Location": "https://private-location.invalid"} + self.server.body = ("invalid-json" if failure == "json" + else '{"access_token":"AT","expires_on":"invalid"}') + environment = dict(os.environ) + if failure == "tls": + environment["IDENTITY_SERVER_THUMBPRINT"] = "00" * 20 + if failure == "endpoint": + environment["IDENTITY_ENDPOINT"] = "http://localhost" + transport = (patch.object(_ServiceFabricHTTPSConnection, "_new_conn", + side_effect=requests.exceptions.ConnectionError()) + if failure == "transport" else nullcontext()) + with patch.dict(os.environ, environment, clear=True), transport: + if cached and failure != "options": + self.assertEqual("cached", app.acquire_token_for_client(resource="R")["access_token"]) + else: + with self.assertRaises(Exception) as error: + app.acquire_token_for_client(resource="R") + for private in ("private-location", "private-value", "unknown-private-key", + "service-fabric-secret"): + self.assertNotIn(private, str(error.exception)) + + def test_concurrent_and_sequential_acquisitions_do_not_share_configuration(self): + from urllib3.poolmanager import pool_classes_by_scheme + before = pool_classes_by_scheme.copy() + proxy, address = self._proxy() + clients = [self._app({"headers": {"X-Client": str(index)}, + "proxies": {"https": address} if index else {}}) for index in range(2)] + sessions = [] + factory = _create_owned_service_fabric_http_client + def capture(*args): + session = factory(*args) + sessions.append(session) + return session + with patch("msal.managed_identity._create_owned_service_fabric_http_client", side_effect=capture): + with ThreadPoolExecutor(max_workers=2) as executor: + results = list(executor.map(lambda app: app.acquire_token_for_client(resource="R"), clients)) + for app in clients: + app.acquire_token_for_client(resource="R", claims_challenge="refresh") + self.assertEqual(["AT", "AT"], [result["access_token"] for result in results]) + self.assertEqual(4, len({id(session) for session in sessions})) + self.assertEqual(4, len({id(session.headers) for session in sessions})) + self.assertEqual(4, len({id(session.proxies) for session in sessions})) + self.assertEqual(["0", "0", "1", "1"], + sorted(request["headers"]["X-Client"] for request in self.server.requests)) + self.assertEqual(2, len(proxy.connects)) + self.assertEqual(before, pool_classes_by_scheme) + + def test_distinct_endpoint_pins_are_isolated_during_concurrent_acquisitions(self): + other = _ServiceFabricTlsFixture() + other.setUp() + self.addCleanup(other.doCleanups) + self.addCleanup(other.tearDown) + barrier = threading.Barrier(2) + def acquire(fixture): + barrier.wait(timeout=5) + return _obtain_token_on_service_fabric( + _UnopenedHttpClient(), fixture.endpoint, "service-fabric-secret", + fixture.thumbprint, "R", service_fabric_http_options={ + "headers": {"X-Pin": fixture.thumbprint}}) + with ThreadPoolExecutor(max_workers=2) as executor: + results = list(executor.map(acquire, (self, other))) + self.assertEqual(["AT", "AT"], [result["access_token"] for result in results]) + for fixture in (self, other): + self.assertEqual(1, len(fixture.server.requests)) + self.assertEqual(fixture.thumbprint, fixture.server.requests[0]["headers"]["X-Pin"]) + + def test_error_mapping_and_http_cache_parity(self): + for code, error in [("SecretHeaderNotFound", "unauthorized_client"), + ("ManagedIdentityNotFound", "invalid_client"), + ("ArgumentNullOrEmpty", "invalid_scope"), ("Other", "invalid_request")]: + self.server.status = 500 + self.server.body = json.dumps({"error": {"code": code}}) + cache = {} + result = self._app(http_cache=cache).acquire_token_for_client(resource="R") + self.assertEqual({"error": error, "error_description": self.server.body}, result) + self.assertEqual({}, cache, "Legacy Service Fabric also bypasses HTTP throttling") + + def test_new_diagnostics_do_not_log_private_configuration_or_redirect_values(self): + for options in ({"private-key": "private-value"}, + {"headers": {"Secret": "private-header"}}, + {"proxies": {"https": "socks5://user:private-password@host"}}): + with self.assertLogs("msal", level="DEBUG") as logs, self.assertRaises( + ManagedIdentityError) as error: + self._app(options).acquire_token_for_client(resource="R") + self.assertNotIn("private-", str(error.exception) + "\n".join(logs.output)) + self.server.status = 302 + self.server.response_headers = {"Location": "https://private-location.invalid"} + with self.assertLogs("msal", level="DEBUG") as logs, self.assertRaises( + ManagedIdentityError) as error: + self._app().acquire_token_for_client(resource="R") + self.assertNotIn("private-", str(error.exception) + "\n".join(logs.output)) + + def test_explicit_none_preserves_source_session_and_adapter_subclasses(self): + for adapter_type in (HTTPAdapter, _ServiceFabricSourceHTTPAdapter): + source = self._new_session() + source.headers["X-Legacy"] = "legacy" + source.params["legacy"] = "parameter" + source.mount("https://", adapter_type()) + source.close = Mock(wraps=source.close) + app = ManagedIdentityClient(SystemAssignedManagedIdentity(), + http_client=source, service_fabric_http_options=None) + app.acquire_token_for_client(resource="R") + self.assertEqual("legacy", self.server.requests[-1]["headers"]["X-Legacy"]) + self.assertIn("legacy=parameter", self.server.requests[-1]["path"]) + source.close.assert_not_called() + source = self._new_session() + source.mount("https://", Mock(spec=BaseAdapter)) + with self.assertRaises(ManagedIdentityError): + ManagedIdentityClient(SystemAssignedManagedIdentity(), http_client=source, + service_fabric_http_options=None).acquire_token_for_client(resource="R") + + +class ServiceFabricOptionsCompatibilityTestCase(unittest.TestCase): + def test_omission_and_none_preserve_legacy_contract(self): + with patch.dict(os.environ, { + "IDENTITY_ENDPOINT": "https://localhost", "IDENTITY_HEADER": "secret", + "IDENTITY_SERVER_THUMBPRINT": "ab" * 20}, clear=True): + for kwargs in ({}, {"service_fabric_http_options": None}): + app = ManagedIdentityClient(SystemAssignedManagedIdentity(), + http_client=_UnopenedHttpClient(), **kwargs) + with self.assertRaisesRegex(ManagedIdentityError, "requests.Session"): + app.acquire_token_for_client(resource="R") + + def test_other_providers_ignore_all_service_fabric_options(self): + environments = [ + {}, + {"IDENTITY_ENDPOINT": "http://localhost", "IDENTITY_HEADER": "secret"}, + {"MSI_ENDPOINT": "http://localhost", "MSI_SECRET": "secret"}, + {"IDENTITY_ENDPOINT": "http://localhost", "IMDS_ENDPOINT": "http://localhost"}, + ] + for environment in environments: + for options in ({}, {"timeout": 2}, {"timeout": None}, False): + with self.subTest(environment=environment, options=options), patch.dict( + os.environ, environment, clear=True), patch( + "msal.managed_identity.os.path.exists", return_value=False), patch( + "msal.managed_identity.requests.Session") as allocation: + consumer = Mock() + consumer.get.return_value = MinimalResponse(status_code=400, text='{"error":"expected"}') + app = ManagedIdentityClient(SystemAssignedManagedIdentity(), + http_client=consumer, service_fabric_http_options=options) + result = app.acquire_token_for_client(resource="R") + self.assertIn("error", result) + consumer.get.assert_called_once() + consumer.close.assert_not_called() + allocation.assert_not_called() + + def test_other_provider_success_protocol_and_cache_are_unchanged(self): + environments = [ + ({}, "2018-02-01", "msi_res_id"), + ({"IDENTITY_ENDPOINT": "http://localhost", "IDENTITY_HEADER": "secret"}, + "2019-08-01", "mi_res_id"), + ({"MSI_ENDPOINT": "http://localhost", "MSI_SECRET": "secret"}, + "2017-09-01", "msi_res_id"), + ({"IDENTITY_ENDPOINT": "http://localhost", "IMDS_ENDPOINT": "http://localhost"}, + "2020-06-01", "msi_res_id"), + ] + for environment, version, selector in environments: + for options in ({}, {"timeout": 2}, {"timeout": None}): + with self.subTest(environment=environment, options=options), patch.dict( + os.environ, environment, clear=True), patch( + "msal.managed_identity.os.path.exists", return_value=False), patch( + "msal.managed_identity.requests.Session") as allocation: + consumer = Mock() + success = MinimalResponse(status_code=200, text=json.dumps({ + "access_token": "AT", "expires_in": 3600, "token_type": "Bearer", + "expires_on": int(time.time()) + 3600, "msi_res_id": "resource-id"})) + consumer.get.side_effect = ([MinimalResponse(status_code=401, + text="", headers={"www-authenticate": "Basic realm=challenge"}), success] + if version == "2020-06-01" else [success]) + app = ManagedIdentityClient(UserAssignedManagedIdentity(resource_id="resource-id"), + http_client=consumer, service_fabric_http_options=options) + with patch.dict(_supported_arc_platforms_and_their_prefixes, + {sys.platform: os.getcwd()}), patch("builtins.open", + mock_open(read_data="secret")), patch( + "msal.managed_identity.os.stat", return_value=Mock(st_size=6)): + self.assertEqual("AT", app.acquire_token_for_client(resource="R")["access_token"]) + self.assertEqual({"api-version": version, "resource": "R", + selector: "resource-id"}, consumer.get.call_args.kwargs["params"]) + self.assertEqual("cache", app.acquire_token_for_client(resource="R")["token_source"]) + consumer.close.assert_not_called() + allocation.assert_not_called() + + @patch.dict(os.environ, { "IDENTITY_ENDPOINT": "http://localhost/token", "IMDS_ENDPOINT": "http://localhost", From c3877159a4f55539531388c2d75fa43c9c5d482a Mon Sep 17 00:00:00 2001 From: Nilesh Choudhary Date: Thu, 24 Sep 2026 11:38:08 +0100 Subject: [PATCH 2/4] test: address Service Fabric TLS review findings Add explicit TLS 1.2 fixture minimums, exact DER proxy-certificate comparison, and a distinct-hostname SNI regression while preserving production pinning. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- tests/test_mi.py | 35 +++++++++++++++++++++++++++++++++-- 1 file changed, 33 insertions(+), 2 deletions(-) diff --git a/tests/test_mi.py b/tests/test_mi.py index e343b8b3..f4633d67 100644 --- a/tests/test_mi.py +++ b/tests/test_mi.py @@ -797,6 +797,7 @@ def _proxy(self, *, tls=False, target=None, hostname="localhost"): .sign(ca_key, hashes.SHA256())) prefix = ".service-fabric-proxy-test-" + uuid.uuid4().hex server.ca_path = os.path.abspath(prefix + "-ca.pem") + server.certificate_der = certificate.public_bytes(serialization.Encoding.DER) server.thumbprint = certificate.fingerprint(hashes.SHA1()).hex() certificate_path, key_path = prefix + ".pem", prefix + ".key" for path, data in ( @@ -808,6 +809,7 @@ def _proxy(self, *, tls=False, target=None, hostname="localhost"): with open(path, "wb") as output: output.write(data) context = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER) + context.minimum_version = ssl.TLSVersion.TLSv1_2 context.load_cert_chain(certificate_path, key_path) os.remove(certificate_path) os.remove(key_path) @@ -1047,12 +1049,41 @@ def test_explicit_http_and_https_connect_proxy_pins_endpoint(self): self.assertTrue(all("Proxy-Authorization" not in request["headers"] for request in self.server.requests)) + def test_https_proxy_authenticates_proxy_hostname_not_endpoint_hostname(self): + proxy, address = self._proxy(tls=True) + server_names = [] + proxy.socket.context.set_servername_callback( + lambda sock, name, context: server_names.append(name)) + address = address.replace("://", "://proxy-user:proxy-password@") + endpoint = self.endpoint.replace("localhost", "127.0.0.1") + options = {"proxies": {"https": address}, "timeout": 2, "max_retries": 2} + with patch.dict(os.environ, {"IDENTITY_ENDPOINT": endpoint}), patch( + "requests.adapters.DEFAULT_CA_BUNDLE_PATH", proxy.ca_path): + self.assertEqual("AT", self._app(options).acquire_token_for_client( + resource="R")["access_token"]) + with patch.dict(os.environ, {"IDENTITY_SERVER_THUMBPRINT": proxy.thumbprint}): + with self.assertRaises((SSLError, requests.exceptions.ProxyError)): + self._app(options).acquire_token_for_client(resource="R") + self.assertEqual(["localhost", "localhost"], server_names) + self.assertEqual(2, len(proxy.connects)) + for target, headers in proxy.connects: + self.assertEqual("127.0.0.1:{}".format(self.server.server_port), target) + self.assertEqual(requests.auth._basic_auth_str("proxy-user", "proxy-password"), + headers["Proxy-Authorization"]) + self.assertNotIn("Secret", headers) + self.assertEqual([], proxy.requests) + self.assertEqual(1, len(self.server.requests)) + self.assertEqual("service-fabric-secret", self.server.requests[0]["headers"]["Secret"]) + self.assertNotIn("Proxy-Authorization", self.server.requests[0]["headers"]) + self.assertNotIn(b"service-fabric-secret", b"".join(proxy.tunnel_data)) + def test_https_proxy_strict_verification_accepts_valid_chain_and_checks_hostname(self): create_default_context = ssl.create_default_context contexts = [] def strict_context(*args, **kwargs): context = create_default_context(*args, **kwargs) + context.minimum_version = ssl.TLSVersion.TLSv1_2 context.verify_flags |= ssl.VERIFY_X509_STRICT contexts.append(context) return context @@ -1066,8 +1097,8 @@ def strict_context(*args, **kwargs): # valid name before testing the actual proxy hostname below. with socket.create_connection(("localhost", proxy.server_port), timeout=2) as sock: with context.wrap_socket(sock, server_hostname=hostname) as tls_socket: - self.assertEqual(proxy.thumbprint, - hashlib.sha1(tls_socket.getpeercert(binary_form=True)).hexdigest()) + self.assertEqual(proxy.certificate_der, + tls_socket.getpeercert(binary_form=True)) self.assertEqual([], proxy.connects) self.assertEqual([], proxy.requests) before = len(self.server.requests) From 34510b343160dea46b2db65e6579467fce137ddd Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Thu, 24 Sep 2026 12:24:36 +0000 Subject: [PATCH 3/4] fix: authenticate selected HTTPS proxy hostname Co-authored-by: 4gust <107404295+4gust@users.noreply.github.com> --- msal/managed_identity.py | 4 +++- tests/test_mi.py | 21 ++++++++++++++++++++- 2 files changed, 23 insertions(+), 2 deletions(-) diff --git a/msal/managed_identity.py b/msal/managed_identity.py index deae0bac..820df42b 100644 --- a/msal/managed_identity.py +++ b/msal/managed_identity.py @@ -863,11 +863,12 @@ class _ServiceFabricHTTPSConnectionPool(HTTPSConnectionPool): class _ServiceFabricAuthenticatedProxyConnection(_ServiceFabricHTTPSConnection): _proxy_ssl_context = None + _proxy_hostname = None def _connect_tls_proxy(self, hostname, sock): # This hook runs before CONNECT. The endpoint's CERT_NONE must never # configure the outer TLS handshake, which authenticates proxy credentials. - hostname = hostname.strip("[]").split("%", 1)[0].rstrip(".") + hostname = self._proxy_hostname.strip("[]").split("%", 1)[0].rstrip(".") proxy_socket = self._proxy_ssl_context.wrap_socket(sock, server_hostname=hostname) self.proxy_is_verified = True return proxy_socket @@ -944,6 +945,7 @@ def send(self, request, stream=False, timeout=None, verify=True, cert=None, prox capath=ca_bundle if os.path.isdir(ca_bundle) else None, cafile=None if os.path.isdir(ca_bundle) else ca_bundle) self._connection_pool_class.ConnectionCls._proxy_ssl_context = context + self._connection_pool_class.ConnectionCls._proxy_hostname = parse_url(proxy).host return super(_ServiceFabricHTTPAdapter, self).send( request, stream=stream, timeout=timeout, verify=verify, cert=cert, proxies=proxies) diff --git a/tests/test_mi.py b/tests/test_mi.py index f4633d67..5ea911c7 100644 --- a/tests/test_mi.py +++ b/tests/test_mi.py @@ -1077,6 +1077,23 @@ def test_https_proxy_authenticates_proxy_hostname_not_endpoint_hostname(self): self.assertNotIn("Proxy-Authorization", self.server.requests[0]["headers"]) self.assertNotIn(b"service-fabric-secret", b"".join(proxy.tunnel_data)) + def test_https_proxy_hostname_is_independent_of_tls_hook_argument(self): + session = _create_owned_service_fabric_http_client( + {"headers": {}, "proxies": {}, "trust_env": False, "max_retries": 0}, + self.thumbprint) + self.addCleanup(session.close) + adapter = session.get_adapter(self.endpoint) + with patch("msal.managed_identity.ssl.create_default_context") as context, patch.object( + HTTPAdapter, "send"): + adapter.send(requests.Request("GET", self.endpoint).prepare(), + proxies={"https": "https://proxy.example:8443"}) + connection = adapter._connection_pool_class.ConnectionCls("origin.example") + sock = Mock() + self.assertIs(context.return_value.wrap_socket.return_value, + connection._connect_tls_proxy("origin.example", sock)) + context.return_value.wrap_socket.assert_called_once_with( + sock, server_hostname="proxy.example") + def test_https_proxy_strict_verification_accepts_valid_chain_and_checks_hostname(self): create_default_context = ssl.create_default_context contexts = [] @@ -1150,7 +1167,9 @@ def test_https_proxy_rejects_untrusted_chain_and_wrong_hostname_before_connect(s trust = (patch("requests.adapters.DEFAULT_CA_BUNDLE_PATH", proxy.ca_path) if failure == "hostname" else nullcontext()) with self.subTest(failure=failure, environment=environment, cached=cached), trust, patch.dict( - os.environ, {"HTTPS_PROXY": address} if environment else {}), patch( + os.environ, dict( + {"HTTPS_PROXY": address} if environment else {}, + IDENTITY_ENDPOINT=self.endpoint.replace("localhost", "wrong.invalid"))), patch( "msal.managed_identity.requests.Session", return_value=session), patch.object( session, "close", wraps=session.close) as close, patch.object( _ServiceFabricHTTPSConnection, "connect", autospec=True, From 787b89bdb24651b80f3dbf0bddd1ee7e4e571f48 Mon Sep 17 00:00:00 2001 From: Nilesh Choudhary Date: Fri, 25 Sep 2026 12:05:01 +0100 Subject: [PATCH 4/4] refactor: use urllib3 fingerprint validation for Service Fabric Delegate endpoint certificate pinning to native urllib3 connections while retaining independent HTTPS proxy authentication. Preserve proxy URL normalization and Requests TLS context isolation. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- docs/index.rst | 6 +- msal/managed_identity.py | 83 +++++++------------- tests/test_mi.py | 159 +++++++++++++++++++++++++++++++-------- 3 files changed, 161 insertions(+), 87 deletions(-) diff --git a/docs/index.rst b/docs/index.rst index 024e725d..8ac2e548 100644 --- a/docs/index.rst +++ b/docs/index.rst @@ -281,8 +281,10 @@ environment failures retain the existing eligible cached-token fallback. MSAL owns and closes a separate session and response for each network acquisition, including its retries. There is no new close API or persistent pool. Every connection authenticates the actual endpoint certificate against -``IDENTITY_SERVER_THUMBPRINT`` before sending the environment-sourced -``Secret``. Self-signed certificates remain supported by this exact pin; +``IDENTITY_SERVER_THUMBPRINT`` using urllib3's ``assert_fingerprint`` before +sending the environment-sourced ``Secret``. Direct and HTTP CONNECT routes +use urllib3's standard HTTPS connection and pool classes. +Self-signed certificates remain supported by this exact pin; there is no caller TLS override. Endpoints must use HTTPS. Redirects are never followed, even on the same origin; all HTTP 300-399 responses raise :class:`msal.ManagedIdentityError` without exposing the redirect target. diff --git a/msal/managed_identity.py b/msal/managed_identity.py index 820df42b..4dbad161 100644 --- a/msal/managed_identity.py +++ b/msal/managed_identity.py @@ -4,7 +4,6 @@ # This code is licensed under the MIT License. import copy import hashlib -import hmac import ipaddress import json import logging @@ -836,32 +835,7 @@ def _normalize_service_fabric_thumbprint(server_thumbprint): return normalized.lower() -class _ServiceFabricHTTPSConnection(HTTPSConnection): - """An HTTPS connection that authenticates the Service Fabric endpoint certificate.""" - _server_thumbprint = None - - def connect(self): - super(_ServiceFabricHTTPSConnection, self).connect() - if getattr(self, "proxy_is_forwarding", False): - self.close() - raise ssl.SSLCertVerificationError( - "Cannot validate the Service Fabric endpoint certificate through " - "a forwarding proxy.") - certificate = self.sock.getpeercert(binary_form=True) - actual_thumbprint = hashlib.sha1(certificate).hexdigest() - if not hmac.compare_digest(actual_thumbprint, self._server_thumbprint): - self.close() - raise ssl.SSLCertVerificationError( - "Service Fabric endpoint certificate thumbprint does not match " - "IDENTITY_SERVER_THUMBPRINT.") - self.is_verified = True - - -class _ServiceFabricHTTPSConnectionPool(HTTPSConnectionPool): - ConnectionCls = _ServiceFabricHTTPSConnection - - -class _ServiceFabricAuthenticatedProxyConnection(_ServiceFabricHTTPSConnection): +class _ServiceFabricAuthenticatedProxyConnection(HTTPSConnection): _proxy_ssl_context = None _proxy_hostname = None @@ -891,37 +865,39 @@ class _ServiceFabricHTTPAdapter(HTTPAdapter): """Use certificate-thumbprint authentication for the Service Fabric endpoint.""" def __init__(self, server_thumbprint, *args, _restrict_proxies=False, **kwargs): + self._server_thumbprint = server_thumbprint self._restrict_proxies = _restrict_proxies - connection_class = type( - "_PinnedServiceFabricHTTPSConnection", - (_ServiceFabricAuthenticatedProxyConnection if _restrict_proxies - else _ServiceFabricHTTPSConnection,), - {"_server_thumbprint": server_thumbprint}, - ) - self._connection_pool_class = type( - "_PinnedServiceFabricHTTPSConnectionPool", - (_ServiceFabricHTTPSConnectionPool,), - {"ConnectionCls": connection_class}, - ) super(_ServiceFabricHTTPAdapter, self).__init__(*args, **kwargs) - def _configure_pool_manager(self, pool_manager): - # PoolManager's mapping is module-global by default, so copy it before - # replacing HTTPS only for this derived Service Fabric session. - pool_manager.pool_classes_by_scheme = pool_manager.pool_classes_by_scheme.copy() - pool_manager.pool_classes_by_scheme["https"] = self._connection_pool_class - def init_poolmanager(self, connections, maxsize, block=False, **pool_kwargs): + pool_kwargs["assert_fingerprint"] = self._server_thumbprint super(_ServiceFabricHTTPAdapter, self).init_poolmanager( connections, maxsize, block=block, **pool_kwargs) - self._configure_pool_manager(self.poolmanager) def proxy_manager_for(self, proxy, **proxy_kwargs): if self._restrict_proxies and not _valid_service_fabric_proxy(proxy): raise ManagedIdentityError("Unsupported Service Fabric proxy configuration.") + if proxy_kwargs.get("use_forwarding_for_https"): + raise requests.exceptions.SSLError( + "Cannot validate the Service Fabric endpoint certificate through " + "a forwarding proxy.") + proxy_kwargs["assert_fingerprint"] = self._server_thumbprint + new_manager = proxy not in self.proxy_manager pool_manager = super(_ServiceFabricHTTPAdapter, self).proxy_manager_for( proxy, **proxy_kwargs) - self._configure_pool_manager(pool_manager) + if new_manager and self._restrict_proxies and urlsplit(proxy).scheme.lower() == "https": + connection_class = type( + "_ServiceFabricProxyConnection", + (_ServiceFabricAuthenticatedProxyConnection,), + {"_proxy_hostname": parse_url(proxy).host}, + ) + # Isolate outer-proxy trust without changing direct or HTTP CONNECT pools. + pool_manager.pool_classes_by_scheme = pool_manager.pool_classes_by_scheme.copy() + pool_manager.pool_classes_by_scheme["https"] = type( + "_ServiceFabricProxyConnectionPool", + (HTTPSConnectionPool,), + {"ConnectionCls": connection_class}, + ) return pool_manager def send(self, request, stream=False, timeout=None, verify=True, cert=None, proxies=None): @@ -944,16 +920,13 @@ def send(self, request, stream=False, timeout=None, verify=True, cert=None, prox context = ssl.create_default_context( capath=ca_bundle if os.path.isdir(ca_bundle) else None, cafile=None if os.path.isdir(ca_bundle) else ca_bundle) - self._connection_pool_class.ConnectionCls._proxy_ssl_context = context - self._connection_pool_class.ConnectionCls._proxy_hostname = parse_url(proxy).host + proxy = requests.utils.prepend_scheme_if_needed(proxy, "http") + pool_manager = self.proxy_manager_for(proxy) + pool_manager.pool_classes_by_scheme["https"].ConnectionCls._proxy_ssl_context = context + # Select pin-only pools before Requests can attach its shared CA context. + # The original verify value above applies only to the outer proxy TLS. return super(_ServiceFabricHTTPAdapter, self).send( - request, stream=stream, timeout=timeout, verify=verify, cert=cert, proxies=proxies) - - def cert_verify(self, conn, url, verify, cert): - # The exact Service Fabric certificate thumbprint is the trust anchor. - # CA configuration applies only to the owned transport's outer proxy TLS. - super(_ServiceFabricHTTPAdapter, self).cert_verify( - conn, url, verify=False, cert=cert) + request, stream=stream, timeout=timeout, verify=False, cert=cert, proxies=proxies) def _create_owned_service_fabric_http_client(options, server_thumbprint): diff --git a/tests/test_mi.py b/tests/test_mi.py index 5ea911c7..b7ef6528 100644 --- a/tests/test_mi.py +++ b/tests/test_mi.py @@ -24,6 +24,8 @@ import requests from requests.adapters import BaseAdapter, HTTPAdapter from requests.exceptions import SSLError +from urllib3.connection import HTTPSConnection +from urllib3.connectionpool import HTTPSConnectionPool from urllib3.util.retry import Retry from cryptography import x509 from cryptography.hazmat.primitives import hashes, serialization @@ -52,7 +54,6 @@ _create_service_fabric_http_client, _obtain_token_on_service_fabric, _ServiceFabricHTTPAdapter, - _ServiceFabricHTTPSConnection, _create_owned_service_fabric_http_client, ) from msal.token_cache import is_subdict_of @@ -643,9 +644,12 @@ def test_derived_client_preserves_standard_session_settings_without_mutation(sel 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) + pool = derived_adapter.poolmanager.connection_from_url(self.endpoint) + self.assertIs(type(pool), HTTPSConnectionPool) + self.assertIs(pool.ConnectionCls, HTTPSConnection) + self.assertEqual(self.thumbprint, pool.assert_fingerprint) + self.assertIsNone( + source_adapter.poolmanager.connection_from_url(self.endpoint).assert_fingerprint) self.assertEqual(source_pool_classes, source_adapter.poolmanager.pool_classes_by_scheme) response = derived.get( self.endpoint, @@ -948,7 +952,7 @@ def test_transport_invalid_proxy_idna_is_private_and_preallocation(self): with self.subTest(scheme=scheme, cached=cached), patch( "msal.managed_identity.requests.Session", wraps=requests.Session ) as allocation, patch.object( - _ServiceFabricHTTPSConnection, "_new_conn") as connect, self.assertLogs( + HTTPSConnection, "_new_conn") as connect, self.assertLogs( level="DEBUG") as logs: app = self._app({"proxies": { "https": "{}://user:{}@\u2603.example".format(scheme, private)}}) @@ -989,6 +993,51 @@ def test_matching_pin_normalization(self): self.assertEqual("AT", self._app().acquire_token_for_client(resource="R")["access_token"]) self.assertEqual(3, len(self.server.requests)) + def test_native_fingerprint_configuration_is_isolated_by_route_and_adapter(self): + from urllib3.poolmanager import pool_classes_by_scheme + original_pool_classes = pool_classes_by_scheme.copy() + for restricted in (False, True): + adapters = [_ServiceFabricHTTPAdapter(pin, _restrict_proxies=restricted) + for pin in (self.thumbprint, "00" * 20)] + for adapter in adapters: + self.addCleanup(adapter.close) + for proxy in (None, "http://localhost:8080", "https://localhost:8443"): + for adapter, pin in zip(adapters, (self.thumbprint, "00" * 20)): + with self.subTest(restricted=restricted, proxy=proxy, pin=pin): + manager = (adapter.proxy_manager_for(proxy) if proxy + else adapter.poolmanager) + pool = manager.connection_from_url(self.endpoint) + connection = pool._new_conn() + self.assertEqual(pin, pool.assert_fingerprint) + self.assertEqual(pin, connection.assert_fingerprint) + self.assertIs(type(connection).connect, HTTPSConnection.connect) + if restricted and proxy and proxy.startswith("https:"): + self.assertIsNot(type(pool), HTTPSConnectionPool) + self.assertIsNot(manager.pool_classes_by_scheme, pool_classes_by_scheme) + else: + self.assertIs(type(pool), HTTPSConnectionPool) + self.assertIs(type(connection), HTTPSConnection) + if proxy: + self.assertIs(manager, adapter.proxy_manager_for(proxy)) + self.assertEqual(original_pool_classes, pool_classes_by_scheme) + + def test_direct_endpoint_uses_pin_instead_of_ca_or_hostname_validation(self): + with patch.dict(os.environ, { + "IDENTITY_ENDPOINT": self.endpoint.replace("localhost", "127.0.0.1")}): + self.assertEqual("AT", self._app().acquire_token_for_client(resource="R")["access_token"]) + self.assertEqual(1, len(self.server.requests)) + self.assertEqual("service-fabric-secret", self.server.requests[0]["headers"]["Secret"]) + + def test_pinning_does_not_change_standard_requests_tls_verification(self): + context = ssl.create_default_context() + with patch.object(requests.adapters, "_preloaded_ssl_context", context, create=True): + self.assertEqual("AT", self._app().acquire_token_for_client(resource="R")["access_token"]) + self.assertEqual(ssl.CERT_REQUIRED, context.verify_mode) + self.assertTrue(context.check_hostname) + with self._new_session() as session, self.assertRaises(SSLError): + session.get(self.endpoint, timeout=2) + self.assertEqual(1, len(self.server.requests)) + def test_bad_pin_and_endpoint_never_send_or_retry_secret(self): for pin, endpoint, error in [ ("00" * 20, self.endpoint, SSLError), @@ -997,8 +1046,8 @@ def test_bad_pin_and_endpoint_never_send_or_retry_secret(self): (self.thumbprint, self.endpoint.replace("https:", "http:"), ManagedIdentityError)]: with self.subTest(pin=pin), patch.dict(os.environ, { "IDENTITY_SERVER_THUMBPRINT": pin, "IDENTITY_ENDPOINT": endpoint}), patch.object( - _ServiceFabricHTTPSConnection, "connect", autospec=True, - side_effect=_ServiceFabricHTTPSConnection.connect) as connect: + HTTPSConnection, "connect", autospec=True, + side_effect=HTTPSConnection.connect) as connect: with self.assertRaises(error): self._app({"max_retries": 2}).acquire_token_for_client(resource="R") self.assertEqual(1 if error is SSLError else 0, connect.call_count) @@ -1087,13 +1136,44 @@ def test_https_proxy_hostname_is_independent_of_tls_hook_argument(self): HTTPAdapter, "send"): adapter.send(requests.Request("GET", self.endpoint).prepare(), proxies={"https": "https://proxy.example:8443"}) - connection = adapter._connection_pool_class.ConnectionCls("origin.example") + pool_manager = adapter.proxy_manager_for("https://proxy.example:8443") + connection = pool_manager.connection_from_url(self.endpoint)._new_conn() sock = Mock() self.assertIs(context.return_value.wrap_socket.return_value, connection._connect_tls_proxy("origin.example", sock)) context.return_value.wrap_socket.assert_called_once_with( sock, server_hostname="proxy.example") + def test_https_proxy_normalization_preserves_authenticated_context(self): + from urllib3.util import connection as urllib3_connection + original_new_socket = urllib3_connection.create_connection + def connect(address, *args, **kwargs): + host, port = address + if host == "xn--bcher-kva.example": + address = ("localhost", port) + return original_new_socket(address, *args, **kwargs) + for hostname, authority in [ + ("localhost", "https://LOCALHOST"), + ("localhost", "HTTPS://localhost"), + ("xn--bcher-kva.example", "https://b\u00fccher.example")]: + proxy, address = self._proxy(tls=True, hostname=hostname) + selected = address.replace("https://localhost", authority) + server_names = [] + proxy.socket.context.set_servername_callback( + lambda sock, name, context: server_names.append(name)) + with patch("requests.adapters.DEFAULT_CA_BUNDLE_PATH", proxy.ca_path), patch.object( + urllib3_connection, "create_connection", side_effect=connect): + for environment in (False, True): + with self.subTest(proxy=selected, environment=environment), patch.dict( + os.environ, {"HTTPS_PROXY": selected} if environment else {}): + self.assertEqual("AT", self._app({ + "trust_env": environment, + "proxies": {} if environment else {"https": selected}, + }).acquire_token_for_client(resource="R")["access_token"]) + self.assertEqual([hostname, hostname], server_names) + self.assertEqual(2, len(proxy.connects)) + self.assertEqual(6, len(self.server.requests)) + def test_https_proxy_strict_verification_accepts_valid_chain_and_checks_hostname(self): create_default_context = ssl.create_default_context contexts = [] @@ -1172,8 +1252,8 @@ def test_https_proxy_rejects_untrusted_chain_and_wrong_hostname_before_connect(s IDENTITY_ENDPOINT=self.endpoint.replace("localhost", "wrong.invalid"))), patch( "msal.managed_identity.requests.Session", return_value=session), patch.object( session, "close", wraps=session.close) as close, patch.object( - _ServiceFabricHTTPSConnection, "connect", autospec=True, - side_effect=_ServiceFabricHTTPSConnection.connect) as connect, patch( + HTTPSConnection, "connect", autospec=True, + side_effect=HTTPSConnection.connect) as connect, patch( "urllib3.util.retry.time.sleep") as sleep, self.assertLogs(level="DEBUG") as logs: app = self._app({"trust_env": environment, "max_retries": 2, "proxies": {} if environment else {"https": address}}) @@ -1262,7 +1342,7 @@ def test_https_proxy_missing_tls_tunneling_capability_fails_before_io(self): {"HTTPS_PROXY": secure_address} if environment else {}), patch( "msal.managed_identity.requests.Session", return_value=session), patch.object( session, "close", wraps=session.close) as close, patch.object( - _ServiceFabricHTTPSConnection, "_new_conn") as connect, self.assertLogs( + HTTPSConnection, "_new_conn") as connect, self.assertLogs( level="DEBUG") as logs: app = self._app({"trust_env": environment, "max_retries": 2, "proxies": {} if environment else {"https": secure_address}}) @@ -1293,8 +1373,8 @@ def test_reconnection_reauthenticates_changed_endpoint_certificate(self): app.acquire_token_for_client(resource="R") self.tls_context.set_servername_callback( lambda sock, name, context: setattr(sock, "context", other.tls_context)) - with patch.object(_ServiceFabricHTTPSConnection, "connect", autospec=True, - side_effect=_ServiceFabricHTTPSConnection.connect) as connect: + with patch.object(HTTPSConnection, "connect", autospec=True, + side_effect=HTTPSConnection.connect) as connect: with self.assertRaises(SSLError): app.acquire_token_for_client(resource="R", claims_challenge="refresh") self.assertEqual(1, connect.call_count) @@ -1325,15 +1405,18 @@ def test_proxy_selection_precedence_and_credentials(self): "https://[::1]/token", options["proxies"])) def test_forwarding_mode_fails_closed(self): - connection = _ServiceFabricHTTPSConnection("localhost") - connection._server_thumbprint = self.thumbprint - connection.sock = Mock() - with patch("urllib3.connection.HTTPSConnection.connect"), patch.object( - _ServiceFabricHTTPSConnection, "proxy_is_forwarding", - new_callable=lambda: property(lambda self: True), create=True): - with self.assertRaises(ssl.SSLCertVerificationError): - connection.connect() - self.assertIsNone(connection.sock) + for restricted in (False, True): + adapter = _ServiceFabricHTTPAdapter(self.thumbprint, _restrict_proxies=restricted) + self.addCleanup(adapter.close) + for cached in (False, True): + with self.subTest(restricted=restricted, cached=cached), patch.object( + HTTPSConnection, "_new_conn") as connect: + if cached: + adapter.proxy_manager_for("https://localhost:8443") + with self.assertRaisesRegex(SSLError, "forwarding proxy"): + adapter.proxy_manager_for( + "https://localhost:8443", use_forwarding_for_https=True) + connect.assert_not_called() self.assertEqual([], self.server.requests) def test_environment_routing_and_netrc_are_explicit_opt_in(self): @@ -1373,7 +1456,7 @@ def test_malformed_selected_environment_proxies_have_private_errors(self): with self.subTest(variable=variable, scheme=scheme, authority=authority, cached=cached), patch.dict(os.environ, { variable: "{}://user:{}@{}".format(scheme, private, authority) - }), patch.object(_ServiceFabricHTTPSConnection, "_new_conn") as connect, patch.object( + }), patch.object(HTTPSConnection, "_new_conn") as connect, patch.object( HTTPAdapter, "send", autospec=True, side_effect=original_send ) as send, self.assertLogs( level="DEBUG") as logs: @@ -1427,7 +1510,7 @@ def test_environment_proxy_transport_error_retains_retry_and_fallback(self): for cached in (False, True): with self.subTest(cached=cached), patch.dict(os.environ, { "HTTPS_PROXY": "http://localhost:9"}), patch.object( - _ServiceFabricHTTPSConnection, "_new_conn", + HTTPSConnection, "_new_conn", side_effect=NewConnectionError(None, "controlled connection failure")) as connect: app = self._app({"trust_env": True, "max_retries": 1}) if cached: @@ -1463,7 +1546,7 @@ def test_real_netrc_and_environment_proxy_still_require_endpoint_pin(self): def test_connect_retry_boundaries_timeouts_and_eventual_pin_success(self): from urllib3.exceptions import NewConnectionError, ConnectTimeoutError - original = _ServiceFabricHTTPSConnection._new_conn + original = HTTPSConnection._new_conn for retries in (0, 1, 2): for eventual_success in (False, True): for exception_type in (NewConnectionError, ConnectTimeoutError): @@ -1474,7 +1557,7 @@ def connect(connection): return original(connection) raise exception_type(connection, "controlled pre-send failure") with self.subTest(retries=retries, success=eventual_success, error=exception_type), patch.object( - _ServiceFabricHTTPSConnection, "_new_conn", connect), patch( + HTTPSConnection, "_new_conn", connect), patch( "urllib3.util.retry.time.sleep") as sleep: app = self._app({"max_retries": retries, "timeout": (1.25, 2)}) if eventual_success: @@ -1505,7 +1588,7 @@ def test_non_connection_errors_are_not_retried_or_reclassified(self): if first_connection_fails else []) + [failure, failure, failure] with self.subTest(failure=type(failure), cached=cached, first_connection_fails=first_connection_fails), patch.object( - _ServiceFabricHTTPSConnection, "_new_conn", side_effect=failures) as connect, patch( + HTTPSConnection, "_new_conn", side_effect=failures) as connect, patch( "urllib3.util.retry.time.sleep") as sleep: app = self._app({"max_retries": 2}) if cached: @@ -1532,8 +1615,8 @@ def test_response_failures_and_retry_after_are_not_retried(self): def test_other_tls_handshake_failure_is_not_retried(self): server, address = self._proxy() with patch.dict(os.environ, {"IDENTITY_ENDPOINT": address.replace("http:", "https:")}), patch.object( - _ServiceFabricHTTPSConnection, "connect", autospec=True, - side_effect=_ServiceFabricHTTPSConnection.connect) as connect: + HTTPSConnection, "connect", autospec=True, + side_effect=HTTPSConnection.connect) as connect: with self.assertRaises(SSLError): self._app({"max_retries": 2}).acquire_token_for_client(resource="R") self.assertEqual(1, connect.call_count) @@ -1622,7 +1705,7 @@ def test_fallback_matrix_and_diagnostic_privacy(self): environment["IDENTITY_SERVER_THUMBPRINT"] = "00" * 20 if failure == "endpoint": environment["IDENTITY_ENDPOINT"] = "http://localhost" - transport = (patch.object(_ServiceFabricHTTPSConnection, "_new_conn", + transport = (patch.object(HTTPSConnection, "_new_conn", side_effect=requests.exceptions.ConnectionError()) if failure == "transport" else nullcontext()) with patch.dict(os.environ, environment, clear=True), transport: @@ -1680,6 +1763,22 @@ def acquire(fixture): self.assertEqual(1, len(fixture.server.requests)) self.assertEqual(fixture.thumbprint, fixture.server.requests[0]["headers"]["X-Pin"]) + def test_concurrent_clients_with_different_pins_cannot_share_verified_connections(self): + barrier = threading.Barrier(2) + def acquire(pin): + barrier.wait(timeout=5) + return _obtain_token_on_service_fabric( + _UnopenedHttpClient(), self.endpoint, "service-fabric-secret", + pin, "R", service_fabric_http_options={"headers": {"X-Pin": pin}}) + with ThreadPoolExecutor(max_workers=2) as executor: + matching = executor.submit(acquire, self.thumbprint) + mismatching = executor.submit(acquire, "00" * 20) + self.assertEqual("AT", matching.result(timeout=5)["access_token"]) + with self.assertRaises(SSLError): + mismatching.result(timeout=5) + self.assertEqual(1, len(self.server.requests)) + self.assertEqual(self.thumbprint, self.server.requests[0]["headers"]["X-Pin"]) + def test_error_mapping_and_http_cache_parity(self): for code, error in [("SecretHeaderNotFound", "unauthorized_client"), ("ManagedIdentityNotFound", "invalid_client"),