diff --git a/sdk/identity/azure-identity/CHANGELOG.md b/sdk/identity/azure-identity/CHANGELOG.md index 58a3acf806e1..756671552282 100644 --- a/sdk/identity/azure-identity/CHANGELOG.md +++ b/sdk/identity/azure-identity/CHANGELOG.md @@ -5,6 +5,7 @@ ### Features Added - Changed parameter from `instance_discovery` to `disable_instance_discovery` to make it more explicit. +- Service principal credentials now enable support for [Continuous Access Evaluation (CAE)](https://learn.microsoft.com/azure/active-directory/conditional-access/concept-continuous-access-evaluation-workload). This indicates to Azure Active Directory that your application can handle CAE claims challenges. ### Breaking Changes diff --git a/sdk/identity/azure-identity/azure/identity/_constants.py b/sdk/identity/azure-identity/azure/identity/_constants.py index 846a26757ae7..545786cb4211 100644 --- a/sdk/identity/azure-identity/azure/identity/_constants.py +++ b/sdk/identity/azure-identity/azure/identity/_constants.py @@ -50,3 +50,5 @@ class EnvironmentVariables: AZURE_FEDERATED_TOKEN_FILE = "AZURE_FEDERATED_TOKEN_FILE" WORKLOAD_IDENTITY_VARS = (AZURE_AUTHORITY_HOST, AZURE_TENANT_ID, AZURE_FEDERATED_TOKEN_FILE) + + AZURE_IDENTITY_DISABLE_CP1 = "AZURE_IDENTITY_DISABLE_CP1" diff --git a/sdk/identity/azure-identity/azure/identity/_credentials/silent.py b/sdk/identity/azure-identity/azure/identity/_credentials/silent.py index a61040b40c57..a0a31b9b1f6c 100644 --- a/sdk/identity/azure-identity/azure/identity/_credentials/silent.py +++ b/sdk/identity/azure-identity/azure/identity/_credentials/silent.py @@ -18,6 +18,7 @@ from .._internal.msal_client import MsalClient from .._internal.shared_token_cache import NO_TOKEN from .._persistent_cache import _load_persistent_cache, TokenCachePersistenceOptions +from .._constants import EnvironmentVariables from .. import AuthenticationRecord @@ -85,7 +86,7 @@ def _get_client_application(self, **kwargs: Any): ) if tenant_id not in self._client_applications: # CP1 = can handle claims challenges (CAE) - capabilities = None if "AZURE_IDENTITY_DISABLE_CP1" in os.environ else ["CP1"] + capabilities = None if EnvironmentVariables.AZURE_IDENTITY_DISABLE_CP1 in os.environ else ["CP1"] self._client_applications[tenant_id] = PublicClientApplication( client_id=self._auth_record.client_id, authority="https://{}/{}".format(self._auth_record.authority, tenant_id), diff --git a/sdk/identity/azure-identity/azure/identity/_internal/aad_client_base.py b/sdk/identity/azure-identity/azure/identity/_internal/aad_client_base.py index f1eb96e605fa..a394b399e5dc 100644 --- a/sdk/identity/azure-identity/azure/identity/_internal/aad_client_base.py +++ b/sdk/identity/azure-identity/azure/identity/_internal/aad_client_base.py @@ -5,6 +5,7 @@ import abc import base64 import json +import os import time from uuid import uuid4 from typing import TYPE_CHECKING, List, Any, Iterable, Optional, Union, Dict @@ -17,9 +18,11 @@ from azure.core.pipeline.transport import HttpRequest from azure.core.credentials import AccessToken from azure.core.exceptions import ClientAuthenticationError +from .._constants import EnvironmentVariables from .utils import get_default_authority, normalize_authority, resolve_tenant from .aadclient_certificate import AadClientCertificate + if TYPE_CHECKING: from azure.core.pipeline import AsyncPipeline, Pipeline from azure.core.pipeline.policies import AsyncHTTPPolicy, HTTPPolicy, SansIOHTTPPolicy @@ -51,6 +54,8 @@ def __init__( self._cache = cache or TokenCache() self._client_id = client_id + # CP1 = can handle claims challenges (CAE) + self._capabilities = None if EnvironmentVariables.AZURE_IDENTITY_DISABLE_CP1 in os.environ else ["CP1"] self._additionally_allowed_tenants = additionally_allowed_tenants or [] self._pipeline = self._build_pipeline(**kwargs) @@ -171,6 +176,10 @@ def _get_auth_code_request( "redirect_uri": redirect_uri, "scope": " ".join(scopes), } + + claims = _merge_claims_challenge_and_capabilities(self._capabilities, kwargs.get("claims")) + if claims: + data["claims"] = claims if client_secret: data["client_secret"] = client_secret @@ -191,6 +200,10 @@ def _get_jwt_assertion_request( "scope": " ".join(scopes), } + claims = _merge_claims_challenge_and_capabilities(self._capabilities, kwargs.get("claims")) + if claims: + data["claims"] = claims + request = self._post(data, **kwargs) return request @@ -233,6 +246,11 @@ def _get_client_secret_request(self, scopes: Iterable[str], secret: str, **kwarg "grant_type": "client_credentials", "scope": " ".join(scopes), } + + claims = _merge_claims_challenge_and_capabilities(self._capabilities, kwargs.get("claims")) + if claims: + data["claims"] = claims + request = self._post(data, **kwargs) return request @@ -250,6 +268,11 @@ def _get_on_behalf_of_request( "requested_token_use": "on_behalf_of", "scope": " ".join(scopes), } + + claims = _merge_claims_challenge_and_capabilities(self._capabilities, kwargs.get("claims")) + if claims: + data["claims"] = claims + if isinstance(client_credential, AadClientCertificate): data["client_assertion"] = self._get_client_certificate_assertion(client_credential) data["client_assertion_type"] = JWT_BEARER_ASSERTION @@ -272,6 +295,11 @@ def _get_refresh_token_request( "client_id": self._client_id, "client_info": 1, # request AAD include home_account_id in its response } + + claims = _merge_claims_challenge_and_capabilities(self._capabilities, kwargs.get("claims")) + if claims: + data["claims"] = claims + request = self._post(data, **kwargs) return request @@ -289,6 +317,10 @@ def _get_refresh_token_on_behalf_of_request( "client_id": self._client_id, "client_info": 1, # request AAD include home_account_id in its response } + claims = _merge_claims_challenge_and_capabilities(self._capabilities, kwargs.get("claims")) + if claims: + data["claims"] = claims + if isinstance(client_credential, AadClientCertificate): data["client_assertion"] = self._get_client_certificate_assertion(client_credential) data["client_assertion_type"] = JWT_BEARER_ASSERTION @@ -310,6 +342,17 @@ def _post(self, data: Dict, **kwargs: Any) -> HttpRequest: return HttpRequest("POST", url, data=data, headers={"Content-Type": "application/x-www-form-urlencoded"}) +def _merge_claims_challenge_and_capabilities(capabilities, claims_challenge): + # Represent capabilities as {"access_token": {"xms_cc": {"values": capabilities}}} + # and then merge/add it into incoming claims + if not capabilities: + return claims_challenge + claims_dict = json.loads(claims_challenge) if claims_challenge else {} + for key in ["access_token"]: + claims_dict.setdefault(key, {}).update(xms_cc={"values": capabilities}) + return json.dumps(claims_dict) + + def _scrub_secrets(response: Dict) -> None: for secret in ("access_token", "refresh_token"): if secret in response: diff --git a/sdk/identity/azure-identity/azure/identity/_internal/client_credential_base.py b/sdk/identity/azure-identity/azure/identity/_internal/client_credential_base.py index 4987a8ba1d81..401b1e50388f 100644 --- a/sdk/identity/azure-identity/azure/identity/_internal/client_credential_base.py +++ b/sdk/identity/azure-identity/azure/identity/_internal/client_credential_base.py @@ -20,7 +20,8 @@ class ClientCredentialBase(MsalCredential, GetTokenMixin): def _acquire_token_silently(self, *scopes: str, **kwargs: Any) -> Optional[AccessToken]: app = self._get_app(**kwargs) request_time = int(time.time()) - result = app.acquire_token_silent_with_error(list(scopes), account=None, **kwargs) + result = app.acquire_token_silent_with_error( + list(scopes), account=None, claims_challenge=kwargs.pop("claims", None), **kwargs) if result and "access_token" in result and "expires_in" in result: return AccessToken(result["access_token"], request_time + int(result["expires_in"])) return None @@ -29,7 +30,7 @@ def _acquire_token_silently(self, *scopes: str, **kwargs: Any) -> Optional[Acces def _request_token(self, *scopes: str, **kwargs: Any) -> Optional[AccessToken]: app = self._get_app(**kwargs) request_time = int(time.time()) - result = app.acquire_token_for_client(list(scopes)) + result = app.acquire_token_for_client(list(scopes), claims_challenge=kwargs.pop("claims", None)) if "access_token" not in result: message = "Authentication failed: {}".format(result.get("error_description") or result.get("error")) raise ClientAuthenticationError(message=message) diff --git a/sdk/identity/azure-identity/azure/identity/_internal/msal_credentials.py b/sdk/identity/azure-identity/azure/identity/_internal/msal_credentials.py index 7fb3e9a8aee7..b691d06f71aa 100644 --- a/sdk/identity/azure-identity/azure/identity/_internal/msal_credentials.py +++ b/sdk/identity/azure-identity/azure/identity/_internal/msal_credentials.py @@ -71,7 +71,7 @@ def _get_app(self, **kwargs): ) if tenant_id not in self._client_applications: # CP1 = can handle claims challenges (CAE) - capabilities = None if "AZURE_IDENTITY_DISABLE_CP1" in os.environ else ["CP1"] + capabilities = None if EnvironmentVariables.AZURE_IDENTITY_DISABLE_CP1 in os.environ else ["CP1"] cls = msal.ConfidentialClientApplication if self._client_credential else msal.PublicClientApplication self._client_applications[tenant_id] = cls( client_id=self._client_id, diff --git a/sdk/identity/azure-identity/tests/helpers.py b/sdk/identity/azure-identity/tests/helpers.py index cc9fd82f21f6..e7ce29403ca1 100644 --- a/sdk/identity/azure-identity/tests/helpers.py +++ b/sdk/identity/azure-identity/tests/helpers.py @@ -199,3 +199,9 @@ def urlsafeb64_decode(s): padding_needed = 4 - len(s) % 4 return base64.urlsafe_b64decode(s + b"=" * padding_needed) + + +def get_token_payload_contents(token: str): + _, payload, _ = token.split(".") + decoded_payload = urlsafeb64_decode(payload).decode() + return json.loads(decoded_payload) diff --git a/sdk/identity/azure-identity/tests/test_aad_client.py b/sdk/identity/azure-identity/tests/test_aad_client.py index 700d3cae28cc..042643bc3a8e 100644 --- a/sdk/identity/azure-identity/tests/test_aad_client.py +++ b/sdk/identity/azure-identity/tests/test_aad_client.py @@ -21,6 +21,16 @@ from mock import Mock, patch # type: ignore +BASE_CLASS_METHODS = [ + ("_get_auth_code_request", ("code", "redirect_uri")), + ("_get_client_secret_request", ("secret", )), + ("_get_jwt_assertion_request", ("assertion", )), + ("_get_refresh_token_request", ("refresh_token", )), + ("_get_on_behalf_of_request", ("client_credential", "user_assertion")), + ("_get_refresh_token_on_behalf_of_request", ("client_credential", "refresh_token")) +] + + def test_error_reporting(): error_name = "everything's sideways" error_description = "something went wrong" @@ -316,3 +326,41 @@ def test_multitenant_cache(): assert client_d.get_cached_access_token([scope]) is None with pytest.raises(ClientAuthenticationError, match=message): client_d.get_cached_access_token([scope], tenant_id=tenant_a) + + +@pytest.mark.parametrize("method,args", BASE_CLASS_METHODS) +def test_claims(method, args): + + scopes = ["scope"] + claims = '{"access_token": {"essential": "true"}}' + + client = AadClient("tenant_id", "client_id") + + expected_merged_claims = '{"access_token": {"essential": "true", "xms_cc": {"values": ["CP1"]}}}' + + with patch.object(AadClient, "_post") as post_mock: + func = getattr(client, method) + func(scopes, *args, claims=claims) + + assert post_mock.call_count == 1 + data, _ = post_mock.call_args + assert len(data) == 1 + assert data[0]["claims"] == expected_merged_claims + + +@pytest.mark.parametrize("method,args", BASE_CLASS_METHODS) +def test_claims_disable_capabilities(method, args): + scopes = ["scope"] + claims = '{"access_token": {"essential": "true"}}' + + with patch.dict("os.environ", {"AZURE_IDENTITY_DISABLE_CP1": "true"}): + client = AadClient("tenant_id", "client_id") + + with patch.object(AadClient, "_post") as post_mock: + func = getattr(client, method) + func(scopes, *args, claims=claims) + + assert post_mock.call_count == 1 + data, _ = post_mock.call_args + assert len(data) == 1 + assert data[0]["claims"] == claims diff --git a/sdk/identity/azure-identity/tests/test_certificate_credential.py b/sdk/identity/azure-identity/tests/test_certificate_credential.py index eed94a74206c..40be31528673 100644 --- a/sdk/identity/azure-identity/tests/test_certificate_credential.py +++ b/sdk/identity/azure-identity/tests/test_certificate_credential.py @@ -24,6 +24,8 @@ from helpers import ( build_aad_response, + build_id_token, + id_token_claims, get_discovery_response, urlsafeb64_decode, mock_response, @@ -428,3 +430,53 @@ def send(request, **_): token = credential.get_token("scope", tenant_id="un" + expected_tenant) assert token.token == expected_token + + +def test_client_capabilities(): + """The credential should configure MSAL for capability CP1 unless AZURE_IDENTITY_DISABLE_CP1 is set""" + + transport = Mock(send=Mock(side_effect=Exception("this test mocks MSAL, so no request should be sent"))) + + credential = CertificateCredential("tenant-id", "client-id", PEM_CERT_PATH, transport=transport) + with patch("msal.ConfidentialClientApplication") as ConfidentialClientApplication: + credential._get_app() + + assert ConfidentialClientApplication.call_count == 1 + _, kwargs = ConfidentialClientApplication.call_args + assert kwargs["client_capabilities"] == ["CP1"] + + credential = CertificateCredential("tenant-id", "client-id", PEM_CERT_PATH, transport=transport) + with patch.dict("os.environ", {"AZURE_IDENTITY_DISABLE_CP1": "true"}): + with patch("msal.ConfidentialClientApplication") as ConfidentialClientApplication: + credential._get_app() + + assert ConfidentialClientApplication.call_count == 1 + _, kwargs = ConfidentialClientApplication.call_args + assert kwargs["client_capabilities"] is None + + +def test_claims_challenge(): + """get_token should pass any claims challenge to MSAL token acquisition APIs""" + + msal_acquire_token_result = dict( + build_aad_response(access_token="**", id_token=build_id_token()), + id_token_claims=id_token_claims("issuer", "subject", "audience", upn="upn"), + ) + expected_claims = '{"access_token": {"essential": "true"}}' + + transport = Mock(send=Mock(side_effect=Exception("this test mocks MSAL, so no request should be sent"))) + credential = CertificateCredential("tenant-id", "client-id", PEM_CERT_PATH, transport=transport) + with patch.object(CertificateCredential, "_get_app") as get_mock_app: + msal_app = get_mock_app() + msal_app.acquire_token_silent_with_error.return_value = None + msal_app.acquire_token_for_client.return_value = msal_acquire_token_result + + credential.get_token("scope", claims=expected_claims) + + assert msal_app.acquire_token_silent_with_error.call_count == 1 + args, kwargs = msal_app.acquire_token_silent_with_error.call_args + assert kwargs["claims_challenge"] == expected_claims + + assert msal_app.acquire_token_for_client.call_count == 1 + args, kwargs = msal_app.acquire_token_for_client.call_args + assert kwargs["claims_challenge"] == expected_claims diff --git a/sdk/identity/azure-identity/tests/test_client_secret_credential.py b/sdk/identity/azure-identity/tests/test_client_secret_credential.py index 598d0011865d..91f0c1c419dd 100644 --- a/sdk/identity/azure-identity/tests/test_client_secret_credential.py +++ b/sdk/identity/azure-identity/tests/test_client_secret_credential.py @@ -12,7 +12,15 @@ import pytest from six.moves.urllib_parse import urlparse -from helpers import build_aad_response, get_discovery_response, mock_response, msal_validating_transport, Request +from helpers import ( + build_aad_response, + build_id_token, + get_discovery_response, + id_token_claims, + mock_response, + msal_validating_transport, + Request +) try: from unittest.mock import Mock, patch @@ -274,3 +282,53 @@ def send(request, **_): with patch.dict("os.environ", {EnvironmentVariables.AZURE_IDENTITY_DISABLE_MULTITENANTAUTH: "true"}): token = credential.get_token("scope", tenant_id="un" + expected_tenant) assert token.token == expected_token + + +def test_client_capabilities(): + """The credential should configure MSAL for capability CP1 unless AZURE_IDENTITY_DISABLE_CP1 is set""" + + transport = Mock(send=Mock(side_effect=Exception("this test mocks MSAL, so no request should be sent"))) + + credential = ClientSecretCredential("tenant-id", "client-id", "client-secret", transport=transport) + with patch("msal.ConfidentialClientApplication") as ConfidentialClientApplication: + credential._get_app() + + assert ConfidentialClientApplication.call_count == 1 + _, kwargs = ConfidentialClientApplication.call_args + assert kwargs["client_capabilities"] == ["CP1"] + + credential = ClientSecretCredential("tenant-id", "client-id", "client-secret", transport=transport) + with patch.dict("os.environ", {"AZURE_IDENTITY_DISABLE_CP1": "true"}): + with patch("msal.ConfidentialClientApplication") as ConfidentialClientApplication: + credential._get_app() + + assert ConfidentialClientApplication.call_count == 1 + _, kwargs = ConfidentialClientApplication.call_args + assert kwargs["client_capabilities"] is None + + +def test_claims_challenge(): + """get_token should pass any claims challenge to MSAL token acquisition APIs""" + + msal_acquire_token_result = dict( + build_aad_response(access_token="**", id_token=build_id_token()), + id_token_claims=id_token_claims("issuer", "subject", "audience", upn="upn"), + ) + expected_claims = '{"access_token": {"essential": "true"}}' + + transport = Mock(send=Mock(side_effect=Exception("this test mocks MSAL, so no request should be sent"))) + credential = ClientSecretCredential("tenant-id", "client-id", "client-secret", transport=transport) + with patch.object(ClientSecretCredential, "_get_app") as get_mock_app: + msal_app = get_mock_app() + msal_app.acquire_token_silent_with_error.return_value = None + msal_app.acquire_token_for_client.return_value = msal_acquire_token_result + + credential.get_token("scope", claims=expected_claims) + + assert msal_app.acquire_token_silent_with_error.call_count == 1 + args, kwargs = msal_app.acquire_token_silent_with_error.call_args + assert kwargs["claims_challenge"] == expected_claims + + assert msal_app.acquire_token_for_client.call_count == 1 + args, kwargs = msal_app.acquire_token_for_client.call_args + assert kwargs["claims_challenge"] == expected_claims diff --git a/sdk/identity/azure-identity/tests/test_live.py b/sdk/identity/azure-identity/tests/test_live.py index 5b2e0a2061b6..5d7d0ef57588 100644 --- a/sdk/identity/azure-identity/tests/test_live.py +++ b/sdk/identity/azure-identity/tests/test_live.py @@ -13,6 +13,8 @@ ) from azure.identity._constants import DEVELOPER_SIGN_ON_CLIENT_ID +from helpers import get_token_payload_contents + ARM_SCOPE = "https://management.azure.com/.default" @@ -21,6 +23,7 @@ def get_token(credential): assert token assert token.token assert token.expires_on + return token @pytest.mark.parametrize("certificate_fixture", ("live_pem_certificate", "live_pfx_certificate")) @@ -42,7 +45,9 @@ def test_certificate_credential(certificate_fixture, request): credential = CertificateCredential( tenant_id, client_id, certificate_data=cert["cert_with_password_bytes"], password=cert["password"] ) - get_token(credential) + token = get_token(credential) + parsed_payload = get_token_payload_contents(token.token) + assert "xms_cc" in parsed_payload and "CP1" in parsed_payload["xms_cc"] def test_client_secret_credential(live_service_principal): @@ -51,7 +56,9 @@ def test_client_secret_credential(live_service_principal): live_service_principal["client_id"], live_service_principal["client_secret"], ) - get_token(credential) + token = get_token(credential) + parsed_payload = get_token_payload_contents(token.token) + assert "xms_cc" in parsed_payload and "CP1" in parsed_payload["xms_cc"] def test_default_credential(live_service_principal): diff --git a/sdk/identity/azure-identity/tests/test_live_async.py b/sdk/identity/azure-identity/tests/test_live_async.py index f96896ec90d2..7a5fafaf84bd 100644 --- a/sdk/identity/azure-identity/tests/test_live_async.py +++ b/sdk/identity/azure-identity/tests/test_live_async.py @@ -6,6 +6,8 @@ from azure.identity.aio import DefaultAzureCredential, CertificateCredential, ClientSecretCredential +from helpers import get_token_payload_contents + ARM_SCOPE = "https://management.azure.com/.default" @@ -14,6 +16,7 @@ async def get_token(credential): assert token assert token.token assert token.expires_on + return token @pytest.mark.asyncio @@ -36,7 +39,10 @@ async def test_certificate_credential(certificate_fixture, request): credential = CertificateCredential( tenant_id, client_id, certificate_data=cert["cert_with_password_bytes"], password=cert["password"] ) - await get_token(credential) + token = await get_token(credential) + parsed_payload = get_token_payload_contents(token.token) + assert "xms_cc" in parsed_payload and "CP1" in parsed_payload["xms_cc"] + @pytest.mark.asyncio @@ -46,7 +52,10 @@ async def test_client_secret_credential(live_service_principal): live_service_principal["client_id"], live_service_principal["client_secret"], ) - await get_token(credential) + token = await get_token(credential) + parsed_payload = get_token_payload_contents(token.token) + assert "xms_cc" in parsed_payload and "CP1" in parsed_payload["xms_cc"] + @pytest.mark.asyncio