diff --git a/sdk/identity/azure-identity/CHANGELOG.md b/sdk/identity/azure-identity/CHANGELOG.md index 4928a5f84707..d434d7017122 100644 --- a/sdk/identity/azure-identity/CHANGELOG.md +++ b/sdk/identity/azure-identity/CHANGELOG.md @@ -1,9 +1,11 @@ # Release History -## 1.25.2 (Unreleased) +## 1.26.0b1 (Unreleased) ### Features Added +- Added support for `WorkloadIdentityCredential` identity binding mode in AKS environments. This feature addresses Entra's limitation on the number of federated identity credentials (FICs) per managed identity by utilizing an AKS proxy that handles FIC exchanges on behalf of pods. ([#43287](https://github.com/Azure/azure-sdk-for-python/pull/43287)) + ### Breaking Changes ### Bugs Fixed diff --git a/sdk/identity/azure-identity/azure/identity/_constants.py b/sdk/identity/azure-identity/azure/identity/_constants.py index 09732f3de9ec..4f86f3d94c48 100644 --- a/sdk/identity/azure-identity/azure/identity/_constants.py +++ b/sdk/identity/azure-identity/azure/identity/_constants.py @@ -68,5 +68,10 @@ class EnvironmentVariables: AZURE_REGIONAL_AUTHORITY_NAME = "AZURE_REGIONAL_AUTHORITY_NAME" AZURE_FEDERATED_TOKEN_FILE = "AZURE_FEDERATED_TOKEN_FILE" + AZURE_KUBERNETES_SNI_NAME = "AZURE_KUBERNETES_SNI_NAME" + AZURE_KUBERNETES_TOKEN_PROXY = "AZURE_KUBERNETES_TOKEN_PROXY" + AZURE_KUBERNETES_CA_FILE = "AZURE_KUBERNETES_CA_FILE" + AZURE_KUBERNETES_CA_DATA = "AZURE_KUBERNETES_CA_DATA" + AZURE_TOKEN_CREDENTIALS = "AZURE_TOKEN_CREDENTIALS" WORKLOAD_IDENTITY_VARS = (AZURE_AUTHORITY_HOST, AZURE_TENANT_ID, AZURE_FEDERATED_TOKEN_FILE) diff --git a/sdk/identity/azure-identity/azure/identity/_credentials/workload_identity.py b/sdk/identity/azure-identity/azure/identity/_credentials/workload_identity.py index db15c4ed633c..ef2453e2079d 100644 --- a/sdk/identity/azure-identity/azure/identity/_credentials/workload_identity.py +++ b/sdk/identity/azure-identity/azure/identity/_credentials/workload_identity.py @@ -9,6 +9,7 @@ from .client_assertion import ClientAssertionCredential from .._constants import EnvironmentVariables +from .._internal import within_credential_chain WORKLOAD_CONFIG_ERROR = ( @@ -16,6 +17,10 @@ "configured. See the troubleshooting guide for more information: " "https://aka.ms/azsdk/python/identity/workloadidentitycredential/troubleshoot" ) +CA_DATA_FILE_ERROR = "Both AZURE_KUBERNETES_CA_FILE and AZURE_KUBERNETES_CA_DATA are set. Only one should be set." +CUSTOM_PROXY_ENV_ERROR = ( + "AZURE_KUBERNETES_TOKEN_PROXY is not set but other custom endpoint-related environment variables are present." +) class TokenFileMixin: @@ -99,6 +104,33 @@ def __init__( assert token_file_path is not None self._token_file_path = token_file_path + + if kwargs.pop("use_token_proxy", False) and not within_credential_chain.get(): + token_proxy_endpoint = os.environ.get(EnvironmentVariables.AZURE_KUBERNETES_TOKEN_PROXY) + sni = os.environ.get(EnvironmentVariables.AZURE_KUBERNETES_SNI_NAME) + ca_file = os.environ.get(EnvironmentVariables.AZURE_KUBERNETES_CA_FILE) + ca_data = os.environ.get(EnvironmentVariables.AZURE_KUBERNETES_CA_DATA) + if token_proxy_endpoint: + if ca_file and ca_data: + raise ValueError(CA_DATA_FILE_ERROR) + + transport = _get_transport( + sni=sni, + token_proxy_endpoint=token_proxy_endpoint, + ca_file=ca_file, + ca_data=ca_data, + ) + + if transport: + kwargs["transport"] = transport + else: + raise ValueError( + "Transport creation failed. Ensure that the requests package is installed to enable token " + "proxy usage in this credential." + ) + elif sni or ca_file or ca_data: + raise ValueError(CUSTOM_PROXY_ENV_ERROR) + super(WorkloadIdentityCredential, self).__init__( tenant_id=tenant_id, client_id=client_id, @@ -106,3 +138,18 @@ def __init__( token_file_path=token_file_path, **kwargs, ) + + +def _get_transport(sni, token_proxy_endpoint, ca_file, ca_data): + try: + from .._internal.token_binding_transport_requests import CustomRequestsTransport + + return CustomRequestsTransport( + sni=sni, + proxy_endpoint=token_proxy_endpoint, + ca_file=ca_file, + ca_data=ca_data, + ) + + except ImportError: + return None diff --git a/sdk/identity/azure-identity/azure/identity/_internal/token_binding_transport_mixin.py b/sdk/identity/azure-identity/azure/identity/_internal/token_binding_transport_mixin.py new file mode 100644 index 000000000000..d790c25d6d24 --- /dev/null +++ b/sdk/identity/azure-identity/azure/identity/_internal/token_binding_transport_mixin.py @@ -0,0 +1,123 @@ +# ------------------------------------ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. +# ------------------------------------ +# cspell:ignore cafile +import os +import urllib.parse +from typing import Optional, Any + +from azure.core.rest import HttpRequest + + +class TokenBindingTransportMixin: + """Mixin class providing URL validation, CA file tracking, and proxy URL functionality for transport classes.""" + + def __init__(self, **kwargs: Any) -> None: + """Initialize CA file tracking and proxy attributes.""" + self._ca_file = kwargs.pop("ca_file", None) + self._ca_data = kwargs.pop("ca_data", None) + self._proxy_endpoint = kwargs.pop("proxy_endpoint", None) + self._sni = kwargs.pop("sni", None) + + self._ca_file_mtime: Optional[float] = None + + if self._ca_file and self._ca_data: + raise ValueError("Both ca_file and ca_data are set. Only one should be set") + + if self._proxy_endpoint: + self._validate_url(self._proxy_endpoint) + + # If we have a ca_file, read it once and store as ca_data + if self._ca_file: + self._load_ca_file_to_data() + + super().__init__() + + def _validate_url(self, url: str) -> None: + """Validate that a URL meets security requirements for HTTPS connections. + + :param url: The URL to validate. + :type url: str + :raises ValueError: If the URL does not meet security requirements. + """ + parsed_url = urllib.parse.urlparse(url) + if parsed_url.scheme != "https": + raise ValueError(f"Endpoint URL ({url}) must use the 'https' scheme. Got '{parsed_url.scheme}' instead.") + if parsed_url.username or parsed_url.password: + raise ValueError(f"Endpoint URL ({url}) must not contain username or password.") + if parsed_url.fragment: + raise ValueError(f"Endpoint URL ({url}) must not contain a fragment.") + if parsed_url.query: + raise ValueError(f"Endpoint URL ({url}) must not contain query parameters.") + + def _load_ca_file_to_data(self) -> None: + """Load CA file content into ca_data and track modification time. + + :raises ValueError: If the CA file is empty on first read. + """ + try: + with open(self._ca_file, "r", encoding="utf-8") as f: + content = f.read() + + # Check if the file is empty + if not content: + # If no prior ca_data exists (first read), fail + if self._ca_data is None: + raise ValueError(f"CA file ({self._ca_file}) is empty. Cannot establish secure connection.") + # If we had prior ca_data, keep it (mid-rotation scenario) + return + + # File has content, update ca_data and tracking + self._ca_data = content + self._ca_file_mtime = os.path.getmtime(self._ca_file) + except (OSError, IOError) as e: + # If no prior ca_data exists (first read), fail + if self._ca_data is None: + raise ValueError(f"Failed to read CA file ({self._ca_file}): {e}") from e + # If we can't read the file, keep existing ca_data but clear mtime + # so we'll try to reload on the next change check + self._ca_file_mtime = None + + def _has_ca_file_changed(self) -> bool: + """Check if the CA file has changed since last load. + + :return: True if the CA file has changed, False otherwise. + :rtype: bool + """ + if not self._ca_file: + return False + + if not os.path.exists(self._ca_file): + # File was deleted, consider this a change if we had data before + return self._ca_data is not None or self._ca_file_mtime is not None + + try: + # Check modification time + current_mtime = os.path.getmtime(self._ca_file) + return self._ca_file_mtime != current_mtime + except (OSError, IOError): + # If we can't read the file stats, assume it changed + return True + + def _update_request_url(self, request: HttpRequest) -> None: + """Update the request URL to use proxy endpoint if configured. + + :param request: The HTTP request object to update. + :type request: ~azure.core.rest.HttpRequest + """ + if self._proxy_endpoint: + parsed_request_url = urllib.parse.urlparse(request.url) + parsed_proxy_url = urllib.parse.urlparse(self._proxy_endpoint) + combined_path = parsed_proxy_url.path.rstrip("/") + "/" + parsed_request_url.path.lstrip("/") + new_url = urllib.parse.urlunparse( + ( + parsed_proxy_url.scheme, + parsed_proxy_url.netloc, + combined_path, + parsed_request_url.params, + parsed_request_url.query, + parsed_request_url.fragment, + ) + ) + request.url = new_url diff --git a/sdk/identity/azure-identity/azure/identity/_internal/token_binding_transport_requests.py b/sdk/identity/azure-identity/azure/identity/_internal/token_binding_transport_requests.py new file mode 100644 index 000000000000..19d15abae1d9 --- /dev/null +++ b/sdk/identity/azure-identity/azure/identity/_internal/token_binding_transport_requests.py @@ -0,0 +1,61 @@ +# ------------------------------------ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. +# ------------------------------------ +""" +Requests transport class for WorkloadIdentityCredential with token proxy support. +""" +import ssl +from typing import Any, Optional + +from requests.adapters import HTTPAdapter +from requests import Session +from azure.core.pipeline.transport import ( # pylint: disable=non-abstract-transport-import, no-name-in-module + RequestsTransport, +) +from azure.core.rest import HttpRequest + +from .token_binding_transport_mixin import TokenBindingTransportMixin + + +class SNIAdapter(HTTPAdapter): + """A custom HTTPAdapter that allows setting a custom SNI hostname.""" + + def __init__(self, server_hostname: Optional[str], ca_data: Optional[str], **kwargs: Any) -> None: + self.server_hostname = server_hostname + self.ca_data = ca_data + super().__init__(**kwargs) + + def init_poolmanager(self, connections: int, maxsize: int, block: bool = False, **pool_kwargs: Any) -> None: + if self.server_hostname: + pool_kwargs["server_hostname"] = self.server_hostname + pool_kwargs["ssl_context"] = ssl.create_default_context(cadata=self.ca_data) + super().init_poolmanager(connections, maxsize, block, **pool_kwargs) + + +class CustomRequestsTransport(TokenBindingTransportMixin, RequestsTransport): + """Custom RequestsTransport with SNI and CA certificate support for WorkloadIdentityCredential.""" + + def __init__(self, *args: Any, **kwargs: Any) -> None: + self.session: Optional[Session] = None + super().__init__(*args, **kwargs) + self._update_adaptor() + + def _update_adaptor(self) -> None: + """Update the session's adapter with the current SNI and CA data.""" + if not self.session: + self.session = Session() + + adapter = SNIAdapter(self._sni, self._ca_data) + self.session.mount("https://", adapter) + + def send(self, request: HttpRequest, **kwargs: Any) -> Any: + self._update_request_url(request) + + # Check if CA file has changed and reload ca_data if needed + if self._ca_file and self._has_ca_file_changed(): + self._load_ca_file_to_data() + # If ca_data was updated, recreate SSL context with the new data + if self._ca_data: + self._update_adaptor() + return super().send(request, **kwargs) diff --git a/sdk/identity/azure-identity/azure/identity/_version.py b/sdk/identity/azure-identity/azure/identity/_version.py index f80f2f3644d7..92c78eb1a721 100644 --- a/sdk/identity/azure-identity/azure/identity/_version.py +++ b/sdk/identity/azure-identity/azure/identity/_version.py @@ -2,4 +2,4 @@ # Copyright (c) Microsoft Corporation. # Licensed under the MIT License. # ------------------------------------ -VERSION = "1.25.2" +VERSION = "1.26.0b1" diff --git a/sdk/identity/azure-identity/azure/identity/aio/_credentials/workload_identity.py b/sdk/identity/azure-identity/azure/identity/aio/_credentials/workload_identity.py index 8c44369da6ff..47fe90a7d8fc 100644 --- a/sdk/identity/azure-identity/azure/identity/aio/_credentials/workload_identity.py +++ b/sdk/identity/azure-identity/azure/identity/aio/_credentials/workload_identity.py @@ -2,11 +2,19 @@ # Copyright (c) Microsoft Corporation. # Licensed under the MIT License. # ------------------------------------ +# cspell:ignore cafile import os from typing import Any, Optional + from .client_assertion import ClientAssertionCredential -from ..._credentials.workload_identity import TokenFileMixin, WORKLOAD_CONFIG_ERROR +from ..._credentials.workload_identity import ( + TokenFileMixin, + WORKLOAD_CONFIG_ERROR, + CA_DATA_FILE_ERROR, + CUSTOM_PROXY_ENV_ERROR, +) from ..._constants import EnvironmentVariables +from ..._internal import within_credential_chain class WorkloadIdentityCredential(ClientAssertionCredential, TokenFileMixin): @@ -72,6 +80,33 @@ def __init__( assert token_file_path is not None self._token_file_path = token_file_path + + if kwargs.pop("use_token_proxy", False) and not within_credential_chain.get(): + token_proxy_endpoint = os.environ.get(EnvironmentVariables.AZURE_KUBERNETES_TOKEN_PROXY) + sni = os.environ.get(EnvironmentVariables.AZURE_KUBERNETES_SNI_NAME) + ca_file = os.environ.get(EnvironmentVariables.AZURE_KUBERNETES_CA_FILE) + ca_data = os.environ.get(EnvironmentVariables.AZURE_KUBERNETES_CA_DATA) + if token_proxy_endpoint: + if ca_file and ca_data: + raise ValueError(CA_DATA_FILE_ERROR) + + transport = _get_transport( + sni=sni, + token_proxy_endpoint=token_proxy_endpoint, + ca_file=ca_file, + ca_data=ca_data, + ) + + if transport: + kwargs["transport"] = transport + else: + raise ValueError( + "Async transport creation failed. Ensure that the aiohttp or requests package is installed to " + "enable token proxy usage in this credential." + ) + elif sni or ca_file or ca_data: + raise ValueError(CUSTOM_PROXY_ENV_ERROR) + super().__init__( tenant_id=tenant_id, client_id=client_id, @@ -79,3 +114,28 @@ def __init__( token_file_path=token_file_path, **kwargs, ) + + +def _get_transport(sni, token_proxy_endpoint, ca_file, ca_data): + try: + from .._internal.token_binding_transport_aiohttp import CustomAioHttpTransport + + return CustomAioHttpTransport( + sni=sni, + proxy_endpoint=token_proxy_endpoint, + ca_file=ca_file, + ca_data=ca_data, + ) + except ImportError: + # Fallback to async-wrapped requests transport + try: + from .._internal.token_binding_transport_asyncio import CustomAsyncioRequestsTransport + + return CustomAsyncioRequestsTransport( + sni=sni, + proxy_endpoint=token_proxy_endpoint, + ca_file=ca_file, + ca_data=ca_data, + ) + except ImportError: + return None diff --git a/sdk/identity/azure-identity/azure/identity/aio/_internal/token_binding_transport_aiohttp.py b/sdk/identity/azure-identity/azure/identity/aio/_internal/token_binding_transport_aiohttp.py new file mode 100644 index 000000000000..1a21c5b13d54 --- /dev/null +++ b/sdk/identity/azure-identity/azure/identity/aio/_internal/token_binding_transport_aiohttp.py @@ -0,0 +1,37 @@ +# ------------------------------------ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. +# ------------------------------------ +""" +Aiohttp transport class for the asynchronous WorkloadIdentityCredential with token proxy support. +""" +import ssl +from typing import Any + +from azure.core.pipeline.transport import ( # pylint: disable=non-abstract-transport-import, no-name-in-module + AioHttpTransport, +) +from azure.core.rest import HttpRequest + +from ..._internal.token_binding_transport_mixin import TokenBindingTransportMixin + + +class CustomAioHttpTransport(TokenBindingTransportMixin, AioHttpTransport): + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self._ssl_context = ssl.create_default_context(cadata=self._ca_data) + + async def send(self, request: HttpRequest, **kwargs: Any) -> Any: + self._update_request_url(request) + kwargs.setdefault("server_hostname", self._sni) + + # Check if CA file has changed and reload ca_data if needed + if self._ca_file and self._has_ca_file_changed(): + self._load_ca_file_to_data() + # If ca_data was updated, recreate SSL context with the new data + if self._ca_data: + self._ssl_context = ssl.create_default_context(cadata=self._ca_data) + + kwargs.setdefault("ssl", self._ssl_context) + return await super().send(request, **kwargs) diff --git a/sdk/identity/azure-identity/azure/identity/aio/_internal/token_binding_transport_asyncio.py b/sdk/identity/azure-identity/azure/identity/aio/_internal/token_binding_transport_asyncio.py new file mode 100644 index 000000000000..fa236c2e4d4f --- /dev/null +++ b/sdk/identity/azure-identity/azure/identity/aio/_internal/token_binding_transport_asyncio.py @@ -0,0 +1,45 @@ +# ------------------------------------ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. +# ------------------------------------ +""" +Asyncio Requests transport class for the asynchronous WorkloadIdentityCredential with token proxy support. +""" +from typing import Any, Optional + +from requests import Session +from azure.core.pipeline.transport import ( # pylint: disable=non-abstract-transport-import, no-name-in-module + AsyncioRequestsTransport, +) +from azure.core.rest import HttpRequest + +from ..._internal.token_binding_transport_mixin import TokenBindingTransportMixin +from ..._internal.token_binding_transport_requests import SNIAdapter + + +class CustomAsyncioRequestsTransport(TokenBindingTransportMixin, AsyncioRequestsTransport): + + def __init__(self, *args, **kwargs): + self.session: Optional[Session] = None + super().__init__(*args, **kwargs) + self._update_adaptor() + + def _update_adaptor(self) -> None: + """Update the session's adapter with the current SNI and CA data.""" + if not self.session: + self.session = Session() + + adapter = SNIAdapter(self._sni, self._ca_data) + self.session.mount("https://", adapter) + + async def send(self, request: HttpRequest, **kwargs: Any) -> Any: + self._update_request_url(request) + + # Check if CA file has changed and reload ca_data if needed + if self._ca_file and self._has_ca_file_changed(): + self._load_ca_file_to_data() + # If ca_data was updated, recreate SSL context with the new data + if self._ca_data: + self._update_adaptor() + + return await super().send(request, **kwargs) diff --git a/sdk/identity/azure-identity/pyproject.toml b/sdk/identity/azure-identity/pyproject.toml index 5d61880b9f88..c9fa89cb338a 100644 --- a/sdk/identity/azure-identity/pyproject.toml +++ b/sdk/identity/azure-identity/pyproject.toml @@ -12,7 +12,7 @@ keywords = ["azure", "azure sdk"] requires-python = ">=3.9" license = "MIT" classifiers = [ - "Development Status :: 5 - Production/Stable", + "Development Status :: 4 - Beta", "Programming Language :: Python", "Programming Language :: Python :: 3 :: Only", "Programming Language :: Python :: 3", @@ -47,4 +47,4 @@ pytyped = ["py.typed"] [tool.azure-sdk-build] pyright = false verifytypes = true -black = true \ No newline at end of file +black = true diff --git a/sdk/identity/azure-identity/tests/proxy_server.py b/sdk/identity/azure-identity/tests/proxy_server.py new file mode 100644 index 000000000000..a6e6e1e3ae61 --- /dev/null +++ b/sdk/identity/azure-identity/tests/proxy_server.py @@ -0,0 +1,328 @@ +# ------------------------------------ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. +# ------------------------------------ +# cspell:ignore ests +""" +Local test server for token binding proxy testing. + +This server simulates a token proxy that can: +1. Accept HTTPS requests with custom SNI and CA certificates +2. Route requests to downstream services +3. Handle various error scenarios for testing +4. Support certificate rotation scenarios +""" + +import argparse +import ipaddress +import json +import logging +import os +import ssl +import tempfile +import threading +import time +import datetime +from http.server import HTTPServer, BaseHTTPRequestHandler +from socketserver import ThreadingMixIn +import uuid + +from cryptography import x509 +from cryptography.x509.oid import NameOID +from cryptography.hazmat.primitives import hashes, serialization +from cryptography.hazmat.primitives.asymmetric import rsa + + +class TokenProxyHandler(BaseHTTPRequestHandler): + """HTTP request handler that simulates a token proxy.""" + + def log_message(self, format, *args): + """Override to use proper logging instead of stderr.""" + logging.info(f"{self.address_string()} - {format % args}") + + def do_GET(self): + """Handle GET requests.""" + self._handle_request() + + def do_POST(self): + """Handle POST requests.""" + self._handle_request() + + def do_PUT(self): + """Handle PUT requests.""" + self._handle_request() + + def do_DELETE(self): + """Handle DELETE requests.""" + self._handle_request() + + def do_PATCH(self): + """Handle PATCH requests.""" + self._handle_request() + + def _handle_request(self): + """Common request handling logic.""" + path = self.path + headers = dict(self.headers) + + # Read request body if present + content_length = int(headers.get("content-length", 0)) + body = self.rfile.read(content_length) if content_length > 0 else b"" + + logging.info(f"Received {self.command} {path}") + logging.info(f"Headers: {headers}") + + # Simulate different responses based on path + if path == "/health": + self._send_health_response() + elif path.endswith("/oauth2/v2.0/token"): + self._send_token_response(body) + elif path == "/error/500": + self._send_error_response(500, "Internal Server Error") + elif path == "/error/ssl": + # Simulate SSL error by closing connection + self.wfile.close() + return + else: + self._send_proxy_response(path, body, headers) + + def _send_health_response(self): + """Send a health check response.""" + response = {"status": "healthy", "timestamp": time.time(), "server": "token-proxy-test-server"} + self._send_json_response(response) + + def _send_token_response(self, body): + """Send a mock OAuth token response.""" + # Parse the request body to extract grant type, etc. + try: + if body: + body_str = body.decode("utf-8") + logging.info(f"Token request body: {body_str}") + except Exception as e: + logging.warning(f"Could not decode request body: {e}") + + # Mock token response + response = { + "access_token": f"mock_token_{uuid.uuid4().hex[:16]}", + "token_type": "Bearer", + "expires_in": 3600, + "scope": "https://graph.microsoft.com/.default", + } + self._send_json_response(response) + + def _send_proxy_response(self, path, body, headers): + """Send a generic proxy response.""" + response = { + "proxied_path": path, + "method": self.command, + "headers_received": dict(headers), + "body_length": len(body), + "proxy_server": "token-proxy-test-server", + } + self._send_json_response(response) + + def _send_json_response(self, data, status_code=200): + """Send a JSON response.""" + response_json = json.dumps(data, indent=2) + + self.send_response(status_code) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(response_json))) + self.end_headers() + + self.wfile.write(response_json.encode("utf-8")) + + def _send_error_response(self, status_code, message): + """Send an error response.""" + self.send_response(status_code) + self.send_header("Content-Type", "text/plain") + self.send_header("Content-Length", str(len(message))) + self.end_headers() + + self.wfile.write(message.encode("utf-8")) + + +class ThreadedHTTPServer(ThreadingMixIn, HTTPServer): + """Threaded HTTP server for handling multiple concurrent requests.""" + + allow_reuse_address = True + daemon_threads = True + + +class TokenProxyTestServer: + """Test server that can be configured with SSL/TLS and custom certificates.""" + + def __init__(self, host="localhost", port=0, use_ssl=True): + self.host = host + self.port = port + self.use_ssl = use_ssl + self.server = None + self.server_thread = None + self.cert_file = None + self.key_file = None + self.ca_file = None + self._temp_files = [] + + def generate_test_certificates(self): + """Generate self-signed certificates for testing.""" + + # Generate private key + private_key = rsa.generate_private_key( + public_exponent=65537, + key_size=2048, + ) + + # Create certificate + subject = issuer = x509.Name( + [ + x509.NameAttribute(NameOID.COUNTRY_NAME, "US"), + x509.NameAttribute(NameOID.STATE_OR_PROVINCE_NAME, "Test"), + x509.NameAttribute(NameOID.LOCALITY_NAME, "Test"), + x509.NameAttribute(NameOID.ORGANIZATION_NAME, "Test Proxy Server"), + x509.NameAttribute(NameOID.COMMON_NAME, self.host), + ] + ) + + cert = ( + x509.CertificateBuilder() + .subject_name(subject) + .issuer_name(issuer) + .public_key(private_key.public_key()) + .serial_number(x509.random_serial_number()) + .not_valid_before(datetime.datetime.now(datetime.timezone.utc)) + .not_valid_after(datetime.datetime.now(datetime.timezone.utc) + datetime.timedelta(days=365)) + .add_extension( + x509.SubjectAlternativeName( + [ + x509.DNSName(self.host), + x509.DNSName("localhost"), + x509.DNSName("1234.ests.aks"), + x509.IPAddress(ipaddress.IPv4Address("127.0.0.1")), + ] + ), + critical=False, + ) + .sign(private_key, hashes.SHA256()) + ) + + # Write certificate to temp file + cert_fd, self.cert_file = tempfile.mkstemp(suffix=".pem", prefix="test_cert_") + with os.fdopen(cert_fd, "wb") as f: + f.write(cert.public_bytes(serialization.Encoding.PEM)) + self._temp_files.append(self.cert_file) + + # Write private key to temp file + key_fd, self.key_file = tempfile.mkstemp(suffix=".pem", prefix="test_key_") + with os.fdopen(key_fd, "wb") as f: + f.write( + private_key.private_bytes( + encoding=serialization.Encoding.PEM, + format=serialization.PrivateFormat.PKCS8, + encryption_algorithm=serialization.NoEncryption(), + ) + ) + self._temp_files.append(self.key_file) + + # Use the same cert as CA for simplicity + self.ca_file = self.cert_file + + logging.info(f"Generated test certificate: {self.cert_file}") + logging.info(f"Generated test key: {self.key_file}") + + def start(self): + """Start the test server.""" + if self.use_ssl: + self.generate_test_certificates() + + # Create server + self.server = ThreadedHTTPServer((self.host, self.port), TokenProxyHandler) + + if self.use_ssl: + # Configure SSL context + context = ssl.create_default_context(ssl.Purpose.CLIENT_AUTH) + if self.cert_file and self.key_file: + context.load_cert_chain(self.cert_file, self.key_file) + + self.server.socket = context.wrap_socket(self.server.socket, server_side=True) + + # Update port if it was 0 (auto-assigned) + self.port = self.server.server_address[1] + + # Start server in background thread + self.server_thread = threading.Thread(target=self.server.serve_forever, daemon=True) + self.server_thread.start() + + scheme = "https" if self.use_ssl else "http" + logging.info(f"Test server started at {scheme}://{self.host}:{self.port}") + + return f"{scheme}://{self.host}:{self.port}" + + def stop(self): + """Stop the test server and clean up.""" + if self.server: + self.server.shutdown() + self.server.server_close() + + if self.server_thread: + self.server_thread.join(timeout=5) + + # Clean up temporary files + for temp_file in self._temp_files: + try: + os.unlink(temp_file) + except OSError: + pass + self._temp_files.clear() # Clear the list after cleanup + + logging.info("Test server stopped and cleaned up") + + def __enter__(self): + """Context manager entry.""" + self.start() + return self + + def __exit__(self, exc_type, exc_val, exc_tb): + """Context manager exit.""" + self.stop() + + @property + def base_url(self): + """Get the base URL of the server.""" + scheme = "https" if self.use_ssl else "http" + return f"{scheme}://{self.host}:{self.port}" + + +def main(): + """Run the test server standalone.""" + parser = argparse.ArgumentParser(description="Token Proxy Test Server") + parser.add_argument("--host", default="localhost", help="Host to bind to") + parser.add_argument("--port", type=int, default=8443, help="Port to bind to") + parser.add_argument("--no-ssl", action="store_true", help="Disable SSL/TLS") + parser.add_argument("--verbose", "-v", action="store_true", help="Enable verbose logging") + + args = parser.parse_args() + + # Configure logging + log_level = logging.DEBUG if args.verbose else logging.INFO + logging.basicConfig(level=log_level, format="%(asctime)s - %(levelname)s - %(message)s") + + # Start server + with TokenProxyTestServer(host=args.host, port=args.port, use_ssl=not args.no_ssl) as server: + print(f"Server running at {server.base_url}") + print("Available endpoints:") + print(f" {server.base_url}/health - Health check") + print(f" {server.base_url}/oauth2/v2.0/token - Mock OAuth token endpoint") + print(f" {server.base_url}/error/500 - Simulate server error") + print(f" {server.base_url}/error/ssl - Simulate SSL error") + print(f" {server.base_url}/ - Generic proxy response") + print("\nPress Ctrl+C to stop") + + try: + while True: + time.sleep(1) + except KeyboardInterrupt: + print("\nShutting down...") + + +if __name__ == "__main__": + main() diff --git a/sdk/identity/azure-identity/tests/test_workload_identity_credential.py b/sdk/identity/azure-identity/tests/test_workload_identity_credential.py index 1db0874a77ce..1bcdcefbe06d 100644 --- a/sdk/identity/azure-identity/tests/test_workload_identity_credential.py +++ b/sdk/identity/azure-identity/tests/test_workload_identity_credential.py @@ -2,12 +2,30 @@ # Copyright (c) Microsoft Corporation. # Licensed under the MIT License. # ------------------------------------ +# cspell:ignore cafile ests +import os +import tempfile +import time from unittest.mock import mock_open, MagicMock, patch import pytest +from azure.core.rest import HttpRequest +from azure.core.exceptions import ServiceRequestError, ServiceResponseError from azure.identity import WorkloadIdentityCredential +from azure.identity._credentials.workload_identity import _get_transport from helpers import mock_response, build_aad_response, GET_TOKEN_METHODS +from proxy_server import TokenProxyTestServer + + +PEM_CERT_PATH = os.path.join(os.path.dirname(__file__), "certificate.pem") + + +@pytest.fixture(scope="module") +def ca_data() -> str: + """Read CA certificate data from a PEM file for testing.""" + with open(PEM_CERT_PATH, "r", encoding="utf-8") as f: + return f.read() def test_workload_identity_credential_initialize(): @@ -44,3 +62,540 @@ def send(request, **kwargs): assert token.token == access_token open_mock.assert_called_once_with(token_file_path, encoding="utf-8") + + +class TestWorkloadIdentityCredentialTokenProxy: + """Test cases for WorkloadIdentityCredential with use_token_proxy=True.""" + + def test_use_token_proxy_creates_custom_transport(self): + """Test that use_token_proxy=True creates a custom transport with correct parameters.""" + tenant_id = "tenant-id" + client_id = "client-id" + token_file_path = "foo-path" + proxy_endpoint = "https://proxy.example.com:8080" + sni_hostname = "sni.example.com" + ca_file_path = "/path/to/ca.pem" + + env_vars = { + "AZURE_KUBERNETES_TOKEN_PROXY": proxy_endpoint, + "AZURE_KUBERNETES_SNI_NAME": sni_hostname, + "AZURE_KUBERNETES_CA_FILE": ca_file_path, + } + + with patch.dict(os.environ, env_vars, clear=False): + with patch("azure.identity._credentials.workload_identity._get_transport") as mock_get_transport: + mock_transport_instance = MagicMock() + mock_get_transport.return_value = mock_transport_instance + + WorkloadIdentityCredential( + tenant_id=tenant_id, + client_id=client_id, + token_file_path=token_file_path, + use_token_proxy=True, + ) + + mock_get_transport.assert_called_once_with( + sni=sni_hostname, + token_proxy_endpoint=proxy_endpoint, + ca_file=ca_file_path, + ca_data=None, + ) + + def test_use_token_proxy_with_ca_data(self): + """Test use_token_proxy with CA data instead of CA file.""" + tenant_id = "tenant-id" + client_id = "client-id" + token_file_path = "foo-path" + proxy_endpoint = "https://proxy.example.com:8080" + ca_data = "-----BEGIN CERTIFICATE-----\nTest CA data\n-----END CERTIFICATE-----" + + env_vars = { + "AZURE_KUBERNETES_TOKEN_PROXY": proxy_endpoint, + "AZURE_KUBERNETES_CA_DATA": ca_data, + } + + with patch.dict(os.environ, env_vars, clear=False): + with patch("azure.identity._credentials.workload_identity._get_transport") as mock_get_transport: + mock_transport_instance = MagicMock() + mock_get_transport.return_value = mock_transport_instance + + WorkloadIdentityCredential( + tenant_id=tenant_id, + client_id=client_id, + token_file_path=token_file_path, + use_token_proxy=True, + ) + + mock_get_transport.assert_called_once_with( + sni=None, + token_proxy_endpoint=proxy_endpoint, + ca_file=None, + ca_data=ca_data, + ) + + def test_use_token_proxy_minimal_config(self): + """Test use_token_proxy with minimal configuration (only proxy endpoint).""" + tenant_id = "tenant-id" + client_id = "client-id" + token_file_path = "foo-path" + proxy_endpoint = "https://proxy.example.com:8080" + + env_vars = { + "AZURE_KUBERNETES_TOKEN_PROXY": proxy_endpoint, + } + + with patch.dict(os.environ, env_vars, clear=False): + with patch("azure.identity._credentials.workload_identity._get_transport") as mock_get_transport: + mock_transport_instance = MagicMock() + mock_get_transport.return_value = mock_transport_instance + + WorkloadIdentityCredential( + tenant_id=tenant_id, + client_id=client_id, + token_file_path=token_file_path, + use_token_proxy=True, + ) + + mock_get_transport.assert_called_once_with( + sni=None, + token_proxy_endpoint=proxy_endpoint, + ca_file=None, + ca_data=None, + ) + + def test_use_token_proxy_missing_proxy_endpoint(self): + """Test that use_token_proxy=True without proxy endpoint uses the normal transport.""" + tenant_id = "tenant-id" + client_id = "client-id" + token_file_path = "foo-path" + + # Ensure proxy endpoint env var is not set + with patch.dict(os.environ, {}, clear=False): + if "AZURE_KUBERNETES_TOKEN_PROXY" in os.environ: + del os.environ["AZURE_KUBERNETES_TOKEN_PROXY"] + + with patch("azure.identity._credentials.workload_identity._get_transport") as mock_get_transport: + WorkloadIdentityCredential( + tenant_id=tenant_id, + client_id=client_id, + token_file_path=token_file_path, + use_token_proxy=True, + ) + mock_get_transport.assert_not_called() + + def test_use_token_proxy_both_ca_file_and_data_raises_error(self): + """Test that setting both CA file and CA data raises ValueError.""" + tenant_id = "tenant-id" + client_id = "client-id" + token_file_path = "foo-path" + proxy_endpoint = "https://proxy.example.com:8080" + ca_file_path = "/path/to/ca.pem" + ca_data = "-----BEGIN CERTIFICATE-----\nTest CA data\n-----END CERTIFICATE-----" + + env_vars = { + "AZURE_KUBERNETES_TOKEN_PROXY": proxy_endpoint, + "AZURE_KUBERNETES_CA_FILE": ca_file_path, + "AZURE_KUBERNETES_CA_DATA": ca_data, + } + + with patch.dict(os.environ, env_vars, clear=False): + with pytest.raises(ValueError, match="Both AZURE_KUBERNETES_CA_FILE and AZURE_KUBERNETES_CA_DATA are set"): + WorkloadIdentityCredential( + tenant_id=tenant_id, + client_id=client_id, + token_file_path=token_file_path, + use_token_proxy=True, + ) + + def test_use_token_proxy_missing_endpoint_with_custom_env_vars_raises_error(self): + """Test that use_token_proxy=True without proxy endpoint but with other custom env vars raises ValueError.""" + tenant_id = "tenant-id" + client_id = "client-id" + token_file_path = "foo-path" + sni_hostname = "sni.example.com" + ca_file_path = "/path/to/ca.pem" + ca_data = "-----BEGIN CERTIFICATE-----\nTest CA data\n-----END CERTIFICATE-----" + + # Ensure proxy endpoint is not set + if "AZURE_KUBERNETES_TOKEN_PROXY" in os.environ: + del os.environ["AZURE_KUBERNETES_TOKEN_PROXY"] + + # Test with SNI set but no proxy endpoint + env_vars_sni = { + "AZURE_KUBERNETES_SNI_NAME": sni_hostname, + } + with patch.dict(os.environ, env_vars_sni, clear=False): + with pytest.raises(ValueError): + WorkloadIdentityCredential( + tenant_id=tenant_id, + client_id=client_id, + token_file_path=token_file_path, + use_token_proxy=True, + ) + + # Test with CA file set but no proxy endpoint + env_vars_ca_file = { + "AZURE_KUBERNETES_CA_FILE": ca_file_path, + } + with patch.dict(os.environ, env_vars_ca_file, clear=False): + with pytest.raises(ValueError): + WorkloadIdentityCredential( + tenant_id=tenant_id, + client_id=client_id, + token_file_path=token_file_path, + use_token_proxy=True, + ) + + # Test with CA data set but no proxy endpoint + env_vars_ca_data = { + "AZURE_KUBERNETES_CA_DATA": ca_data, + } + with patch.dict(os.environ, env_vars_ca_data, clear=False): + with pytest.raises(ValueError): + WorkloadIdentityCredential( + tenant_id=tenant_id, + client_id=client_id, + token_file_path=token_file_path, + use_token_proxy=True, + ) + + def test_use_token_proxy_false_does_not_create_transport(self): + """Test that use_token_proxy=False (default) does not create a custom transport.""" + tenant_id = "tenant-id" + client_id = "client-id" + token_file_path = "foo-path" + + with patch("azure.identity._credentials.workload_identity._get_transport") as mock_get_transport: + WorkloadIdentityCredential( + tenant_id=tenant_id, + client_id=client_id, + token_file_path=token_file_path, + use_token_proxy=False, + ) + mock_get_transport.assert_not_called() + + +class TestCustomRequestsTransport: + """Test cases for the custom RequestsTransport used by WorkloadIdentityCredential.""" + + def test_get_transport_creates_custom_requests_transport(self, ca_data): + """Test that _get_transport creates CustomRequestsTransport with correct parameters.""" + sni = "test.sni.com" + proxy_endpoint = "https://proxy.example.com:8080" + ca_file = PEM_CERT_PATH + + transport = _get_transport( + sni=sni, + token_proxy_endpoint=proxy_endpoint, + ca_file=ca_file, + ca_data=None, + ) + + assert transport is not None + assert hasattr(transport, "_sni") + assert hasattr(transport, "_proxy_endpoint") + assert hasattr(transport, "_ca_file") + assert hasattr(transport, "_ca_data") + assert transport._sni == sni + assert transport._proxy_endpoint == proxy_endpoint + assert transport._ca_file == ca_file + assert transport._ca_data == ca_data + + def test_get_transport_with_minimal_config(self): + """Test _get_transport with minimal configuration.""" + proxy_endpoint = "https://proxy.example.com:8080" + + transport = _get_transport( + sni=None, + token_proxy_endpoint=proxy_endpoint, + ca_file=None, + ca_data=None, + ) + + assert transport is not None + assert transport._sni is None + assert transport._proxy_endpoint == proxy_endpoint + assert transport._ca_file is None + assert transport._ca_data is None + + def test_custom_requests_transport_inherits_from_token_binding_mixin(self): + """Test that CustomRequestsTransport inherits from TokenBindingTransportMixin.""" + transport = _get_transport( + sni="test.sni.com", + token_proxy_endpoint="https://proxy.example.com:8080", + ca_file=None, + ca_data=None, + ) + + assert transport is not None + + # Verify inheritance from TokenBindingTransportMixin + from azure.identity._internal.token_binding_transport_mixin import TokenBindingTransportMixin + + assert isinstance(transport, TokenBindingTransportMixin) + + # Verify TokenBindingTransportMixin methods are available + assert hasattr(transport, "_update_request_url") + assert hasattr(transport, "_has_ca_file_changed") + assert hasattr(transport, "_load_ca_file_to_data") + assert hasattr(transport, "_validate_url") + + +class TestCustomRequestsTransportWithLocalServer: + """Integration tests using a local test server for CustomRequestsTransport.""" + + def test_basic_https_request(self): + """Test basic HTTPS request to test server.""" + with TokenProxyTestServer(use_ssl=True) as server: + # Create transport with server's CA certificate + transport = _get_transport(sni=None, token_proxy_endpoint=None, ca_file=server.ca_file, ca_data=None) + assert transport is not None + request = HttpRequest("GET", f"{server.base_url}/health") + + response = transport.send(request) + + assert response.status_code == 200 + data = response.json() + assert data["status"] == "healthy" + assert "timestamp" in data + + def test_post_request_with_body(self): + """Test POST request with request body.""" + with TokenProxyTestServer(use_ssl=True) as server: + transport = _get_transport(sni=None, token_proxy_endpoint=None, ca_file=server.ca_file, ca_data=None) + assert transport is not None + + # Prepare OAuth-like request + body = "grant_type=client_credentials&scope=https://graph.microsoft.com/.default" + request = HttpRequest( + "POST", + f"{server.base_url}/tenant/oauth2/v2.0/token", + headers={"Content-Type": "application/x-www-form-urlencoded"}, + data=body.encode("utf-8"), + ) + + response = transport.send(request) + + assert response.status_code == 200 + data = response.json() + assert "access_token" in data + assert data["token_type"] == "Bearer" + assert data["expires_in"] == 3600 + + def test_proxy_endpoint_comprehensive(self): + """Test comprehensive proxy endpoint functionality with various HTTP methods and scenarios.""" + with TokenProxyTestServer(use_ssl=True) as server: + # Configure transport with proxy endpoint + transport = _get_transport( + sni=None, token_proxy_endpoint=server.base_url, ca_file=server.ca_file, ca_data=None + ) + assert transport is not None + + # Test 1: POST request with JSON body through proxy + post_data = {"grant_type": "client_credentials", "scope": "https://graph.microsoft.com/.default"} + post_request = HttpRequest( + "POST", + "https://login.microsoftonline.com/tenant/oauth2/v2.0/token2", + headers={"Content-Type": "application/json"}, + json=post_data, + ) + + post_response = transport.send(post_request) + assert post_response.status_code == 200 + post_data_response = post_response.json() + assert post_data_response["method"] == "POST" + assert post_data_response["proxied_path"] == "/tenant/oauth2/v2.0/token2" + + # Test 2: PUT request through proxy + put_request = HttpRequest( + "PUT", + "https://graph.microsoft.com/v1.0/me/profile", + headers={"Content-Type": "application/json"}, + json={"displayName": "Test User"}, + ) + + put_response = transport.send(put_request) + assert put_response.status_code == 200 + put_data_response = put_response.json() + assert put_data_response["method"] == "PUT" + assert put_data_response["proxied_path"] == "/v1.0/me/profile" + + # Test 3: DELETE request through proxy + delete_request = HttpRequest("DELETE", "https://graph.microsoft.com/v1.0/applications/app-id") + + delete_response = transport.send(delete_request) + assert delete_response.status_code == 200 + delete_data_response = delete_response.json() + assert delete_data_response["method"] == "DELETE" + assert delete_data_response["proxied_path"] == "/v1.0/applications/app-id" + + # Test 4: Complex URL with multiple path segments and query parameters + complex_url = "https://management.azure.com/subscriptions/sub-id/resourceGroups/rg/providers/Microsoft.Storage/storageAccounts/account?api-version=2021-04-01&expand=properties" + complex_request = HttpRequest("GET", complex_url) + + complex_response = transport.send(complex_request) + assert complex_response.status_code == 200 + complex_data_response = complex_response.json() + assert complex_data_response["method"] == "GET" + expected_path = "/subscriptions/sub-id/resourceGroups/rg/providers/Microsoft.Storage/storageAccounts/account?api-version=2021-04-01&expand=properties" + assert complex_data_response["proxied_path"] == expected_path + + # Test 5: Request with custom headers through proxy + headers_request = HttpRequest( + "GET", + "https://vault.azure.net/secrets/test-secret?api-version=7.3", + headers={ + "Authorization": "Bearer test-token", + "X-Custom-Header": "proxy-test-value", + "User-Agent": "Azure-SDK-For-Python", + }, + ) + + headers_response = transport.send(headers_request) + assert headers_response.status_code == 200 + headers_data_response = headers_response.json() + assert headers_data_response["method"] == "GET" + assert headers_data_response["proxied_path"] == "/secrets/test-secret?api-version=7.3" + + # Verify headers were forwarded through proxy + received_headers = headers_data_response["headers_received"] + assert "Authorization" in received_headers + assert "X-Custom-Header" in received_headers + assert received_headers["Authorization"] == "Bearer test-token" + assert received_headers["X-Custom-Header"] == "proxy-test-value" + + def test_sni_with_custom_hostname(self): + """Test SNI (Server Name Indication) with custom hostname.""" + with TokenProxyTestServer(use_ssl=True) as server: + # Use SNI with a different hostname than the server + transport = _get_transport( + sni="1234.ests.aks", token_proxy_endpoint=None, ca_file=server.ca_file, ca_data=None + ) + assert transport is not None + + request = HttpRequest("GET", f"{server.base_url}/health") + response = transport.send(request) + + assert response.status_code == 200 + data = response.json() + assert data["status"] == "healthy" + + # Check an invalid SNI hostname + transport = _get_transport( + sni="unmatched.sni.hostname", token_proxy_endpoint=None, ca_file=server.ca_file, ca_data=None + ) + assert transport is not None + request = HttpRequest("GET", f"{server.base_url}/health") + with pytest.raises(ServiceRequestError): + transport.send(request) + + def test_ca_file_change_detection(self): + """Test CA file change detection with real certificates.""" + with TokenProxyTestServer(use_ssl=True) as server: + # Create a copy of the CA file that we can modify + ca_file = server.ca_file + if ca_file is None: + pytest.skip("CA file not available") + + assert ca_file is not None # Type hint for mypy + with tempfile.NamedTemporaryFile(mode="w", delete=False, suffix=".pem") as temp_ca: + with open(ca_file, "r") as original: + original_content = original.read() + temp_ca.write(original_content) + temp_ca_path = temp_ca.name + + try: + transport = _get_transport(sni=None, token_proxy_endpoint=None, ca_file=temp_ca_path, ca_data=None) + assert transport is not None + + # First request should work + request = HttpRequest("GET", f"{server.base_url}/health") + response1 = transport.send(request) + assert response1.status_code == 200 + + # Modify the CA file (add some content) + time.sleep(0.1) + with open(temp_ca_path, "a") as f: + f.write("\n# Modified for testing\n") + + # Second request should still work (using the same cert content) + response2 = transport.send(request) + assert transport._ca_data != original_content + assert "Modified for testing" in transport._ca_data + assert response2.status_code == 200 + + finally: + os.unlink(temp_ca_path) + + def test_ssl_error_handling(self): + """Test SSL error handling.""" + with TokenProxyTestServer(use_ssl=True) as server: + # Create transport without proper CA file (will cause SSL error) + transport = _get_transport(sni=None, token_proxy_endpoint=None, ca_file=None, ca_data=None) + assert transport is not None + + request = HttpRequest("GET", f"{server.base_url}/health") + + # Should raise SSL-related error + with pytest.raises((ServiceRequestError, ServiceResponseError)): + transport.send(request) + + def test_server_error_response(self): + """Test handling of server error responses.""" + with TokenProxyTestServer(use_ssl=True) as server: + transport = _get_transport(sni=None, token_proxy_endpoint=None, ca_file=server.ca_file, ca_data=None) + assert transport is not None + + request = HttpRequest("GET", f"{server.base_url}/error/500") + + response = transport.send(request) + + assert response.status_code == 500 + # Should not raise exception, just return error response + + def test_custom_headers_preserved(self): + """Test that custom headers are preserved and sent to server.""" + with TokenProxyTestServer(use_ssl=True) as server: + transport = _get_transport(sni=None, token_proxy_endpoint=None, ca_file=server.ca_file, ca_data=None) + assert transport is not None + + custom_headers = { + "Authorization": "Bearer test-token", + "User-Agent": "CustomRequestsTransport/1.0", + "X-Custom-Header": "test-value", + } + + request = HttpRequest("GET", f"{server.base_url}/proxy/test", headers=custom_headers) + + response = transport.send(request) + + assert response.status_code == 200 + data = response.json() + + # Server echoes back the headers it received + received_headers = data["headers_received"] + assert "Authorization" in received_headers + assert "User-Agent" in received_headers + assert "X-Custom-Header" in received_headers + assert received_headers["Authorization"] == "Bearer test-token" + + def test_query_parameters_preserved(self): + """Test that query parameters are preserved in proxy requests.""" + with TokenProxyTestServer(use_ssl=True) as server: + transport = _get_transport( + sni=None, token_proxy_endpoint=server.base_url, ca_file=server.ca_file, ca_data=None + ) + assert transport is not None + + # Request with query parameters + original_url = "https://example.com/api/data?scope=read&limit=10&format=json" + request = HttpRequest("GET", original_url) + + response = transport.send(request) + + assert response.status_code == 200 + data = response.json() + + # The path should include the query parameters + expected_path = "/api/data?scope=read&limit=10&format=json" + assert data["proxied_path"] == expected_path diff --git a/sdk/identity/azure-identity/tests/test_workload_identity_credential_async.py b/sdk/identity/azure-identity/tests/test_workload_identity_credential_async.py index fc6f3e8c6cb5..6580af09648b 100644 --- a/sdk/identity/azure-identity/tests/test_workload_identity_credential_async.py +++ b/sdk/identity/azure-identity/tests/test_workload_identity_credential_async.py @@ -2,12 +2,34 @@ # Copyright (c) Microsoft Corporation. # Licensed under the MIT License. # ------------------------------------ +# cspell:ignore cafile aexit ests +import os +import tempfile +import time +import asyncio +from time import sleep as real_sleep from unittest.mock import mock_open, patch, MagicMock import pytest +from azure.core.rest import HttpRequest +from azure.core.exceptions import ServiceRequestError, ServiceResponseError from azure.identity.aio import WorkloadIdentityCredential +from azure.identity.aio._credentials.workload_identity import _get_transport +from azure.identity.aio._internal.token_binding_transport_aiohttp import CustomAioHttpTransport +from azure.identity.aio._internal.token_binding_transport_asyncio import CustomAsyncioRequestsTransport from helpers import mock_response, build_aad_response, GET_TOKEN_METHODS +from proxy_server import TokenProxyTestServer + + +PEM_CERT_PATH = os.path.join(os.path.dirname(__file__), "certificate.pem") + + +@pytest.fixture(scope="module") +def ca_data() -> str: + """Read CA certificate data from a PEM file for testing.""" + with open(PEM_CERT_PATH, "r", encoding="utf-8") as f: + return f.read() def test_workload_identity_credential_initialize(): @@ -45,3 +67,810 @@ async def send(request, **kwargs): assert token.token == access_token open_mock.assert_called_once_with(token_file_path, encoding="utf-8") + + +class TestWorkloadIdentityCredentialTokenProxyAsync: + """Async test cases for WorkloadIdentityCredential with use_token_proxy=True.""" + + def test_use_token_proxy_creates_custom_aiohttp_transport(self): + """Test that use_token_proxy=True creates a custom aiohttp transport with correct parameters.""" + tenant_id = "tenant-id" + client_id = "client-id" + token_file_path = "foo-path" + proxy_endpoint = "https://proxy.example.com:8080" + sni_hostname = "sni.example.com" + ca_file_path = "/path/to/ca.pem" + + env_vars = { + "AZURE_KUBERNETES_TOKEN_PROXY": proxy_endpoint, + "AZURE_KUBERNETES_SNI_NAME": sni_hostname, + "AZURE_KUBERNETES_CA_FILE": ca_file_path, + } + + with patch.dict(os.environ, env_vars, clear=False): + with patch("azure.identity.aio._credentials.workload_identity._get_transport") as mock_get_transport: + mock_transport_instance = MagicMock() + mock_get_transport.return_value = mock_transport_instance + + WorkloadIdentityCredential( + tenant_id=tenant_id, + client_id=client_id, + token_file_path=token_file_path, + use_token_proxy=True, + ) + + mock_get_transport.assert_called_once_with( + sni=sni_hostname, + token_proxy_endpoint=proxy_endpoint, + ca_file=ca_file_path, + ca_data=None, + ) + + def test_use_token_proxy_with_ca_data(self): + """Test use_token_proxy with CA data instead of CA file.""" + tenant_id = "tenant-id" + client_id = "client-id" + token_file_path = "foo-path" + proxy_endpoint = "https://proxy.example.com:8080" + ca_data = "-----BEGIN CERTIFICATE-----\nTest CA data\n-----END CERTIFICATE-----" + + env_vars = { + "AZURE_KUBERNETES_TOKEN_PROXY": proxy_endpoint, + "AZURE_KUBERNETES_CA_DATA": ca_data, + } + + with patch.dict(os.environ, env_vars, clear=False): + with patch("azure.identity.aio._credentials.workload_identity._get_transport") as mock_get_transport: + mock_transport_instance = MagicMock() + mock_get_transport.return_value = mock_transport_instance + + WorkloadIdentityCredential( + tenant_id=tenant_id, + client_id=client_id, + token_file_path=token_file_path, + use_token_proxy=True, + ) + mock_get_transport.assert_called_once_with( + sni=None, + token_proxy_endpoint=proxy_endpoint, + ca_file=None, + ca_data=ca_data, + ) + + def test_use_token_proxy_minimal_config(self): + """Test use_token_proxy with minimal configuration (only proxy endpoint).""" + tenant_id = "tenant-id" + client_id = "client-id" + token_file_path = "foo-path" + proxy_endpoint = "https://proxy.example.com:8080" + + env_vars = { + "AZURE_KUBERNETES_TOKEN_PROXY": proxy_endpoint, + } + + with patch.dict(os.environ, env_vars, clear=False): + with patch("azure.identity.aio._credentials.workload_identity._get_transport") as mock_get_transport: + mock_transport_instance = MagicMock() + mock_get_transport.return_value = mock_transport_instance + + WorkloadIdentityCredential( + tenant_id=tenant_id, + client_id=client_id, + token_file_path=token_file_path, + use_token_proxy=True, + ) + + mock_get_transport.assert_called_once_with( + sni=None, + token_proxy_endpoint=proxy_endpoint, + ca_file=None, + ca_data=None, + ) + + def test_use_token_proxy_missing_proxy_endpoint(self): + """Test that use_token_proxy=True without proxy endpoint uses the normal transport.""" + tenant_id = "tenant-id" + client_id = "client-id" + token_file_path = "foo-path" + + # Ensure proxy endpoint env var is not set + with patch.dict(os.environ, {}, clear=False): + if "AZURE_KUBERNETES_TOKEN_PROXY" in os.environ: + del os.environ["AZURE_KUBERNETES_TOKEN_PROXY"] + + with patch("azure.identity.aio._credentials.workload_identity._get_transport") as mock_get_transport: + WorkloadIdentityCredential( + tenant_id=tenant_id, + client_id=client_id, + token_file_path=token_file_path, + use_token_proxy=True, + ) + mock_get_transport.assert_not_called() + + def test_use_token_proxy_both_ca_file_and_data_raises_error(self): + """Test that setting both CA file and CA data raises ValueError.""" + tenant_id = "tenant-id" + client_id = "client-id" + token_file_path = "foo-path" + proxy_endpoint = "https://proxy.example.com:8080" + ca_file_path = "/path/to/ca.pem" + ca_data = "-----BEGIN CERTIFICATE-----\nTest CA data\n-----END CERTIFICATE-----" + + env_vars = { + "AZURE_KUBERNETES_TOKEN_PROXY": proxy_endpoint, + "AZURE_KUBERNETES_CA_FILE": ca_file_path, + "AZURE_KUBERNETES_CA_DATA": ca_data, + } + + with patch.dict(os.environ, env_vars, clear=False): + with pytest.raises(ValueError, match="Both AZURE_KUBERNETES_CA_FILE and AZURE_KUBERNETES_CA_DATA are set"): + WorkloadIdentityCredential( + tenant_id=tenant_id, + client_id=client_id, + token_file_path=token_file_path, + use_token_proxy=True, + ) + + def test_use_token_proxy_missing_endpoint_with_custom_env_vars_raises_error(self): + """Test that use_token_proxy=True without proxy endpoint but with other custom env vars raises ValueError.""" + tenant_id = "tenant-id" + client_id = "client-id" + token_file_path = "foo-path" + sni_hostname = "sni.example.com" + ca_file_path = "/path/to/ca.pem" + ca_data = "-----BEGIN CERTIFICATE-----\nTest CA data\n-----END CERTIFICATE-----" + + # Ensure proxy endpoint is not set + if "AZURE_KUBERNETES_TOKEN_PROXY" in os.environ: + del os.environ["AZURE_KUBERNETES_TOKEN_PROXY"] + + # Test with SNI set but no proxy endpoint + env_vars_sni = { + "AZURE_KUBERNETES_SNI_NAME": sni_hostname, + } + with patch.dict(os.environ, env_vars_sni, clear=False): + with pytest.raises(ValueError): + WorkloadIdentityCredential( + tenant_id=tenant_id, + client_id=client_id, + token_file_path=token_file_path, + use_token_proxy=True, + ) + + # Test with CA file set but no proxy endpoint + env_vars_ca_file = { + "AZURE_KUBERNETES_CA_FILE": ca_file_path, + } + with patch.dict(os.environ, env_vars_ca_file, clear=False): + with pytest.raises(ValueError): + WorkloadIdentityCredential( + tenant_id=tenant_id, + client_id=client_id, + token_file_path=token_file_path, + use_token_proxy=True, + ) + + # Test with CA data set but no proxy endpoint + env_vars_ca_data = { + "AZURE_KUBERNETES_CA_DATA": ca_data, + } + with patch.dict(os.environ, env_vars_ca_data, clear=False): + with pytest.raises(ValueError): + WorkloadIdentityCredential( + tenant_id=tenant_id, + client_id=client_id, + token_file_path=token_file_path, + use_token_proxy=True, + ) + + @pytest.mark.asyncio + @pytest.mark.parametrize("get_token_method", GET_TOKEN_METHODS) + async def test_use_token_proxy_get_token_success(self, get_token_method): + """Test successful token acquisition when using token proxy.""" + tenant_id = "tenant-id" + client_id = "client-id" + access_token = "foo-access-token" + token_file_path = "foo-path" + assertion = "foo-assertion" + proxy_endpoint = "https://proxy.example.com:8080" + + async def send(request, **kwargs): + assert "claims" not in kwargs + assert "tenant_id" not in kwargs + assert request.data.get("client_assertion") == assertion + return mock_response(json_payload=build_aad_response(access_token=access_token)) + + mock_transport_instance = MagicMock(send=send) + + env_vars = { + "AZURE_KUBERNETES_TOKEN_PROXY": proxy_endpoint, + } + + with patch.dict(os.environ, env_vars, clear=False): + with patch("azure.identity.aio._credentials.workload_identity._get_transport") as mock_get_transport: + mock_get_transport.return_value = mock_transport_instance + + credential = WorkloadIdentityCredential( + tenant_id=tenant_id, + client_id=client_id, + token_file_path=token_file_path, + use_token_proxy=True, + ) + + open_mock = mock_open(read_data=assertion) + with patch("builtins.open", open_mock): + token = await getattr(credential, get_token_method)("scope") + assert token.token == access_token + + open_mock.assert_called_once_with(token_file_path, encoding="utf-8") + + def test_use_token_proxy_false_does_not_create_transport(self): + """Test that use_token_proxy=False (default) does not create a custom transport.""" + tenant_id = "tenant-id" + client_id = "client-id" + token_file_path = "foo-path" + + with patch("azure.identity.aio._credentials.workload_identity._get_transport") as mock_get_transport: + WorkloadIdentityCredential( + tenant_id=tenant_id, + client_id=client_id, + token_file_path=token_file_path, + use_token_proxy=False, + ) + mock_get_transport.assert_not_called() + + +class TestCustomAioHttpTransport: + """Test cases for the custom AioHttpTransport used by WorkloadIdentityCredential.""" + + def test_get_transport_creates_workload_identity_aiohttp_transport(self, ca_data): + """Test that _get_transport creates WorkloadIdentityAioHttpTransport with correct parameters.""" + sni = "test.sni.com" + proxy_endpoint = "https://proxy.example.com:8080" + ca_file = PEM_CERT_PATH + + transport = _get_transport( + sni=sni, + token_proxy_endpoint=proxy_endpoint, + ca_file=ca_file, + ca_data=None, + ) + + assert type(transport) is CustomAioHttpTransport + assert hasattr(transport, "_sni") + assert hasattr(transport, "_proxy_endpoint") + assert hasattr(transport, "_ca_file") + assert hasattr(transport, "_ca_data") + assert transport._sni == sni + assert transport._proxy_endpoint == proxy_endpoint + assert transport._ca_file == ca_file + assert transport._ca_data == ca_data + + assert hasattr(transport, "_ssl_context") + assert transport._ssl_context is not None + + def test_get_transport_with_minimal_config(self): + """Test _get_transport with minimal configuration.""" + proxy_endpoint = "https://proxy.example.com:8080" + + transport = _get_transport( + sni=None, + token_proxy_endpoint=proxy_endpoint, + ca_file=None, + ca_data=None, + ) + + assert transport is not None + assert transport._sni is None + assert transport._proxy_endpoint == proxy_endpoint + assert transport._ca_file is None + assert transport._ca_data is None + + assert hasattr(transport, "_ssl_context") + assert transport._ssl_context is not None + + def test_workload_identity_aiohttp_transport_inherits_from_token_binding_mixin(self): + """Test that WorkloadIdentityAioHttpTransport inherits from TokenBindingTransportMixin.""" + transport = _get_transport( + sni="test.sni.com", + token_proxy_endpoint="https://proxy.example.com:8080", + ca_file=None, + ca_data=None, + ) + + assert transport is not None + + # Verify inheritance from TokenBindingTransportMixin + from azure.identity._internal.token_binding_transport_mixin import TokenBindingTransportMixin + + assert isinstance(transport, TokenBindingTransportMixin) + + # Verify TokenBindingTransportMixin methods are available + assert hasattr(transport, "_update_request_url") + assert hasattr(transport, "_has_ca_file_changed") + assert hasattr(transport, "_load_ca_file_to_data") + assert hasattr(transport, "_validate_url") + + +class TestCustomAioHttpTransportWithLocalServer: + """Integration tests using a local test server for WorkloadIdentityAioHttpTransport.""" + + @pytest.mark.asyncio + async def test_basic_https_request(self): + """Test basic HTTPS request to test server.""" + with TokenProxyTestServer(use_ssl=True) as server: + # Create transport with server's CA certificate + transport = _get_transport(sni=None, token_proxy_endpoint=None, ca_file=server.ca_file, ca_data=None) + assert transport is not None + request = HttpRequest("GET", f"{server.base_url}/health") + + response = await transport.send(request) + + assert response.status_code == 200 + data = response.json() + assert data["status"] == "healthy" + assert "timestamp" in data + + @pytest.mark.asyncio + async def test_post_request_with_body(self): + """Test POST request with request body.""" + with TokenProxyTestServer(use_ssl=True) as server: + transport = _get_transport(sni=None, token_proxy_endpoint=None, ca_file=server.ca_file, ca_data=None) + assert transport is not None + + # Prepare OAuth-like request + body = "grant_type=client_credentials&scope=https://graph.microsoft.com/.default" + request = HttpRequest( + "POST", + f"{server.base_url}/tenant/oauth2/v2.0/token", + headers={"Content-Type": "application/x-www-form-urlencoded"}, + data=body.encode("utf-8"), + ) + + response = await transport.send(request) + + assert response.status_code == 200 + data = response.json() + assert "access_token" in data + assert data["token_type"] == "Bearer" + assert data["expires_in"] == 3600 + + @pytest.mark.asyncio + async def test_proxy_endpoint_comprehensive(self): + """Test comprehensive proxy endpoint functionality with various HTTP methods and scenarios.""" + with TokenProxyTestServer(use_ssl=True) as server: + # Configure transport with proxy endpoint + transport = _get_transport( + sni=None, token_proxy_endpoint=server.base_url, ca_file=server.ca_file, ca_data=None + ) + assert transport is not None + + # Test 1: POST request with JSON body through proxy + post_data = {"grant_type": "client_credentials", "scope": "https://graph.microsoft.com/.default"} + post_request = HttpRequest( + "POST", + "https://login.microsoftonline.com/tenant/oauth2/v2.0/token2", + headers={"Content-Type": "application/json"}, + json=post_data, + ) + + post_response = await transport.send(post_request) + assert post_response.status_code == 200 + post_data_response = post_response.json() + assert post_data_response["method"] == "POST" + assert post_data_response["proxied_path"] == "/tenant/oauth2/v2.0/token2" + + # Test 2: PUT request through proxy + put_request = HttpRequest( + "PUT", + "https://graph.microsoft.com/v1.0/me/profile", + headers={"Content-Type": "application/json"}, + json={"displayName": "Test User"}, + ) + + put_response = await transport.send(put_request) + assert put_response.status_code == 200 + put_data_response = put_response.json() + assert put_data_response["method"] == "PUT" + assert put_data_response["proxied_path"] == "/v1.0/me/profile" + + # Test 3: DELETE request through proxy + delete_request = HttpRequest("DELETE", "https://graph.microsoft.com/v1.0/applications/app-id") + + delete_response = await transport.send(delete_request) + assert delete_response.status_code == 200 + delete_data_response = delete_response.json() + assert delete_data_response["method"] == "DELETE" + assert delete_data_response["proxied_path"] == "/v1.0/applications/app-id" + + # Test 4: Complex URL with multiple path segments and query parameters + complex_url = "https://management.azure.com/subscriptions/sub-id/resourceGroups/rg/providers/Microsoft.Storage/storageAccounts/account?api-version=2021-04-01&expand=properties" + complex_request = HttpRequest("GET", complex_url) + + complex_response = await transport.send(complex_request) + assert complex_response.status_code == 200 + complex_data_response = complex_response.json() + assert complex_data_response["method"] == "GET" + expected_path = "/subscriptions/sub-id/resourceGroups/rg/providers/Microsoft.Storage/storageAccounts/account?api-version=2021-04-01&expand=properties" + assert complex_data_response["proxied_path"] == expected_path + + # Test 5: Request with custom headers through proxy + headers_request = HttpRequest( + "GET", + "https://vault.azure.net/secrets/test-secret?api-version=7.3", + headers={ + "Authorization": "Bearer test-token", + "X-Custom-Header": "proxy-test-value", + "User-Agent": "Azure-SDK-For-Python", + }, + ) + + headers_response = await transport.send(headers_request) + assert headers_response.status_code == 200 + headers_data_response = headers_response.json() + assert headers_data_response["method"] == "GET" + assert headers_data_response["proxied_path"] == "/secrets/test-secret?api-version=7.3" + + # Verify headers were forwarded through proxy + received_headers = headers_data_response["headers_received"] + assert "Authorization" in received_headers + assert "X-Custom-Header" in received_headers + assert received_headers["Authorization"] == "Bearer test-token" + assert received_headers["X-Custom-Header"] == "proxy-test-value" + + @pytest.mark.asyncio + async def test_sni_with_custom_hostname(self): + """Test SNI (Server Name Indication) with custom hostname.""" + with TokenProxyTestServer(use_ssl=True) as server: + # Use SNI with a different hostname than the server + transport = _get_transport( + sni="1234.ests.aks", token_proxy_endpoint=None, ca_file=server.ca_file, ca_data=None + ) + assert transport is not None + + request = HttpRequest("GET", f"{server.base_url}/health") + + response = await transport.send(request) + + assert response.status_code == 200 + data = response.json() + assert data["status"] == "healthy" + + transport = _get_transport( + sni="unmatched.sni.hostname", token_proxy_endpoint=None, ca_file=server.ca_file, ca_data=None + ) + assert transport is not None + request = HttpRequest("GET", f"{server.base_url}/health") + with pytest.raises(ServiceRequestError): + await transport.send(request) + + @pytest.mark.asyncio + async def test_ca_file_change_detection(self): + """Test CA file change detection with real certificates.""" + with TokenProxyTestServer(use_ssl=True) as server: + # Create a copy of the CA file that we can modify + ca_file = server.ca_file + if ca_file is None: + pytest.skip("CA file not available") + + assert ca_file is not None # Type hint for mypy + with tempfile.NamedTemporaryFile(mode="w", delete=False, suffix=".pem") as temp_ca: + with open(ca_file, "r") as original: + original_content = original.read() + temp_ca.write(original_content) + temp_ca_path = temp_ca.name + + try: + transport = _get_transport(sni=None, token_proxy_endpoint=None, ca_file=temp_ca_path, ca_data=None) + assert transport is not None + + # First request should work + request = HttpRequest("GET", f"{server.base_url}/health") + response1 = await transport.send(request) + assert response1.status_code == 200 + + # Modify the CA file (add some content) + real_sleep(0.1) + with open(temp_ca_path, "a") as f: + f.write("\n# Modified for testing\n") + + # Second request should still work (using the same cert content) + response2 = await transport.send(request) + assert transport._ca_data != original_content + assert response2.status_code == 200 + + finally: + os.unlink(temp_ca_path) + + @pytest.mark.asyncio + async def test_ssl_error_handling(self): + """Test SSL error handling.""" + with TokenProxyTestServer(use_ssl=True) as server: + # Create transport without proper CA file (will cause SSL error) + transport = _get_transport(sni=None, token_proxy_endpoint=None, ca_file=None, ca_data=None) + assert transport is not None + + request = HttpRequest("GET", f"{server.base_url}/health") + + # Should raise SSL-related error + with pytest.raises((ServiceRequestError, ServiceResponseError)): + await transport.send(request) + + @pytest.mark.asyncio + async def test_server_error_response(self): + """Test handling of server error responses.""" + with TokenProxyTestServer(use_ssl=True) as server: + transport = _get_transport(sni=None, token_proxy_endpoint=None, ca_file=server.ca_file, ca_data=None) + assert transport is not None + + request = HttpRequest("GET", f"{server.base_url}/error/500") + + response = await transport.send(request) + + assert response.status_code == 500 + # Should not raise exception, just return error response + + @pytest.mark.asyncio + async def test_custom_headers_preserved(self): + """Test that custom headers are preserved and sent to server.""" + with TokenProxyTestServer(use_ssl=True) as server: + transport = _get_transport(sni=None, token_proxy_endpoint=None, ca_file=server.ca_file, ca_data=None) + assert transport is not None + + custom_headers = { + "Authorization": "Bearer test-token", + "User-Agent": "WorkloadIdentityAioHttpTransport/1.0", + "X-Custom-Header": "test-value", + } + + request = HttpRequest("GET", f"{server.base_url}/proxy/test", headers=custom_headers) + + response = await transport.send(request) + + assert response.status_code == 200 + data = response.json() + + # Server echoes back the headers it received + received_headers = data["headers_received"] + assert "Authorization" in received_headers + assert "User-Agent" in received_headers + assert "X-Custom-Header" in received_headers + assert received_headers["Authorization"] == "Bearer test-token" + + @pytest.mark.asyncio + async def test_query_parameters_preserved(self): + """Test that query parameters are preserved in proxy requests.""" + with TokenProxyTestServer(use_ssl=True) as server: + transport = _get_transport( + sni=None, token_proxy_endpoint=server.base_url, ca_file=server.ca_file, ca_data=None + ) + assert transport is not None + + # Request with query parameters + original_url = "https://example.com/api/data?scope=read&limit=10&format=json" + request = HttpRequest("GET", original_url) + + response = await transport.send(request) + + assert response.status_code == 200 + data = response.json() + + # The path should include the query parameters + expected_path = "/api/data?scope=read&limit=10&format=json" + assert data["proxied_path"] == expected_path + + @pytest.mark.asyncio + async def test_concurrent_requests(self): + """Test handling multiple concurrent requests.""" + with TokenProxyTestServer(use_ssl=True) as server: + transport = _get_transport(sni=None, token_proxy_endpoint=None, ca_file=server.ca_file, ca_data=None) + assert transport is not None + + async def make_request(request_id): + request = HttpRequest("GET", f"{server.base_url}/health") + response = await transport.send(request) + return request_id, response.status_code, response.json() + + # Make 5 concurrent requests + tasks = [make_request(i) for i in range(5)] + results = await asyncio.gather(*tasks) + + # All requests should succeed + assert len(results) == 5 + for request_id, status_code, data in results: + assert status_code == 200 + assert data["status"] == "healthy" + + +class TestCustomAsyncioRequestsTransportFallback: + """Test cases for the custom AsyncioRequestsTransport used by WorkloadIdentityCredential.""" + + def test_get_transport_creates_workload_identity_asyncio_requests_transport(self, ca_data): + """Test that _get_transport creates WorkloadIdentityAsyncioRequestsTransport with correct parameters.""" + sni = "test.sni.com" + proxy_endpoint = "https://proxy.example.com:8080" + ca_file = PEM_CERT_PATH + + with patch.dict("sys.modules", {"azure.identity.aio._internal.token_binding_transport_aiohttp": None}): + transport = _get_transport( + sni=sni, + token_proxy_endpoint=proxy_endpoint, + ca_file=ca_file, + ca_data=None, + ) + + assert type(transport) is CustomAsyncioRequestsTransport + assert hasattr(transport, "_sni") + assert hasattr(transport, "_proxy_endpoint") + assert hasattr(transport, "_ca_file") + assert hasattr(transport, "_ca_data") + assert transport._sni == sni + assert transport._proxy_endpoint == proxy_endpoint + assert transport._ca_file == ca_file + assert transport._ca_data == ca_data + + @pytest.mark.asyncio + async def test_basic_https_request(self): + """Test basic HTTPS request to test server.""" + with TokenProxyTestServer(use_ssl=True) as server: + # Create transport with server's CA certificate + with patch.dict("sys.modules", {"azure.identity.aio._internal.token_binding_transport_aiohttp": None}): + transport = _get_transport(sni=None, token_proxy_endpoint=None, ca_file=server.ca_file, ca_data=None) + assert type(transport) is CustomAsyncioRequestsTransport + request = HttpRequest("GET", f"{server.base_url}/health") + + response = await transport.send(request) + + assert response.status_code == 200 + data = response.json() + assert data["status"] == "healthy" + assert "timestamp" in data + + @pytest.mark.asyncio + async def test_proxy_endpoint_comprehensive(self): + """Test comprehensive proxy endpoint functionality with various HTTP methods and scenarios.""" + with TokenProxyTestServer(use_ssl=True) as server: + with patch.dict("sys.modules", {"azure.identity.aio._internal.token_binding_transport_aiohttp": None}): + transport = _get_transport( + sni=None, token_proxy_endpoint=server.base_url, ca_file=server.ca_file, ca_data=None + ) + assert type(transport) is CustomAsyncioRequestsTransport + + # Test 1: POST request with JSON body through proxy + post_data = {"grant_type": "client_credentials", "scope": "https://graph.microsoft.com/.default"} + post_request = HttpRequest( + "POST", + "https://login.microsoftonline.com/tenant/oauth2/v2.0/token2", + headers={"Content-Type": "application/json"}, + json=post_data, + ) + + post_response = await transport.send(post_request) + assert post_response.status_code == 200 + post_data_response = post_response.json() + assert post_data_response["method"] == "POST" + assert post_data_response["proxied_path"] == "/tenant/oauth2/v2.0/token2" + + # Test 2: PUT request through proxy + put_request = HttpRequest( + "PUT", + "https://graph.microsoft.com/v1.0/me/profile", + headers={"Content-Type": "application/json"}, + json={"displayName": "Test User"}, + ) + + put_response = await transport.send(put_request) + assert put_response.status_code == 200 + put_data_response = put_response.json() + assert put_data_response["method"] == "PUT" + assert put_data_response["proxied_path"] == "/v1.0/me/profile" + + # Test 3: DELETE request through proxy + delete_request = HttpRequest("DELETE", "https://graph.microsoft.com/v1.0/applications/app-id") + + delete_response = await transport.send(delete_request) + assert delete_response.status_code == 200 + delete_data_response = delete_response.json() + assert delete_data_response["method"] == "DELETE" + assert delete_data_response["proxied_path"] == "/v1.0/applications/app-id" + + # Test 4: Complex URL with multiple path segments and query parameters + complex_url = "https://management.azure.com/subscriptions/sub-id/resourceGroups/rg/providers/Microsoft.Storage/storageAccounts/account?api-version=2021-04-01&expand=properties" + complex_request = HttpRequest("GET", complex_url) + + complex_response = await transport.send(complex_request) + assert complex_response.status_code == 200 + complex_data_response = complex_response.json() + assert complex_data_response["method"] == "GET" + expected_path = "/subscriptions/sub-id/resourceGroups/rg/providers/Microsoft.Storage/storageAccounts/account?api-version=2021-04-01&expand=properties" + assert complex_data_response["proxied_path"] == expected_path + + # Test 5: Request with custom headers through proxy + headers_request = HttpRequest( + "GET", + "https://vault.azure.net/secrets/test-secret?api-version=7.3", + headers={ + "Authorization": "Bearer test-token", + "X-Custom-Header": "proxy-test-value", + "User-Agent": "Azure-SDK-For-Python", + }, + ) + + headers_response = await transport.send(headers_request) + assert headers_response.status_code == 200 + headers_data_response = headers_response.json() + assert headers_data_response["method"] == "GET" + assert headers_data_response["proxied_path"] == "/secrets/test-secret?api-version=7.3" + + # Verify headers were forwarded through proxy + received_headers = headers_data_response["headers_received"] + assert "Authorization" in received_headers + assert "X-Custom-Header" in received_headers + assert received_headers["Authorization"] == "Bearer test-token" + assert received_headers["X-Custom-Header"] == "proxy-test-value" + + @pytest.mark.asyncio + async def test_sni_with_custom_hostname(self): + """Test SNI (Server Name Indication) with custom hostname.""" + with TokenProxyTestServer(use_ssl=True) as server: + with patch.dict("sys.modules", {"azure.identity.aio._internal.token_binding_transport_aiohttp": None}): + # Use SNI with a different hostname than the server + transport = _get_transport( + sni="1234.ests.aks", token_proxy_endpoint=None, ca_file=server.ca_file, ca_data=None + ) + assert type(transport) is CustomAsyncioRequestsTransport + + request = HttpRequest("GET", f"{server.base_url}/health") + response = await transport.send(request) + + assert response.status_code == 200 + data = response.json() + assert data["status"] == "healthy" + + # Check an invalid SNI hostname + transport = _get_transport( + sni="unmatched.sni.hostname", token_proxy_endpoint=None, ca_file=server.ca_file, ca_data=None + ) + assert transport is not None + request = HttpRequest("GET", f"{server.base_url}/health") + with pytest.raises(ServiceRequestError): + await transport.send(request) + + @pytest.mark.asyncio + async def test_ca_file_change_detection(self): + """Test CA file change detection with real certificates.""" + with TokenProxyTestServer(use_ssl=True) as server: + with patch.dict("sys.modules", {"azure.identity.aio._internal.token_binding_transport_aiohttp": None}): + # Create a copy of the CA file that we can modify + ca_file = server.ca_file + if ca_file is None: + pytest.skip("CA file not available") + + assert ca_file is not None # Type hint for mypy + with tempfile.NamedTemporaryFile(mode="w", delete=False, suffix=".pem") as temp_ca: + with open(ca_file, "r") as original: + original_content = original.read() + temp_ca.write(original_content) + temp_ca_path = temp_ca.name + + try: + transport = _get_transport(sni=None, token_proxy_endpoint=None, ca_file=temp_ca_path, ca_data=None) + assert type(transport) is CustomAsyncioRequestsTransport + + # First request should work + request = HttpRequest("GET", f"{server.base_url}/health") + response1 = await transport.send(request) + assert response1.status_code == 200 + + # Modify the CA file (add some content) + real_sleep(0.1) + with open(temp_ca_path, "a") as f: + f.write("\n# Modified for testing\n") + + # Second request should still work (using the same cert content) + response2 = await transport.send(request) + assert transport._ca_data != original_content + assert response2.status_code == 200 + + finally: + os.unlink(temp_ca_path)