Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
aad_client_base.py377 linesDownload Raw Back to _internal
1# ------------------------------------2# Copyright (c) Microsoft Corporation.3# Licensed under the MIT License.4# ------------------------------------5import abc6import base647import json8import time9from uuid import uuid410from typing import TYPE_CHECKING, List, Any, Iterable, Optional, Union, Dict11 12from msal import TokenCache13 14from azure.core.pipeline import PipelineResponse15from azure.core.pipeline.policies import ContentDecodePolicy16from azure.core.pipeline.transport import HttpRequest17from azure.core.credentials import AccessToken18from azure.core.exceptions import ClientAuthenticationError19from .utils import get_default_authority, normalize_authority, resolve_tenant20from .aadclient_certificate import AadClientCertificate21from .._persistent_cache import _load_persistent_cache22 23 24if TYPE_CHECKING:25    from azure.core.pipeline import AsyncPipeline, Pipeline26    from azure.core.pipeline.policies import AsyncHTTPPolicy, HTTPPolicy, SansIOHTTPPolicy27    from azure.core.pipeline.transport import AsyncHttpTransport, HttpTransport28 29    PipelineType = Union[AsyncPipeline, Pipeline]30    PolicyType = Union[AsyncHTTPPolicy, HTTPPolicy, SansIOHTTPPolicy]31    TransportType = Union[AsyncHttpTransport, HttpTransport]32 33JWT_BEARER_ASSERTION = "urn:ietf:params:oauth:client-assertion-type:jwt-bearer"34 35 36class AadClientBase(abc.ABC):37    _POST = ["POST"]38 39    def __init__(40        self,41        tenant_id: str,42        client_id: str,43        authority: Optional[str] = None,44        cache: Optional[TokenCache] = None,45        cae_cache: Optional[TokenCache] = None,46        *,47        additionally_allowed_tenants: Optional[List[str]] = None,48        **kwargs: Any49    ) -> None:50        self._authority = normalize_authority(authority) if authority else get_default_authority()51 52        self._tenant_id = tenant_id53        self._client_id = client_id54        self._additionally_allowed_tenants = additionally_allowed_tenants or []55        self._pipeline = self._build_pipeline(**kwargs)56 57        self._cache = cache58        self._cae_cache = cae_cache59        self._cache_options = kwargs.pop("cache_persistence_options", None)60 61    def _get_cache(self, **kwargs: Any) -> TokenCache:62        cache = self._cae_cache if kwargs.get("enable_cae") else self._cache63        if not cache:64            cache = self._initialize_cache(is_cae=bool(kwargs.get("enable_cae")))65        return cache66 67    def _initialize_cache(self, is_cae: bool = False) -> TokenCache:68        if self._cache_options:69            if is_cae:70                self._cae_cache = _load_persistent_cache(self._cache_options, is_cae)71            else:72                self._cache = _load_persistent_cache(self._cache_options, is_cae)73        else:74            if is_cae:75                self._cae_cache = TokenCache()76            else:77                self._cache = TokenCache()78        return self._cae_cache if is_cae else self._cache79 80    def get_cached_access_token(self, scopes: Iterable[str], **kwargs: Any) -> Optional[AccessToken]:81        tenant = resolve_tenant(82            self._tenant_id, additionally_allowed_tenants=self._additionally_allowed_tenants, **kwargs83        )84 85        cache = self._get_cache(**kwargs)86        tokens = cache.find(87            TokenCache.CredentialType.ACCESS_TOKEN,88            target=list(scopes),89            query={"client_id": self._client_id, "realm": tenant},90        )91        for token in tokens:92            expires_on = int(token["expires_on"])93            if expires_on > int(time.time()):94                return AccessToken(token["secret"], expires_on)95        return None96 97    def get_cached_refresh_tokens(self, scopes: Iterable[str], **kwargs) -> List[Dict]:98        # Assumes all cached refresh tokens belong to the same user99        cache = self._get_cache(**kwargs)100        return cache.find(TokenCache.CredentialType.REFRESH_TOKEN, target=list(scopes))101 102    @abc.abstractmethod103    def obtain_token_by_authorization_code(self, scopes, code, redirect_uri, client_secret=None, **kwargs):104        pass105 106    @abc.abstractmethod107    def obtain_token_by_jwt_assertion(self, scopes, assertion, **kwargs):108        pass109 110    @abc.abstractmethod111    def obtain_token_by_client_certificate(self, scopes, certificate, **kwargs):112        pass113 114    @abc.abstractmethod115    def obtain_token_by_client_secret(self, scopes, secret, **kwargs):116        pass117 118    @abc.abstractmethod119    def obtain_token_by_refresh_token(self, scopes, refresh_token, **kwargs):120        pass121 122    @abc.abstractmethod123    def obtain_token_on_behalf_of(self, scopes, client_credential, user_assertion, **kwargs):124        pass125 126    @abc.abstractmethod127    def _build_pipeline(self, **kwargs):128        pass129 130    def _process_response(self, response: PipelineResponse, request_time: int, **kwargs) -> AccessToken:131        content = response.context.get(132            ContentDecodePolicy.CONTEXT_NAME133        ) or ContentDecodePolicy.deserialize_from_http_generics(response.http_response)134 135        cache = self._get_cache(**kwargs)136        if response.http_request.body.get("grant_type") == "refresh_token":137            if content.get("error") == "invalid_grant":138                # the request's refresh token is invalid -> evict it from the cache139                cache_entries = cache.find(140                    TokenCache.CredentialType.REFRESH_TOKEN,141                    query={"secret": response.http_request.body["refresh_token"]},142                )143                for invalid_token in cache_entries:144                    cache.remove_rt(invalid_token)145            if "refresh_token" in content:146                # Microsoft Entra ID returned a new refresh token -> update the cache entry147                cache_entries = cache.find(148                    TokenCache.CredentialType.REFRESH_TOKEN,149                    query={"secret": response.http_request.body["refresh_token"]},150                )151                # If the old token is in multiple cache entries, the cache is in a state we don't152                # expect or know how to reason about, so we update nothing.153                if len(cache_entries) == 1:154                    cache.update_rt(cache_entries[0], content["refresh_token"])155                    del content["refresh_token"]  # prevent caching a redundant entry156 157        _raise_for_error(response, content)158 159        if "expires_on" in content:160            expires_on = int(content["expires_on"])161        elif "expires_in" in content:162            expires_on = request_time + int(content["expires_in"])163        else:164            _scrub_secrets(content)165            raise ClientAuthenticationError(message="Unexpected response from Microsoft Entra ID: {}".format(content))166 167        token = AccessToken(content["access_token"], expires_on)168 169        # caching is the final step because 'add' mutates 'content'170        cache.add(171            event={172                "client_id": self._client_id,173                "response": content,174                "scope": response.http_request.body["scope"].split(),175                "token_endpoint": response.http_request.url,176            },177            now=request_time,178        )179 180        return token181 182    def _get_auth_code_request(183        self, scopes: Iterable[str], code: str, redirect_uri: str, client_secret: Optional[str] = None, **kwargs: Any184    ) -> HttpRequest:185        data = {186            "client_id": self._client_id,187            "code": code,188            "grant_type": "authorization_code",189            "redirect_uri": redirect_uri,190            "scope": " ".join(scopes),191        }192 193        claims = _merge_claims_challenge_and_capabilities(194            ["CP1"] if kwargs.get("enable_cae") else [], kwargs.get("claims")195        )196        if claims:197            data["claims"] = claims198        if client_secret:199            data["client_secret"] = client_secret200 201        request = self._post(data, **kwargs)202        return request203 204    def _get_jwt_assertion_request(self, scopes: Iterable[str], assertion: str, **kwargs: Any) -> HttpRequest:205        data = {206            "client_assertion": assertion,207            "client_assertion_type": JWT_BEARER_ASSERTION,208            "client_id": self._client_id,209            "grant_type": "client_credentials",210            "scope": " ".join(scopes),211        }212 213        claims = _merge_claims_challenge_and_capabilities(214            ["CP1"] if kwargs.get("enable_cae") else [], kwargs.get("claims")215        )216        if claims:217            data["claims"] = claims218 219        request = self._post(data, **kwargs)220        return request221 222    def _get_client_certificate_assertion(self, certificate: AadClientCertificate, **kwargs: Any) -> str:223        now = int(time.time())224        header = json.dumps({"typ": "JWT", "alg": "RS256", "x5t": certificate.thumbprint}).encode("utf-8")225        payload = json.dumps(226            {227                "jti": str(uuid4()),228                "aud": self._get_token_url(**kwargs),229                "iss": self._client_id,230                "sub": self._client_id,231                "nbf": now,232                "exp": now + (60 * 30),233            }234        ).encode("utf-8")235        jws = base64.urlsafe_b64encode(header) + b"." + base64.urlsafe_b64encode(payload)236        signature = certificate.sign(jws)237        jwt_bytes = jws + b"." + base64.urlsafe_b64encode(signature)238        return jwt_bytes.decode("utf-8")239 240    def _get_client_certificate_request(241        self, scopes: Iterable[str], certificate: AadClientCertificate, **kwargs: Any242    ) -> HttpRequest:243        assertion = self._get_client_certificate_assertion(certificate, **kwargs)244        return self._get_jwt_assertion_request(scopes, assertion, **kwargs)245 246    def _get_client_secret_request(self, scopes: Iterable[str], secret: str, **kwargs: Any) -> HttpRequest:247        data = {248            "client_id": self._client_id,249            "client_secret": secret,250            "grant_type": "client_credentials",251            "scope": " ".join(scopes),252        }253 254        claims = _merge_claims_challenge_and_capabilities(255            ["CP1"] if kwargs.get("enable_cae") else [], kwargs.get("claims")256        )257        if claims:258            data["claims"] = claims259 260        request = self._post(data, **kwargs)261        return request262 263    def _get_on_behalf_of_request(264        self,265        scopes: Iterable[str],266        client_credential: Union[str, AadClientCertificate],267        user_assertion: str,268        **kwargs: Any269    ) -> HttpRequest:270        data = {271            "assertion": user_assertion,272            "client_id": self._client_id,273            "grant_type": "urn:ietf:params:oauth:grant-type:jwt-bearer",274            "requested_token_use": "on_behalf_of",275            "scope": " ".join(scopes),276        }277 278        claims = _merge_claims_challenge_and_capabilities(279            ["CP1"] if kwargs.get("enable_cae") else [], kwargs.get("claims")280        )281        if claims:282            data["claims"] = claims283 284        if isinstance(client_credential, AadClientCertificate):285            data["client_assertion"] = self._get_client_certificate_assertion(client_credential)286            data["client_assertion_type"] = JWT_BEARER_ASSERTION287        else:288            data["client_secret"] = client_credential289 290        request = self._post(data, **kwargs)291        return request292 293    def _get_refresh_token_request(self, scopes: Iterable[str], refresh_token: str, **kwargs: Any) -> HttpRequest:294        data = {295            "grant_type": "refresh_token",296            "refresh_token": refresh_token,297            "scope": " ".join(scopes),298            "client_id": self._client_id,299            "client_info": 1,  # request Microsoft Entra ID include home_account_id in its response300        }301 302        claims = _merge_claims_challenge_and_capabilities(303            ["CP1"] if kwargs.get("enable_cae") else [], kwargs.get("claims")304        )305        if claims:306            data["claims"] = claims307 308        request = self._post(data, **kwargs)309        return request310 311    def _get_refresh_token_on_behalf_of_request(312        self,313        scopes: Iterable[str],314        client_credential: Union[str, AadClientCertificate],315        refresh_token: str,316        **kwargs: Any317    ) -> HttpRequest:318        data = {319            "grant_type": "refresh_token",320            "refresh_token": refresh_token,321            "scope": " ".join(scopes),322            "client_id": self._client_id,323            "client_info": 1,  # request Microsoft Entra ID include home_account_id in its response324        }325        claims = _merge_claims_challenge_and_capabilities(326            ["CP1"] if kwargs.get("enable_cae") else [], kwargs.get("claims")327        )328        if claims:329            data["claims"] = claims330 331        if isinstance(client_credential, AadClientCertificate):332            data["client_assertion"] = self._get_client_certificate_assertion(client_credential)333            data["client_assertion_type"] = JWT_BEARER_ASSERTION334        else:335            data["client_secret"] = client_credential336        request = self._post(data, **kwargs)337        return request338 339    def _get_token_url(self, **kwargs: Any) -> str:340        tenant = resolve_tenant(341            self._tenant_id, additionally_allowed_tenants=self._additionally_allowed_tenants, **kwargs342        )343        return "/".join((self._authority, tenant, "oauth2/v2.0/token"))344 345    def _post(self, data: Dict, **kwargs: Any) -> HttpRequest:346        url = self._get_token_url(**kwargs)347        return HttpRequest("POST", url, data=data, headers={"Content-Type": "application/x-www-form-urlencoded"})348 349 350def _merge_claims_challenge_and_capabilities(capabilities, claims_challenge):351    # Represent capabilities as {"access_token": {"xms_cc": {"values": capabilities}}}352    # and then merge/add it into incoming claims353    if not capabilities:354        return claims_challenge355    claims_dict = json.loads(claims_challenge) if claims_challenge else {}356    for key in ["access_token"]:357        claims_dict.setdefault(key, {}).update(xms_cc={"values": capabilities})358    return json.dumps(claims_dict)359 360 361def _scrub_secrets(response: Dict) -> None:362    for secret in ("access_token", "refresh_token"):363        if secret in response:364            response[secret] = "***"365 366 367def _raise_for_error(response: PipelineResponse, content: Dict) -> None:368    if "error" not in content:369        return370 371    _scrub_secrets(content)372    if "error_description" in content:373        message = "Microsoft Entra ID error '({}) {}'".format(content["error"], content["error_description"])374    else:375        message = "Microsoft Entra ID error '{}'".format(content)376    raise ClientAuthenticationError(message=message, response=response.http_response)377 
codekingpro/portable-devtools · Team Ai