diff --git a/packages/google-api-core/tests/unit/operations_v1/test_operations_rest_client.py b/packages/google-api-core/tests/unit/operations_v1/test_operations_rest_client.py index cdb0722474d6..17d9150e5f2b 100644 --- a/packages/google-api-core/tests/unit/operations_v1/test_operations_rest_client.py +++ b/packages/google-api-core/tests/unit/operations_v1/test_operations_rest_client.py @@ -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 diff --git a/packages/google-auth/google/auth/aio/transport/sessions.py b/packages/google-auth/google/auth/aio/transport/sessions.py index d2375604b127..14e950427b5f 100644 --- a/packages/google-auth/google/auth/aio/transport/sessions.py +++ b/packages/google-auth/google/auth/aio/transport/sessions.py @@ -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 @@ -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): + 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, + ) + + 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 @@ -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 @@ -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() @@ -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"): diff --git a/packages/google-auth/tests/transport/aio/test_sessions.py b/packages/google-auth/tests/transport/aio/test_sessions.py index 69793fe49f5b..0b4ce134fa9c 100644 --- a/packages/google-auth/tests/transport/aio/test_sessions.py +++ b/packages/google-auth/tests/transport/aio/test_sessions.py @@ -13,13 +13,17 @@ # limitations under the License. import asyncio +import http.client as http_client +import threading from typing import AsyncGenerator -from unittest.mock import Mock, patch +from unittest.mock import AsyncMock, Mock, patch import pytest # type: ignore from aioresponses import aioresponses # type: ignore -from google.auth.aio.credentials import AnonymousCredentials +import google.auth.credentials +import google.auth.transport.requests +from google.auth.aio.credentials import AnonymousCredentials, Credentials from google.auth.aio.transport import ( _DEFAULT_TIMEOUT_SECONDS, DEFAULT_MAX_RETRY_ATTEMPTS, @@ -28,7 +32,12 @@ Response, sessions, ) -from google.auth.exceptions import InvalidType, TimeoutError, TransportError +from google.auth.exceptions import ( + InvalidType, + RefreshError, + TimeoutError, + TransportError, +) @pytest.fixture @@ -203,9 +212,77 @@ async def test_constructor_raises_incorrect_credentials_error(self): sessions.AsyncAuthorizedSession(credentials) exc.match( - 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`" ) + @pytest.mark.asyncio + async def test_constructor_with_sync_credentials(self): + sync_credentials = google.auth.credentials.AnonymousCredentials() + authed_session = sessions.AsyncAuthorizedSession( + sync_credentials, auth_request=MockRequest() + ) + + # Synchronous credentials are adapted to the asynchronous credentials interface. + assert isinstance(authed_session._credentials, Credentials) + assert isinstance(authed_session._credentials, sessions._SyncCredentialsAdapter) + assert authed_session._credentials._credentials is sync_credentials + with patch.object( + authed_session._credentials, + "close", + wraps=authed_session._credentials.close, + ) as mock_close: + await authed_session.close() + mock_close.assert_called_once() + + @pytest.mark.asyncio + async def test_request_with_sync_credentials_success(self, mocked_content): + sync_credentials = Mock(spec=google.auth.credentials.Credentials) + mocked_response = MockResponse( + status_code=http_client.OK, + headers={"Content-Type": "application/json"}, + content=mocked_content, + ) + auth_request = MockRequest(mocked_response) + authed_session = sessions.AsyncAuthorizedSession(sync_credentials, auth_request) + + response = await authed_session.request( + "GET", self.TEST_URL, headers={"x-test": "value"} + ) + + assert response.status_code == http_client.OK + assert await response.read() == b"Cavefish have no sight." + # The synchronous credentials are invoked with a synchronous transport rather + # than the asynchronous transport of the session. + sync_credentials.before_request.assert_called_once() + request, method, url, headers = sync_credentials.before_request.call_args.args + assert isinstance(request, google.auth.transport.requests.Request) + assert method == "GET" + assert url == self.TEST_URL + assert headers == {"x-test": "value"} + await authed_session.close() + + @pytest.mark.asyncio + async def test_request_with_sync_credentials_refreshes_on_unauthorized(self): + sync_credentials = Mock(spec=google.auth.credentials.Credentials) + unauthorized_response = MockResponse(status_code=http_client.UNAUTHORIZED) + ok_response = MockResponse(status_code=http_client.OK) + auth_request = AsyncMock(side_effect=[unauthorized_response, ok_response]) + authed_session = sessions.AsyncAuthorizedSession( + sync_credentials, auth_request=auth_request + ) + + response = await authed_session.request("GET", self.TEST_URL) + + assert response is ok_response + assert auth_request.call_count == 2 + assert unauthorized_response._close + sync_credentials.refresh.assert_called_once() + (refresh_request,) = sync_credentials.refresh.call_args.args + assert isinstance(refresh_request, google.auth.transport.requests.Request) + # The same synchronous transport is reused for every call to the credentials. + assert refresh_request is sync_credentials.before_request.call_args.args[0] + await authed_session.close() + @pytest.mark.asyncio async def test_request_default_auth_request_success(self): with aioresponses() as m: @@ -368,3 +445,179 @@ def test_mock_request_clone(): request = MockRequest() cloned = request._clone() assert cloned is request + + +class BlockingRefreshCredentials(google.auth.credentials.Credentials): + """Synchronous credentials whose refresh blocks until released by the test.""" + + def __init__(self): + super().__init__() + self.refresh_calls = 0 + self.in_flight_refreshes = 0 + self.max_in_flight_refreshes = 0 + self.refresh_started = threading.Event() + self.release_refresh = threading.Event() + + def refresh(self, request): + self.refresh_calls += 1 + self.in_flight_refreshes += 1 + self.max_in_flight_refreshes = max( + self.max_in_flight_refreshes, self.in_flight_refreshes + ) + self.refresh_started.set() + self.release_refresh.wait(timeout=5) + self.in_flight_refreshes -= 1 + self.token = "token" + + +class TestSyncCredentialsAdapter(object): + TEST_URL = "http://example.com/" + + @pytest.mark.asyncio + async def test_delegates_to_sync_credentials(self): + sync_credentials = Mock(spec=google.auth.credentials.Credentials) + sync_credentials.token = "sync-token" + sync_credentials.expiry = Mock() + sync_credentials.valid = True + sync_credentials.expired = False + adapter = sessions._SyncCredentialsAdapter(sync_credentials) + headers = {} + + await adapter.before_request(Mock(), "GET", self.TEST_URL, headers) + await adapter.refresh(Mock()) + await adapter.apply(headers, token="token") + + assert adapter.token == "sync-token" + assert adapter.expiry is sync_credentials.expiry + assert adapter.valid is True + assert adapter.expired is False + + # The same synchronous transport is used for every call. + sync_request = sync_credentials.before_request.call_args.args[0] + assert isinstance(sync_request, google.auth.transport.requests.Request) + sync_credentials.before_request.assert_called_once_with( + sync_request, "GET", self.TEST_URL, headers + ) + sync_credentials.refresh.assert_called_once_with(sync_request) + sync_credentials.apply.assert_called_once_with(headers, token="token") + with patch.object(sync_request.session, "close") as mock_close: + adapter.close() + mock_close.assert_called_once() + + @pytest.mark.asyncio + async def test_blocking_calls_run_off_the_event_loop_thread(self): + sync_credentials = Mock(spec=google.auth.credentials.Credentials) + thread_ids = [] + sync_credentials.before_request.side_effect = lambda *args: thread_ids.append( + threading.get_ident() + ) + sync_credentials.refresh.side_effect = lambda *args: thread_ids.append( + threading.get_ident() + ) + adapter = sessions._SyncCredentialsAdapter(sync_credentials) + + await adapter.before_request(Mock(), "GET", self.TEST_URL, {}) + await adapter.refresh(Mock()) + + assert len(thread_ids) == 2 + assert threading.get_ident() not in thread_ids + + @pytest.mark.asyncio + async def test_concurrent_before_request_refreshes_once(self): + sync_credentials = BlockingRefreshCredentials() + adapter = sessions._SyncCredentialsAdapter(sync_credentials) + headers = [{} for _ in range(5)] + + tasks = [ + asyncio.create_task(adapter.before_request(Mock(), "GET", self.TEST_URL, h)) + for h in headers + ] + await asyncio.to_thread(sync_credentials.refresh_started.wait, 5) + # Yield to the event loop so all tasks run until they await the refresh. + await asyncio.sleep(0) + sync_credentials.release_refresh.set() + await asyncio.gather(*tasks) + + assert sync_credentials.refresh_calls == 1 + assert all(h["authorization"] == "Bearer token" for h in headers) + + @pytest.mark.asyncio + async def test_refresh_and_before_request_share_one_refresh(self): + sync_credentials = BlockingRefreshCredentials() + adapter = sessions._SyncCredentialsAdapter(sync_credentials) + headers = {} + + refresh_task = asyncio.create_task(adapter.refresh(Mock())) + await asyncio.to_thread(sync_credentials.refresh_started.wait, 5) + before_request_task = asyncio.create_task( + adapter.before_request(Mock(), "GET", self.TEST_URL, headers) + ) + # Yield to the event loop so before_request_task reaches the refresh. + await asyncio.sleep(0) + sync_credentials.release_refresh.set() + await asyncio.gather(refresh_task, before_request_task) + + assert sync_credentials.max_in_flight_refreshes == 1 + assert sync_credentials.refresh_calls == 1 + assert headers["authorization"] == "Bearer token" + + @pytest.mark.asyncio + async def test_cancelled_caller_does_not_abandon_refresh(self): + sync_credentials = BlockingRefreshCredentials() + adapter = sessions._SyncCredentialsAdapter(sync_credentials) + + cancelled_task = asyncio.create_task( + adapter.before_request(Mock(), "GET", self.TEST_URL, {}) + ) + await asyncio.to_thread(sync_credentials.refresh_started.wait, 5) + cancelled_task.cancel() + with pytest.raises(asyncio.CancelledError): + await cancelled_task + + # The refresh started on behalf of the cancelled request is still running: + # a new request must wait for it rather than start a second refresh. + headers = {} + waiting_task = asyncio.create_task( + adapter.before_request(Mock(), "GET", self.TEST_URL, headers) + ) + await asyncio.sleep(0) + assert not waiting_task.done() + sync_credentials.release_refresh.set() + await waiting_task + + assert sync_credentials.refresh_calls == 1 + assert sync_credentials.max_in_flight_refreshes == 1 + assert headers["authorization"] == "Bearer token" + + @pytest.mark.asyncio + async def test_failed_refresh_is_not_reused(self): + sync_credentials = Mock(spec=google.auth.credentials.Credentials) + sync_credentials.refresh.side_effect = [RefreshError("refresh failed"), None] + adapter = sessions._SyncCredentialsAdapter(sync_credentials) + + with pytest.raises(RefreshError): + await adapter.refresh(Mock()) + await adapter.refresh(Mock()) + + assert sync_credentials.refresh.call_count == 2 + + @pytest.mark.asyncio + async def test_before_request_with_valid_credentials_does_not_wait_for_refresh( + self, + ): + sync_credentials = BlockingRefreshCredentials() + sync_credentials.token = "token" + adapter = sessions._SyncCredentialsAdapter(sync_credentials) + headers = {} + + refresh_task = asyncio.create_task(adapter.refresh(Mock())) + await asyncio.to_thread(sync_credentials.refresh_started.wait, 5) + # Requests that already have a valid token must not wait for the refresh. + await asyncio.wait_for( + adapter.before_request(Mock(), "GET", self.TEST_URL, headers), timeout=5 + ) + assert headers["authorization"] == "Bearer token" + + sync_credentials.release_refresh.set() + await refresh_task + assert sync_credentials.refresh_calls == 1