Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
token_cache.py396 linesDownload Raw Back to msal
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 
codekingpro/portable-devtools · Team Ai