Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -53,7 +53,12 @@ def close(self) -> None:

@log_get_token("SharedTokenCacheCredential")
def get_token(
self, *scopes: str, claims: Optional[str] = None, tenant_id: Optional[str] = None, **kwargs: Any
self,
*scopes: str,
claims: Optional[str] = None,
tenant_id: Optional[str] = None,
enable_cae: bool = False,
**kwargs: Any
) -> AccessToken:
"""Get an access token for `scopes` from the shared cache.

Expand All @@ -77,7 +82,7 @@ def get_token(
:raises ~azure.core.exceptions.ClientAuthenticationError: authentication failed. The error's ``message``
attribute gives a reason.
"""
return self._credential.get_token(*scopes, claims=claims, tenant_id=tenant_id, **kwargs)
return self._credential.get_token(*scopes, claims=claims, tenant_id=tenant_id, enable_cae=enable_cae, **kwargs)

@staticmethod
def supported() -> bool:
Expand All @@ -102,15 +107,20 @@ def __exit__(self, *args):
self._client.__exit__(*args)

def get_token(
self, *scopes: str, claims: Optional[str] = None, tenant_id: Optional[str] = None, **kwargs: Any
self,
*scopes: str,
claims: Optional[str] = None,
tenant_id: Optional[str] = None,
enable_cae: bool = False,
**kwargs: Any
) -> AccessToken:
if not scopes:
raise ValueError("'get_token' requires at least one scope")

if not self._client_initialized:
self._initialize_client()

is_cae = bool(kwargs.get("enable_cae", False))
is_cae = enable_cae
token_cache = self._cae_cache if is_cae else self._cache

# Try to load the cache if it is None.
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -54,7 +54,12 @@ def _should_refresh(self, token: AccessToken) -> bool:
return True

def get_token(
self, *scopes: str, claims: Optional[str] = None, tenant_id: Optional[str] = None, **kwargs: Any
self,
*scopes: str,
claims: Optional[str] = None,
tenant_id: Optional[str] = None,
enable_cae: bool = False,
**kwargs: Any
) -> AccessToken:
"""Request an access token for `scopes`.

Expand All @@ -80,14 +85,20 @@ def get_token(
raise ValueError('"get_token" requires at least one scope')

try:
token = self._acquire_token_silently(*scopes, claims=claims, tenant_id=tenant_id, **kwargs)
token = self._acquire_token_silently(
*scopes, claims=claims, tenant_id=tenant_id, enable_cae=enable_cae, **kwargs
)
if not token:
self._last_request_time = int(time.time())
token = self._request_token(*scopes, claims=claims, tenant_id=tenant_id, **kwargs)
token = self._request_token(
*scopes, claims=claims, tenant_id=tenant_id, enable_cae=enable_cae, **kwargs
)
elif self._should_refresh(token):
try:
self._last_request_time = int(time.time())
token = self._request_token(*scopes, claims=claims, tenant_id=tenant_id, **kwargs)
token = self._request_token(
*scopes, claims=claims, tenant_id=tenant_id, enable_cae=enable_cae, **kwargs
)
except Exception: # pylint:disable=broad-except
pass
_LOGGER.log(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,7 @@
import json
import logging
import time
from typing import Any, Optional
from typing import Any, Optional, Iterable
from urllib.parse import urlparse

from azure.core.credentials import AccessToken
Expand Down Expand Up @@ -112,7 +112,12 @@ def __init__(
super(InteractiveCredential, self).__init__(**kwargs)

def get_token(
self, *scopes: str, claims: Optional[str] = None, tenant_id: Optional[str] = None, **kwargs: Any
self,
*scopes: str,
claims: Optional[str] = None,
tenant_id: Optional[str] = None,
enable_cae: bool = False,
**kwargs: Any
) -> AccessToken:
"""Request an access token for `scopes`.

