Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
shared_token_cache.py270 linesDownload Raw Back to _internal
1# ------------------------------------2# Copyright (c) Microsoft Corporation.3# Licensed under the MIT License.4# ------------------------------------5import abc6import platform7import time8from typing import Any, Iterable, List, Mapping, Optional, cast9from urllib.parse import urlparse10import msal11 12from azure.core.credentials import AccessToken13from .. import CredentialUnavailableError14from .._constants import KnownAuthorities15from .._internal import get_default_authority, normalize_authority, wrap_exceptions16from .._persistent_cache import _load_persistent_cache, TokenCachePersistenceOptions17from .._internal import AadClientBase18 19ABC = abc.ABC20CacheItem = Mapping[str, str]21 22MULTIPLE_ACCOUNTS = """SharedTokenCacheCredential authentication unavailable. Multiple accounts23were found in the cache. Use username and tenant id to disambiguate."""24 25MULTIPLE_MATCHING_ACCOUNTS = """SharedTokenCacheCredential authentication unavailable. Multiple accounts26matching the specified{}{} were found in the cache."""27 28NO_ACCOUNTS = """SharedTokenCacheCredential authentication unavailable. No accounts were found in the cache."""29 30NO_MATCHING_ACCOUNTS = """SharedTokenCacheCredential authentication unavailable. No account31matching the specified{}{} was found in the cache."""32 33NO_TOKEN = """Token acquisition failed for user '{}'. To fix, re-authenticate34through developer tooling supporting Azure single sign on"""35 36# build a dictionary {authority: {its known aliases}}, aliases taken from MSAL.NET's KnownMetadataProvider37KNOWN_ALIASES = {38    alias: aliases  # N.B. aliases includes alias itself39    for aliases in (40        frozenset((KnownAuthorities.AZURE_CHINA, "login.partner.microsoftonline.cn")),41        frozenset((KnownAuthorities.AZURE_PUBLIC_CLOUD, "login.windows.net", "login.microsoft.com", "sts.windows.net")),42        frozenset((KnownAuthorities.AZURE_GOVERNMENT, "login.usgovcloudapi.net")),43    )44    for alias in aliases45}46 47 48def _account_to_string(account):49    username = account.get("username")50    home_account_id = account.get("home_account_id", "").split(".")51    tenant_id = home_account_id[-1] if len(home_account_id) == 2 else ""52    return "(username: {}, tenant: {})".format(username, tenant_id)53 54 55def _filtered_accounts(56    accounts: Iterable[CacheItem], username: Optional[str] = None, tenant_id: Optional[str] = None57) -> List[CacheItem]:58    """Return accounts matching username and/or tenant_id.59 60    :param accounts: accounts from the MSAL cache61    :type accounts: Iterable[CacheItem]62    :param str username: an account's username63    :param str tenant_id: an account's tenant ID64    :return: accounts matching username and/or tenant_id65    :rtype: list[CacheItem]66    """67 68    filtered_accounts = []69    for account in accounts:70        if username and account.get("username") != username:71            continue72        if tenant_id:73            try:74                _, tenant = account["home_account_id"].split(".")75                if tenant_id != tenant:76                    continue77            except Exception:  # pylint:disable=broad-except78                continue79        filtered_accounts.append(account)80    return filtered_accounts81 82 83class SharedTokenCacheBase(ABC):  # pylint: disable=too-many-instance-attributes84    def __init__(85        self,86        username: Optional[str] = None,87        *,88        authority: Optional[str] = None,89        tenant_id: Optional[str] = None,90        **kwargs: Any91    ) -> None:  # pylint:disable=unused-argument92        self._authority = normalize_authority(authority) if authority else get_default_authority()93        environment = urlparse(self._authority).netloc94        self._environment_aliases = KNOWN_ALIASES.get(environment) or frozenset((environment,))95        self._username = username96        self._tenant_id = tenant_id97        self._cache = kwargs.pop("_cache", None)98        self._cae_cache = kwargs.pop("_cae_cache", None)99        self._cache_persistence_options = kwargs.pop("cache_persistence_options", None)100        self._client_kwargs = kwargs101        self._client_kwargs["tenant_id"] = "organizations"102        self._client = cast(AadClientBase, None)103        self._client_initialized = False104 105    def _initialize_client(self) -> None:106        if self._client_initialized:107            return108 109        self._client = self._get_auth_client(110            authority=self._authority, cache=self._cache, cae_cache=self._cae_cache, **self._client_kwargs111        )112        self._client_initialized = True113 114    def _initialize_cache(self, is_cae: bool = False) -> Optional[msal.TokenCache]:115 116        # If no cache options were provided, the default cache will be used. This credential accepts the117        # user's default cache regardless of whether it's encrypted. It doesn't create a new cache. If the118        # default cache exists, the user must have created it earlier. If it's unencrypted, the user must119        # have allowed that.120        cache_options = self._cache_persistence_options or TokenCachePersistenceOptions(allow_unencrypted_storage=True)121        if not self.supported():122            raise CredentialUnavailableError(message="Shared token cache is not supported on this platform.")123 124        if not self._cache and not is_cae:125            try:126                self._cache = _load_persistent_cache(cache_options, is_cae)127                self._client._cache = self._cache  # pylint:disable=protected-access128            except Exception:  # pylint:disable=broad-except129                return None130 131        if not self._cae_cache and is_cae:132            try:133                self._cae_cache = _load_persistent_cache(cache_options, is_cae)134                self._client._cae_cache = self._cae_cache  # pylint:disable=protected-access135            except Exception:  # pylint:disable=broad-except136                return None137 138        return self._cae_cache if is_cae else self._cache139 140    @abc.abstractmethod141    def _get_auth_client(self, **kwargs) -> AadClientBase:142        pass143 144    def _get_cache_items_for_authority(145        self, credential_type: msal.TokenCache.CredentialType, is_cae: bool = False146    ) -> List[CacheItem]:147        """Return cache items matching this credential's authority or one of its aliases.148 149        :param credential_type: the type of credential to look for in the cache150        :param bool is_cae: whether to look in the CAE cache151        :type credential_type: msal.TokenCache.CredentialType152        :return: a list of cache items153        :rtype: list[CacheItem]154        """155 156        cache = self._cae_cache if is_cae else self._cache157        items = []158        for item in cache.find(credential_type):159            environment = item.get("environment")160            if environment in self._environment_aliases:161                items.append(item)162        return items163 164    def _get_accounts_having_matching_refresh_tokens(self, is_cae: bool = False) -> Iterable[CacheItem]:165        """Returns an iterable of cached accounts which have a matching refresh token.166 167        :param bool is_cae: whether to look in the CAE cache168        :return: an iterable of cached accounts169        :rtype: Iterable[CacheItem]170        """171 172        refresh_tokens = self._get_cache_items_for_authority(msal.TokenCache.CredentialType.REFRESH_TOKEN, is_cae)173        all_accounts = self._get_cache_items_for_authority(msal.TokenCache.CredentialType.ACCOUNT, is_cae)174 175        accounts = {}176        for refresh_token in refresh_tokens:177            home_account_id = refresh_token.get("home_account_id")178            if not home_account_id:179                continue180            for account in all_accounts:181                # When the token has no family, msal.net falls back to matching client_id,182                # which won't work for the shared cache because we don't know the IDs of183                # all contributing apps. It should be unnecessary anyway because the184                # apps should all belong to the family.185                if home_account_id == account.get("home_account_id") and "family_id" in refresh_token:186                    accounts[account["home_account_id"]] = account187        return accounts.values()188 189    @wrap_exceptions190    def _get_account(191        self, username: Optional[str] = None, tenant_id: Optional[str] = None, is_cae: bool = False192    ) -> CacheItem:193        """Returns exactly one account which has a refresh token and matches username and/or tenant_id.194 195        :param str username: an account's username196        :param str tenant_id: an account's tenant ID197        :param bool is_cae: whether to use the CAE cache198        :return: an account199        :rtype: CacheItem200        """201 202        accounts = self._get_accounts_having_matching_refresh_tokens(is_cae)203        if not accounts:204            # cache is empty or contains no refresh token -> user needs to sign in205            raise CredentialUnavailableError(message=NO_ACCOUNTS)206 207        filtered_accounts = _filtered_accounts(accounts, username, tenant_id)208        if len(filtered_accounts) == 1:209            return filtered_accounts[0]210 211        # no, or multiple, accounts after filtering -> choose the best error message212        cached_accounts = ", ".join(_account_to_string(account) for account in accounts)213        if username or tenant_id:214            username_string = " username: {}".format(username) if username else ""215            tenant_string = " tenant: {}".format(tenant_id) if tenant_id else ""216            if filtered_accounts:217                message = MULTIPLE_MATCHING_ACCOUNTS.format(username_string, tenant_string)218            else:219                message = NO_MATCHING_ACCOUNTS.format(username_string, tenant_string)220        else:221            message = MULTIPLE_ACCOUNTS.format(cached_accounts)222 223        raise CredentialUnavailableError(message=message)224 225    def _get_cached_access_token(226        self, scopes: Iterable[str], account: CacheItem, is_cae: bool = False227    ) -> Optional[AccessToken]:228        if "home_account_id" not in account:229            return None230 231        cache = self._cae_cache if is_cae else self._cache232        try:233            cache_entries = cache.find(234                msal.TokenCache.CredentialType.ACCESS_TOKEN,235                target=list(scopes),236                query={"home_account_id": account["home_account_id"]},237            )238            for token in cache_entries:239                expires_on = int(token["expires_on"])240                if expires_on - 300 > int(time.time()):241                    return AccessToken(token["secret"], expires_on)242        except Exception as ex:  # pylint:disable=broad-except243            message = "Error accessing cached data: {}".format(ex)244            raise CredentialUnavailableError(message=message) from ex245 246        return None247 248    def _get_refresh_tokens(self, account, is_cae: bool = False) -> List[str]:249        if "home_account_id" not in account:250            return []251 252        cache = self._cae_cache if is_cae else self._cache253        try:254            cache_entries = cache.find(255                msal.TokenCache.CredentialType.REFRESH_TOKEN, query={"home_account_id": account["home_account_id"]}256            )257            return [token["secret"] for token in cache_entries if "secret" in token]258        except Exception as ex:  # pylint:disable=broad-except259            message = "Error accessing cached data: {}".format(ex)260            raise CredentialUnavailableError(message=message) from ex261 262    @staticmethod263    def supported() -> bool:264        """Whether the shared token cache is supported on the current platform.265 266        :return: True if the shared token cache is supported on the current platform.267        :rtype: bool268        """269        return platform.system() in {"Darwin", "Linux", "Windows"}270 
codekingpro/portable-devtools · Team Ai