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 @@ -1241,24 +1241,20 @@ def test_operations_base_transport_with_adc():
)
def test_operations_auth_adc(client_class):
# If no credentials are provided, we should use ADC credentials.
is_async = "async" in str(client_class).lower()
if is_async and parse_version_to_tuple(auth_version) < (2, 60, 0):
# Older versions of google-auth do not accept the synchronous credentials
# returned by ADC in `AsyncAuthorizedSession`.
pytest.skip("ADC with the async REST transport requires google-auth >= 2.60.0")

with mock.patch.object(google.auth, "default", autospec=True) as adc:
adc.return_value = (ga_credentials.AnonymousCredentials(), None)

if "async" in str(client_class).lower():
# TODO(): Add support for adc to async REST transport.
# NOTE: Ideally, the logic for adc shouldn't be called if transport
# is set to async REST. If the user does not configure credentials
# of type `google.auth.aio.credentials.Credentials`,
# we should raise an exception to avoid the adc workflow.
with pytest.raises(google.auth.exceptions.InvalidType):
client_class()
else:
client_class()
adc.assert_called_once_with(
scopes=None,
default_scopes=(),
quota_project_id=None,
)
client_class()
adc.assert_called_once_with(
scopes=None,
default_scopes=(),
quota_project_id=None,
)


# TODO(https://github.com/googleapis/python-api-core/issues/705): Add
Expand Down
125 changes: 117 additions & 8 deletions packages/google-auth/google/auth/aio/transport/sessions.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,8 +24,9 @@
from contextlib import asynccontextmanager
from typing import TYPE_CHECKING, Mapping, Optional, Union

import google.auth.credentials
import google.auth.transport._mtls_helper
from google.auth import _exponential_backoff, exceptions
from google.auth import _exponential_backoff, _helpers, exceptions
from google.auth.aio import transport
from google.auth.aio.credentials import Credentials
from google.auth.aio.transport import mtls
Expand Down Expand Up @@ -104,6 +105,106 @@ async def with_timeout(coro):
_remaining_time()


class _SyncCredentialsAdapter(Credentials):
"""Adapts synchronous credentials to the asynchronous credentials interface.

:class:`AsyncAuthorizedSession` wraps :class:`google.auth.credentials.Credentials`
(e.g. application default credentials) with this adapter so that they can be
used with an asynchronous transport. Calls are delegated to the wrapped
credentials using a synchronous transport, and blocking calls such as
refreshing the access token run in a worker thread so that the event loop
is not blocked.

Args:
credentials (google.auth.credentials.Credentials): The synchronous
credentials to adapt.
"""

def __init__(self, credentials: google.auth.credentials.Credentials):
self._credentials = credentials
# Synchronous credentials cannot use the asynchronous transport of the
# session, so they are called with a synchronous transport instead.
self._sync_request_instance = None
# Synchronous credentials are not safe to refresh concurrently, which
# concurrent requests would otherwise do from multiple worker threads.
# Instead, at most one refresh is in flight and concurrent callers share it.
self._pending_refresh: Optional["asyncio.Task[None]"] = None

@property
def _sync_request(self):
if self._sync_request_instance is None:
# Imported here because `requests` is an optional dependency of
# google-auth. It is installed alongside `aiohttp` by the `aiohttp` extra.
from google.auth.transport import requests as sync_requests

self._sync_request_instance = sync_requests.Request()
return self._sync_request_instance

def close(self):
if (
self._sync_request_instance is not None
and hasattr(self._sync_request_instance, "session")
and self._sync_request_instance.session is not None
):
self._sync_request_instance.session.close()

@property
def token(self):
"""Optional[str]: The bearer token that can be used in HTTP headers to make
authenticated requests."""
return self._credentials.token

@property
def expiry(self):
"""Optional[datetime]: When the token expires and is no longer valid.
If this is None, the token is assumed to never expire."""
return self._credentials.expiry

