From f31073a6e09675138be524350cd01856259b4959 Mon Sep 17 00:00:00 2001 From: Paul Van Eck Date: Mon, 20 Feb 2023 13:44:35 -0800 Subject: [PATCH] [Identity] Add additional CAE support Continous access evaluation support (CAE) was already supported in many of the credentials through MSAL (primarily user credentials). This enables it in several other credential types (i.e. service principal creds) such as ClientAssertionCredential and many of the Async Credentials. Signed-off-by: Paul Van Eck --- sdk/identity/azure-identity/CHANGELOG.md | 1 + .../azure/identity/_constants.py | 2 + .../azure/identity/_credentials/silent.py | 3 +- .../identity/_internal/aad_client_base.py | 43 +++++++++++++ .../_internal/client_credential_base.py | 5 +- .../identity/_internal/msal_credentials.py | 2 +- sdk/identity/azure-identity/tests/helpers.py | 6 ++ .../azure-identity/tests/test_aad_client.py | 48 +++++++++++++++ .../tests/test_certificate_credential.py | 52 ++++++++++++++++ .../tests/test_client_secret_credential.py | 60 ++++++++++++++++++- .../azure-identity/tests/test_live.py | 11 +++- .../azure-identity/tests/test_live_async.py | 13 +++- 12 files changed, 237 insertions(+), 9 deletions(-) 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