Expand Down Expand Up @@ -142,7 +147,9 @@ def get_token(

allow_prompt = kwargs.pop("_allow_prompt", not self._disable_automatic_authentication)
try:
token = self._acquire_token_silent(*scopes, claims=claims, tenant_id=tenant_id, **kwargs)
token = self._acquire_token_silent(
*scopes, claims=claims, tenant_id=tenant_id, enable_cae=enable_cae, **kwargs
)
_LOGGER.info("%s.get_token succeeded", self.__class__.__name__)
return token
except Exception as ex: # pylint:disable=broad-except
Expand All @@ -159,7 +166,7 @@ def get_token(
now = int(time.time())

try:
result = self._request_token(*scopes, claims=claims, tenant_id=tenant_id, **kwargs)
result = self._request_token(*scopes, claims=claims, tenant_id=tenant_id, enable_cae=enable_cae, **kwargs)
if "access_token" not in result:
message = "Authentication failed: {}".format(result.get("error_description") or result.get("error"))
response = self._client.get_error_response(result)
Expand All @@ -179,7 +186,9 @@ def get_token(
_LOGGER.info("%s.get_token succeeded", self.__class__.__name__)
return AccessToken(result["access_token"], now + int(result["expires_in"]))

def authenticate(self, **kwargs: Any) -> AuthenticationRecord:
def authenticate(
self, *, scopes: Optional[Iterable[str]] = None, claims: Optional[str] = None, **kwargs: Any
) -> AuthenticationRecord:
"""Interactively authenticate a user.

:keyword Iterable[str] scopes: scopes to request during authentication, such as those provided by
Expand All @@ -192,7 +201,6 @@ def authenticate(self, **kwargs: Any) -> AuthenticationRecord:
attribute gives a reason.
"""

scopes = kwargs.pop("scopes", None)
if not scopes:
if self._authority not in _DEFAULT_AUTHENTICATE_SCOPES:
# the credential is configured to use a cloud whose ARM scope we can't determine
Expand All @@ -202,7 +210,7 @@ def authenticate(self, **kwargs: Any) -> AuthenticationRecord:

scopes = _DEFAULT_AUTHENTICATE_SCOPES[self._authority]

_ = self.get_token(*scopes, _allow_prompt=True, **kwargs)
_ = self.get_token(*scopes, _allow_prompt=True, claims=claims, **kwargs)
return self._auth_record # type: ignore

@wrap_exceptions
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -43,7 +43,12 @@ async def close(self) -> None:

@log_get_token_async
async def get_token(
self, *scopes: str, claims: Optional[str] = None, tenant_id: Optional[str] = None, **kwargs: Any
self,
*scopes: str,
claims: Optional[str] = None,
tenant_id: Optional[str] = None,
enable_cae: bool = False,
**kwargs: Any
) -> AccessToken:
"""Get an access token for `scopes` from the shared cache.

Expand Down Expand Up @@ -74,7 +79,7 @@ async def get_token(
if not self._client_initialized:
self._initialize_client()

is_cae = bool(kwargs.get("enable_cae", False))
is_cae = enable_cae
token_cache = self._cae_cache if is_cae else self._cache

# Try to load the cache if it is None.
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -54,7 +54,12 @@ def _should_refresh(self, token: AccessToken) -> bool:
return True

async def get_token(
self, *scopes: str, claims: Optional[str] = None, tenant_id: Optional[str] = None, **kwargs: Any
self,
*scopes: str,
claims: Optional[str] = None,
tenant_id: Optional[str] = None,
enable_cae: bool = False,
**kwargs: Any
) -> AccessToken:
"""Request an access token for `scopes`.

Expand All @@ -80,14 +85,20 @@ async def get_token(
raise ValueError('"get_token" requires at least one scope')

try:
token = await self._acquire_token_silently(*scopes, claims=claims, tenant_id=tenant_id, **kwargs)
token = await self._acquire_token_silently(
*scopes, claims=claims, tenant_id=tenant_id, enable_cae=enable_cae, **kwargs
)
if not token:
self._last_request_time = int(time.time())
token = await self._request_token(*scopes, claims=claims, tenant_id=tenant_id, **kwargs)
token = await self._request_token(
*scopes, claims=claims, tenant_id=tenant_id, enable_cae=enable_cae, **kwargs
)
elif self._should_refresh(token):
try:
self._last_request_time = int(time.time())
token = await self._request_token(*scopes, claims=claims, tenant_id=tenant_id, **kwargs)
token = await self._request_token(
*scopes, claims=claims, tenant_id=tenant_id, enable_cae=enable_cae, **kwargs
)
except Exception: # pylint:disable=broad-except
pass
_LOGGER.log(
Expand Down
20 changes: 10 additions & 10 deletions sdk/identity/azure-identity/tests/test_get_token_mixin.py
Original file line number Diff line number Diff line change
Expand Up @@ -40,8 +40,8 @@ def test_no_cached_token():
credential = MockCredential()
token = credential.get_token(SCOPE)

credential.acquire_token_silently.assert_called_once_with(SCOPE, claims=None, tenant_id=None)
credential.request_token.assert_called_once_with(SCOPE, claims=None, tenant_id=None)
credential.acquire_token_silently.assert_called_once_with(SCOPE, claims=None, enable_cae=False, tenant_id=None)
credential.request_token.assert_called_once_with(SCOPE, claims=None, enable_cae=False, tenant_id=None)
assert token.token == MockCredential.NEW_TOKEN.token


Expand All @@ -61,7 +61,7 @@ def test_token_acquisition_failure():
with pytest.raises(Exception):
credential.get_token(SCOPE)
assert credential.request_token.call_count == i + 1
credential.request_token.assert_called_with(SCOPE, claims=None, tenant_id=None)
credential.request_token.assert_called_with(SCOPE, claims=None, enable_cae=False, tenant_id=None)


def test_expired_token():
Expand All @@ -71,8 +71,8 @@ def test_expired_token():
credential = MockCredential(cached_token=AccessToken(CACHED_TOKEN, now - 1))
token = credential.get_token(SCOPE)

credential.acquire_token_silently.assert_called_once_with(SCOPE, claims=None, tenant_id=None)
credential.request_token.assert_called_once_with(SCOPE, claims=None, tenant_id=None)
credential.acquire_token_silently.assert_called_once_with(SCOPE, claims=None, enable_cae=False, tenant_id=None)
credential.request_token.assert_called_once_with(SCOPE, claims=None, enable_cae=False, tenant_id=None)
assert token.token == MockCredential.NEW_TOKEN.token


Expand All @@ -82,7 +82,7 @@ def test_cached_token_outside_refresh_window():
credential = MockCredential(cached_token=AccessToken(CACHED_TOKEN, time.time() + DEFAULT_REFRESH_OFFSET + 1))
token = credential.get_token(SCOPE)

credential.acquire_token_silently.assert_called_once_with(SCOPE, claims=None, tenant_id=None)
credential.acquire_token_silently.assert_called_once_with(SCOPE, claims=None, enable_cae=False, tenant_id=None)
assert credential.request_token.call_count == 0
assert token.token == CACHED_TOKEN

Expand All @@ -93,8 +93,8 @@ def test_cached_token_within_refresh_window():
credential = MockCredential(cached_token=AccessToken(CACHED_TOKEN, time.time() + DEFAULT_REFRESH_OFFSET - 1))
token = credential.get_token(SCOPE)

credential.acquire_token_silently.assert_called_once_with(SCOPE, claims=None, tenant_id=None)
credential.request_token.assert_called_once_with(SCOPE, claims=None, tenant_id=None)
credential.acquire_token_silently.assert_called_once_with(SCOPE, claims=None, enable_cae=False, tenant_id=None)
credential.request_token.assert_called_once_with(SCOPE, claims=None, enable_cae=False, tenant_id=None)
assert token.token == MockCredential.NEW_TOKEN.token


Expand All @@ -109,5 +109,5 @@ def test_retry_delay():
for i in range(4):
token = credential.get_token(SCOPE)
assert token.token == CACHED_TOKEN
credential.acquire_token_silently.assert_called_with(SCOPE, claims=None, tenant_id=None)
credential.request_token.assert_called_once_with(SCOPE, claims=None, tenant_id=None)
credential.acquire_token_silently.assert_called_with(SCOPE, claims=None, enable_cae=False, tenant_id=None)
credential.request_token.assert_called_once_with(SCOPE, claims=None, enable_cae=False, tenant_id=None)
20 changes: 10 additions & 10 deletions sdk/identity/azure-identity/tests/test_get_token_mixin_async.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,8 +43,8 @@ async def test_no_cached_token():
credential = MockCredential()
token = await credential.get_token(SCOPE)

credential.acquire_token_silently.assert_called_once_with(SCOPE, claims=None, tenant_id=None)
credential.request_token.assert_called_once_with(SCOPE, claims=None, tenant_id=None)
credential.acquire_token_silently.assert_called_once_with(SCOPE, claims=None, enable_cae=False, tenant_id=None)
credential.request_token.assert_called_once_with(SCOPE, claims=None, enable_cae=False, tenant_id=None)
assert token.token == MockCredential.NEW_TOKEN.token


Expand All @@ -64,7 +64,7 @@ async def test_token_acquisition_failure():
with pytest.raises(Exception):
await credential.get_token(SCOPE)
assert credential.request_token.call_count == i + 1
credential.request_token.assert_called_with(SCOPE, claims=None, tenant_id=None)
credential.request_token.assert_called_with(SCOPE, claims=None, enable_cae=False, tenant_id=None)


async def test_expired_token():
Expand All @@ -74,8 +74,8 @@ async def test_expired_token():
credential = MockCredential(cached_token=AccessToken(CACHED_TOKEN, now - 1))
token = await credential.get_token(SCOPE)

credential.acquire_token_silently.assert_called_once_with(SCOPE, claims=None, tenant_id=None)
credential.request_token.assert_called_once_with(SCOPE, claims=None, tenant_id=None)
credential.acquire_token_silently.assert_called_once_with(SCOPE, claims=None, enable_cae=False, tenant_id=None)
credential.request_token.assert_called_once_with(SCOPE, claims=None, enable_cae=False, tenant_id=None)
assert token.token == MockCredential.NEW_TOKEN.token


Expand All @@ -85,7 +85,7 @@ async def test_cached_token_outside_refresh_window():
credential = MockCredential(cached_token=AccessToken(CACHED_TOKEN, time.time() + DEFAULT_REFRESH_OFFSET + 1))
token = await credential.get_token(SCOPE)

credential.acquire_token_silently.assert_called_once_with(SCOPE, claims=None, tenant_id=None)
credential.acquire_token_silently.assert_called_once_with(SCOPE, claims=None, enable_cae=False, tenant_id=None)
assert credential.request_token.call_count == 0
assert token.token == CACHED_TOKEN

Expand All @@ -96,8 +96,8 @@ async def test_cached_token_within_refresh_window():
credential = MockCredential(cached_token=AccessToken(CACHED_TOKEN, time.time() + DEFAULT_REFRESH_OFFSET - 1))
token = await credential.get_token(SCOPE)

credential.acquire_token_silently.assert_called_once_with(SCOPE, claims=None, tenant_id=None)
credential.request_token.assert_called_once_with(SCOPE, claims=None, tenant_id=None)
credential.acquire_token_silently.assert_called_once_with(SCOPE, claims=None, enable_cae=False, tenant_id=None)
credential.request_token.assert_called_once_with(SCOPE, claims=None, enable_cae=False, tenant_id=None)
assert token.token == MockCredential.NEW_TOKEN.token


Expand All @@ -112,5 +112,5 @@ async def test_retry_delay():
for i in range(4):
token = await credential.get_token(SCOPE)
assert token.token == CACHED_TOKEN
credential.acquire_token_silently.assert_called_with(SCOPE, claims=None, tenant_id=None)
credential.request_token.assert_called_once_with(SCOPE, claims=None, tenant_id=None)
credential.acquire_token_silently.assert_called_with(SCOPE, claims=None, enable_cae=False, tenant_id=None)
credential.request_token.assert_called_once_with(SCOPE, claims=None, enable_cae=False, tenant_id=None)
Loading