codekingpro/portable-devtools
114k
1# ------------------------------------2# Copyright (c) Microsoft Corporation.3# Licensed under the MIT License.4# ------------------------------------5import time6from typing import Any, Optional, Dict7 8from azure.core.credentials import AccessToken9from azure.core.exceptions import ClientAuthenticationError10from .get_token_mixin import GetTokenMixin11 12from . import wrap_exceptions13from .msal_credentials import MsalCredential14 15 16def _get_known_kwargs(kwargs: Dict[str, Any]):17 # Remove kwargs not expected by MSAL. These aren't typically passed by users, but this is a precaution.18 known_kwargs = {"force_refresh", "authority", "correlation_id", "data"}19 return {k: v for k, v in kwargs.items() if k in known_kwargs}20 21 22class ClientCredentialBase(MsalCredential, GetTokenMixin):23 """Base class for credentials authenticating a service principal with a certificate or secret"""24 25 @wrap_exceptions26 def _acquire_token_silently(self, *scopes: str, **kwargs: Any) -> Optional[AccessToken]:27 app = self._get_app(**kwargs)28 request_time = int(time.time())29 result = app.acquire_token_silent_with_error(30 list(scopes), account=None, claims_challenge=kwargs.pop("claims", None), **_get_known_kwargs(kwargs)31 )32 if result and "access_token" in result and "expires_in" in result:33 return AccessToken(result["access_token"], request_time + int(result["expires_in"]))34 return None35 36 @wrap_exceptions37 def _request_token(self, *scopes: str, **kwargs: Any) -> Optional[AccessToken]:38 app = self._get_app(**kwargs)39 request_time = int(time.time())40 result = app.acquire_token_for_client(list(scopes), claims_challenge=kwargs.pop("claims", None))41 if "access_token" not in result:42 message = "Authentication failed: {}".format(result.get("error_description") or result.get("error"))43 raise ClientAuthenticationError(message=message)44 45 return AccessToken(result["access_token"], request_time + int(result["expires_in"]))46 