@property
@_helpers.copy_docstring(google.auth.credentials.Credentials)
def valid(self):
return self._credentials.valid

@property
@_helpers.copy_docstring(google.auth.credentials.Credentials)
def expired(self):
return self._credentials.expired

async def _refresh_shared(self):
"""Refreshes the wrapped credentials, joining a refresh already in flight.

The refresh is shielded from cancellation: a caller that is cancelled
while waiting (e.g. because of a timeout) stops waiting, but the refresh
completes so that the next caller joins it rather than starting a
second, concurrent refresh.
"""
if self._pending_refresh is None or self._pending_refresh.done():
self._pending_refresh = asyncio.create_task(
asyncio.to_thread(self._credentials.refresh, self._sync_request)
)
await asyncio.shield(self._pending_refresh)

@_helpers.copy_docstring(Credentials)
async def apply(self, headers, token=None):
Comment thread
daniel-sanche marked this conversation as resolved.
self._credentials.apply(headers, token=token)

@_helpers.copy_docstring(Credentials)
async def refresh(self, request):
await self._refresh_shared()

@_helpers.copy_docstring(Credentials)
async def before_request(self, request, method, url, headers):
if not self._credentials.valid:
await self._refresh_shared()
await asyncio.to_thread(
self._credentials.before_request,
self._sync_request,
method,
url,
headers,
)
Comment thread
daniel-sanche marked this conversation as resolved.


class AsyncAuthorizedSession:
"""This is an asynchronous implementation of :class:`google.auth.requests.AuthorizedSession` class.
We utilize an instance of a class that implements :class:`google.auth.aio.transport.Request` configured
Expand All @@ -126,8 +227,9 @@ class AsyncAuthorizedSession:
credentials' headers to the request and refreshing credentials as needed.

Args:
credentials (google.auth.aio.credentials.Credentials):
The credentials to add to the request.
credentials (Union[google.auth.aio.credentials.Credentials, google.auth.credentials.Credentials]):
The credentials to add to the request. Synchronous credentials
(e.g. application default credentials) are also supported.
auth_request (Optional[google.auth.aio.transport.Request]):
An instance of a class that implements
:class:`~google.auth.aio.transport.Request` used to make requests
Expand All @@ -139,17 +241,22 @@ class AsyncAuthorizedSession:
- google.auth.exceptions.TransportError: If `auth_request` is `None`
and the external package `aiohttp` is not installed.
- google.auth.exceptions.InvalidType: If the provided credentials are
not of type `google.auth.aio.credentials.Credentials`.
not of type `google.auth.aio.credentials.Credentials` or
`google.auth.credentials.Credentials`.
"""

def __init__(
self, credentials: Credentials, auth_request: Optional[transport.Request] = None
self,
credentials: Union[Credentials, google.auth.credentials.Credentials],
auth_request: Optional[transport.Request] = None,
):
if not isinstance(credentials, Credentials):
if isinstance(credentials, google.auth.credentials.Credentials):
credentials = _SyncCredentialsAdapter(credentials)
elif not isinstance(credentials, Credentials):
raise exceptions.InvalidType(
f"The configured credentials of type {type(credentials)} are invalid and must be of type `google.auth.aio.credentials.Credentials`"
f"The configured credentials of type {type(credentials)} are invalid and must be of type `google.auth.aio.credentials.Credentials` or `google.auth.credentials.Credentials`"
)
self._credentials = credentials
self._credentials: Credentials = credentials
_auth_request = auth_request
if not _auth_request and AIOHTTP_INSTALLED:
_auth_request = AiohttpRequest()
Expand Down Expand Up @@ -846,6 +953,8 @@ async def close(self) -> None:
if inspect.isawaitable(res):
await res
finally:
if hasattr(self._credentials, "close"):
self._credentials.close()
for old_request in self._old_auth_requests:
try:
if hasattr(old_request, "close"):
Expand Down
Loading
Loading