codekingpro/portable-devtools
114k
1import json2import threading3import time4import logging5 6from .authority import canonicalize7from .oauth2cli.oidc import decode_part, decode_id_token8from .oauth2cli.oauth2 import Client9 10 11logger = logging.getLogger(__name__)12_GRANT_TYPE_BROKER = "broker"13 14def is_subdict_of(small, big):15 return dict(big, **small) == big16 17def _get_username(id_token_claims):18 return id_token_claims.get(19 "preferred_username", # AAD20 id_token_claims.get("upn")) # ADFS 201921 22class TokenCache(object):23 """This is considered as a base class containing minimal cache behavior.24 25 Although it maintains tokens using unified schema across all MSAL libraries,26 this class does not serialize/persist them.27 See subclass :class:`SerializableTokenCache` for details on serialization.28 """29 30 class CredentialType:31 ACCESS_TOKEN = "AccessToken"32 REFRESH_TOKEN = "RefreshToken"33 ACCOUNT = "Account" # Not exactly a credential type, but we put it here34 ID_TOKEN = "IdToken"35 APP_METADATA = "AppMetadata"36 37 class AuthorityType:38 ADFS = "ADFS"39 MSSTS = "MSSTS" # MSSTS means AAD v2 for both AAD & MSA40 41 def __init__(self):42 self._lock = threading.RLock()43 self._cache = {}44 self.key_makers = {45 self.CredentialType.REFRESH_TOKEN:46 lambda home_account_id=None, environment=None, client_id=None,47 target=None, **ignored_payload_from_a_real_token:48 "-".join([49 home_account_id or "",50 environment or "",51 self.CredentialType.REFRESH_TOKEN,52 client_id or "",53 "", # RT is cross-tenant in AAD54 target or "", # raw value could be None if deserialized from other SDK55 ]).lower(),56 self.CredentialType.ACCESS_TOKEN:57 lambda home_account_id=None, environment=None, client_id=None,58 realm=None, target=None, **ignored_payload_from_a_real_token:59 "-".join([60 home_account_id or "",61 environment or "",62 self.CredentialType.ACCESS_TOKEN,63 client_id or "",64 realm or "",65 target or "",66 ]).lower(),67 self.CredentialType.ID_TOKEN:68 lambda home_account_id=None, environment=None, client_id=None,69 realm=None, **ignored_payload_from_a_real_token:70 "-".join([71 home_account_id or "",72 environment or "",73 self.CredentialType.ID_TOKEN,74 client_id or "",75 realm or "",76 "" # Albeit irrelevant, schema requires an empty scope here77 ]).lower(),78 self.CredentialType.ACCOUNT:79 lambda home_account_id=None, environment=None, realm=None,80 **ignored_payload_from_a_real_entry:81 "-".join([82 home_account_id or "",83 environment or "",84 realm or "",85 ]).lower(),86 self.CredentialType.APP_METADATA:87 lambda environment=None, client_id=None, **kwargs:88 "appmetadata-{}-{}".format(environment or "", client_id or ""),89 }90 91 def _get_access_token(92 self,93 home_account_id, environment, client_id, realm, target, # Together they form a compound key94 default=None,95 ): # O(1)96 return self._get(97 self.CredentialType.ACCESS_TOKEN,98 self.key_makers[TokenCache.CredentialType.ACCESS_TOKEN](99 home_account_id=home_account_id,100 environment=environment,101 client_id=client_id,102 realm=realm,103 target=" ".join(target),104 ),105 default=default)106 107 def _get_app_metadata(self, environment, client_id, default=None): # O(1)108 return self._get(109 self.CredentialType.APP_METADATA,110 self.key_makers[TokenCache.CredentialType.APP_METADATA](111 environment=environment,112 client_id=client_id,113 ),114 default=default)115 116 def _get(self, credential_type, key, default=None): # O(1)117 with self._lock:118 return self._cache.get(credential_type, {}).get(key, default)119 120 def _find(self, credential_type, target=None, query=None): # O(n) generator121 """Returns a generator of matching entries.122 123 It is O(1) for AT hits, and O(n) for other types.124 Note that it holds a lock during the entire search.125 """126 target = sorted(target or []) # Match the order sorted by add()127 assert isinstance(target, list), "Invalid parameter type"128 129 preferred_result = None130 if (credential_type == self.CredentialType.ACCESS_TOKEN131 and isinstance(query, dict)132 and "home_account_id" in query and "environment" in query133 and "client_id" in query and "realm" in query and target134 ): # Special case for O(1) AT lookup135 preferred_result = self._get_access_token(136 query["home_account_id"], query["environment"],137 query["client_id"], query["realm"], target)138 if preferred_result:139 yield preferred_result140 141 target_set = set(target)142 with self._lock:143 # Since the target inside token cache key is (per schema) unsorted,144 # there is no point to attempt an O(1) key-value search here.145 # So we always do an O(n) in-memory search.146 for entry in self._cache.get(credential_type, {}).values():147 if is_subdict_of(query or {}, entry) and (148 target_set <= set(entry.get("target", "").split())149 if target else True):150 if entry != preferred_result: # Avoid yielding the same entry twice151 yield entry152 153 def find(self, credential_type, target=None, query=None): # Obsolete. Use _find() instead.154 return list(self._find(credential_type, target=target, query=query))155 156 def add(self, event, now=None):157 """Handle a token obtaining event, and add tokens into cache."""158 def make_clean_copy(dictionary, sensitive_fields): # Masks sensitive info159 return {160 k: "********" if k in sensitive_fields else v161 for k, v in dictionary.items()162 }163 clean_event = dict(164 event,165 data=make_clean_copy(event.get("data", {}), (166 "password", "client_secret", "refresh_token", "assertion",167 )),168 response=make_clean_copy(event.get("response", {}), (169 "id_token_claims", # Provided by broker170 "access_token", "refresh_token", "id_token", "username",171 )),172 )173 logger.debug("event=%s", json.dumps(174 # We examined and concluded that this log won't have Log Injection risk,175 # because the event payload is already in JSON so CR/LF will be escaped.176 clean_event,177 indent=4, sort_keys=True,178 default=str, # assertion is in bytes in Python 3179 ))180 return self.__add(event, now=now)181 182 def __parse_account(self, response, id_token_claims):183 """Return client_info and home_account_id"""184 if "client_info" in response: # It happens when client_info and profile are in request185 client_info = json.loads(decode_part(response["client_info"]))186 if "uid" in client_info and "utid" in client_info:187 return client_info, "{uid}.{utid}".format(**client_info)188 # https://github.com/AzureAD/microsoft-authentication-library-for-python/issues/387189 if id_token_claims: # This would be an end user on ADFS-direct scenario190 sub = id_token_claims["sub"] # "sub" always exists, per OIDC specs191 return {"uid": sub}, sub192 # client_credentials flow will reach this code path193 return {}, None194 195 def __add(self, event, now=None):196 # event typically contains: client_id, scope, token_endpoint,197 # response, params, data, grant_type198 environment = realm = None199 if "token_endpoint" in event:200 _, environment, realm = canonicalize(event["token_endpoint"])201 if "environment" in event: # Always available unless in legacy test cases202 environment = event["environment"] # Set by application.py203 response = event.get("response", {})204 data = event.get("data", {})205 access_token = response.get("access_token")206 refresh_token = response.get("refresh_token")207 id_token = response.get("id_token")208 id_token_claims = response.get("id_token_claims") or ( # Prefer the claims from broker209 # Only use decode_id_token() when necessary, it contains time-sensitive validation210 decode_id_token(id_token, client_id=event["client_id"]) if id_token else {})211 client_info, home_account_id = self.__parse_account(response, id_token_claims)212 213 target = ' '.join(sorted(event.get("scope") or [])) # Schema should have required sorting214 215 with self._lock:216 now = int(time.time() if now is None else now)217 218 if access_token:219 default_expires_in = ( # https://www.rfc-editor.org/rfc/rfc6749#section-5.1220 int(response.get("expires_on")) - now # Some Managed Identity emits this221 ) if response.get("expires_on") else 600222 expires_in = int( # AADv1-like endpoint returns a string223 response.get("expires_in", default_expires_in))224 ext_expires_in = int( # AADv1-like endpoint returns a string225 response.get("ext_expires_in", expires_in))226 at = {227 "credential_type": self.CredentialType.ACCESS_TOKEN,228 "secret": access_token,229 "home_account_id": home_account_id,230 "environment": environment,231 "client_id": event.get("client_id"),232 "target": target,233 "realm": realm,234 "token_type": response.get("token_type", "Bearer"),235 "cached_at": str(now), # Schema defines it as a string236 "expires_on": str(now + expires_in), # Same here237 "extended_expires_on": str(now + ext_expires_in) # Same here238 }239 if data.get("key_id"): # It happens in SSH-cert or POP scenario240 at["key_id"] = data.get("key_id")241 if "refresh_in" in response:242 refresh_in = response["refresh_in"] # It is an integer243 at["refresh_on"] = str(now + refresh_in) # Schema wants a string244 self.modify(self.CredentialType.ACCESS_TOKEN, at, at)245 246 if client_info and not event.get("skip_account_creation"):247 account = {248 "home_account_id": home_account_id,249 "environment": environment,250 "realm": realm,251 "local_account_id": event.get(252 "_account_id", # Came from mid-tier code path.253 # Emperically, it is the oid in AAD or cid in MSA.254 id_token_claims.get("oid", id_token_claims.get("sub"))),255 "username": _get_username(id_token_claims)256 or data.get("username") # Falls back to ROPC username257 or event.get("username") # Falls back to Federated ROPC username258 or "", # The schema does not like null259 "authority_type": event.get(260 "authority_type", # Honor caller's choice of authority_type261 self.AuthorityType.ADFS if realm == "adfs"262 else self.AuthorityType.MSSTS),263 # "client_info": response.get("client_info"), # Optional264 }265 grant_types_that_establish_an_account = (266 _GRANT_TYPE_BROKER, "authorization_code", "password",267 Client.DEVICE_FLOW["GRANT_TYPE"])268 if event.get("grant_type") in grant_types_that_establish_an_account:269 account["account_source"] = event["grant_type"]270 self.modify(self.CredentialType.ACCOUNT, account, account)271 272 if id_token:273 idt = {274 "credential_type": self.CredentialType.ID_TOKEN,275 "secret": id_token,276 "home_account_id": home_account_id,277 "environment": environment,278 "realm": realm,279 "client_id": event.get("client_id"),280 # "authority": "it is optional",281 }282 self.modify(self.CredentialType.ID_TOKEN, idt, idt)283 284 if refresh_token:285 rt = {286 "credential_type": self.CredentialType.REFRESH_TOKEN,287 "secret": refresh_token,288 "home_account_id": home_account_id,289 "environment": environment,290 "client_id": event.get("client_id"),291 "target": target, # Optional per schema though292 "last_modification_time": str(now), # Optional. Schema defines it as a string.293 }294 if "foci" in response:295 rt["family_id"] = response["foci"]296 self.modify(self.CredentialType.REFRESH_TOKEN, rt, rt)297 298 app_metadata = {299 "client_id": event.get("client_id"),300 "environment": environment,301 }302 if "foci" in response:303 app_metadata["family_id"] = response.get("foci")304 self.modify(self.CredentialType.APP_METADATA, app_metadata, app_metadata)305 306 def modify(self, credential_type, old_entry, new_key_value_pairs=None):307 # Modify the specified old_entry with new_key_value_pairs,308 # or remove the old_entry if the new_key_value_pairs is None.309 310 # This helper exists to consolidate all token add/modify/remove behaviors,311 # so that the sub-classes will have only one method to work on,312 # instead of patching a pair of update_xx() and remove_xx() per type.313 # You can monkeypatch self.key_makers to support more types on-the-fly.314 key = self.key_makers[credential_type](**old_entry)315 with self._lock:316 if new_key_value_pairs: # Update with them317 entries = self._cache.setdefault(credential_type, {})318 entries[key] = dict(319 old_entry, # Do not use entries[key] b/c it might not exist320 **new_key_value_pairs)321 else: # Remove old_entry322 self._cache.setdefault(credential_type, {}).pop(key, None)323 324 def remove_rt(self, rt_item):325 assert rt_item.get("credential_type") == self.CredentialType.REFRESH_TOKEN326 return self.modify(self.CredentialType.REFRESH_TOKEN, rt_item)327 328 def update_rt(self, rt_item, new_rt):329 assert rt_item.get("credential_type") == self.CredentialType.REFRESH_TOKEN330 return self.modify(self.CredentialType.REFRESH_TOKEN, rt_item, {331 "secret": new_rt,332 "last_modification_time": str(int(time.time())), # Optional. Schema defines it as a string.333 })334 335 def remove_at(self, at_item):336 assert at_item.get("credential_type") == self.CredentialType.ACCESS_TOKEN337 return self.modify(self.CredentialType.ACCESS_TOKEN, at_item)338 339 def remove_idt(self, idt_item):340 assert idt_item.get("credential_type") == self.CredentialType.ID_TOKEN341 return self.modify(self.CredentialType.ID_TOKEN, idt_item)342 343 def remove_account(self, account_item):344 assert "authority_type" in account_item345 return self.modify(self.CredentialType.ACCOUNT, account_item)346 347 348class SerializableTokenCache(TokenCache):349 """This serialization can be a starting point to implement your own persistence.350 351 This class does NOT actually persist the cache on disk/db/etc..352 Depending on your need,353 the following simple recipe for file-based persistence may be sufficient::354 355 import os, atexit, msal356 cache = msal.SerializableTokenCache()357 if os.path.exists("my_cache.bin"):358 cache.deserialize(open("my_cache.bin", "r").read())359 atexit.register(lambda:360 open("my_cache.bin", "w").write(cache.serialize())361 # Hint: The following optional line persists only when state changed362 if cache.has_state_changed else None363 )364 app = msal.ClientApplication(..., token_cache=cache)365 ...366 367 :var bool has_state_changed:368 Indicates whether the cache state in the memory has changed since last369 :func:`~serialize` or :func:`~deserialize` call.370 """371 has_state_changed = False372 373 def add(self, event, **kwargs):374 super(SerializableTokenCache, self).add(event, **kwargs)375 self.has_state_changed = True376 377 def modify(self, credential_type, old_entry, new_key_value_pairs=None):378 super(SerializableTokenCache, self).modify(379 credential_type, old_entry, new_key_value_pairs)380 self.has_state_changed = True381 382 def deserialize(self, state):383 # type: (Optional[str]) -> None384 """Deserialize the cache from a state previously obtained by serialize()"""385 with self._lock:386 self._cache = json.loads(state) if state else {}387 self.has_state_changed = False # reset388 389 def serialize(self):390 # type: () -> str391 """Serialize the current cache state into a string."""392 with self._lock:393 self.has_state_changed = False394 return json.dumps(self._cache, indent=4)395 396 