codekingpro/portable-devtools
114k
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 