diff --git a/packages/google-auth/google/auth/aio/transport/sessions.py b/packages/google-auth/google/auth/aio/transport/sessions.py index d88162667bda..28b331e3723e 100644 --- a/packages/google-auth/google/auth/aio/transport/sessions.py +++ b/packages/google-auth/google/auth/aio/transport/sessions.py @@ -13,8 +13,11 @@ # limitations under the License. import asyncio +import collections.abc from contextlib import asynccontextmanager import functools +import http.client as http_client +import logging import time from typing import Mapping, Optional, TYPE_CHECKING, Union import warnings @@ -37,6 +40,9 @@ except (ImportError, AttributeError): ClientTimeout = None +_LOGGER = logging.getLogger(__name__) +MTLS_URL_PREFIXES = ["mtls.googleapis.com", "mtls.sandbox.googleapis.com"] + # Tracks the internal aiohttp installation and usage try: @@ -148,6 +154,7 @@ def __init__( "`auth_request` must either be configured or the external package `aiohttp` must be installed to use the default value." ) self._auth_request = _auth_request + self._mtls_rotation_lock = asyncio.Lock() async def configure_mtls_channel(self, client_cert_callback=None): """Configure the client certificate and key for SSL connection. @@ -277,7 +284,10 @@ async def request( google.auth.exceptions.TimeoutError: If the method does not complete within the configured `max_allowed_time` or the request exceeds the configured `timeout`. + google.auth.exceptions.MutualTLSChannelError: If mutual TLS + channel reconfiguration fails for any reason during certificate rotation. """ + _auth_retry_count = kwargs.pop("_auth_retry_count", 0) if self._mtls_init_task: try: await self._mtls_init_task @@ -310,8 +320,101 @@ async def request( url, method, data, headers, actual_timeout, **kwargs ) ) + if response.status_code not in transport.DEFAULT_RETRYABLE_STATUS_CODES: break + + if response.status_code == http_client.UNAUTHORIZED: + if _auth_retry_count < 2: + is_streaming = ( + data is not None + and isinstance( + data, (collections.abc.Iterator, collections.abc.AsyncIterable) + ) + or hasattr(data, "read") + ) + if getattr(self, "is_mtls", False) and any( + prefix in url for prefix in MTLS_URL_PREFIXES + ): + # Snapshot the stale certificate state BEFORE acquiring the lock. + # This represents the cert that caused the 401 rejection. + stale_cert = self._cached_cert + + # Wait in line to acquire the lock + async with self._mtls_rotation_lock: + # Check Did another coroutine already reconfigure mTLS + if self._cached_cert != stale_cert: + # Yes! Another request already updated the channel + pass + else: + try: + ( + call_cert_bytes, + call_key_bytes, + cached_fingerprint, + current_cert_fingerprint, + ) = await mtls._run_in_executor( + google.auth.transport._mtls_helper.check_parameters_for_unauthorized_response, + self._cached_cert, + ) + except Exception as e: + _LOGGER.warning( + "Failed to check client certificate parameters: %s. Proceeding with original response.", + e, + ) + else: + if cached_fingerprint != current_cert_fingerprint: + try: + _LOGGER.info( + "Client certificate has changed, reconfiguring mTLS " + "channel." + ) + if ( + self._mtls_init_task + and self._mtls_init_task.done() + ): + self._mtls_init_task = None + await self.configure_mtls_channel( + lambda: (call_cert_bytes, call_key_bytes) + ) + except Exception as e: + _LOGGER.error( + "Failed to reconfigure mTLS channel: %s", e + ) + raise exceptions.MutualTLSChannelError( + "Failed to reconfigure mTLS channel" + ) from e + else: + _LOGGER.info( + "Skipping reconfiguration of mTLS channel because the client" + " certificate has not changed." + ) + if is_streaming: + return response + if hasattr(response, "close"): + if asyncio.iscoroutinefunction(response.close): + await response.close() + else: + response.close() + try: + await self._credentials.refresh(self._auth_request) + except exceptions.RefreshError as e: + _LOGGER.debug( + "Credential refresh failed, returning 401 response. Error: %s", + e, + ) + return response + kwargs["_auth_retry_count"] = _auth_retry_count + 1 + return await self.request( + method, + url, + data=data, + headers=headers, + max_allowed_time=max_allowed_time, + timeout=timeout, + total_attempts=total_attempts, + **kwargs, + ) return response @functools.wraps(request) diff --git a/packages/google-auth/tests/transport/aio/test_sessions_mtls.py b/packages/google-auth/tests/transport/aio/test_sessions_mtls.py index b68766ca5b5d..3e5ff3fc3f86 100644 --- a/packages/google-auth/tests/transport/aio/test_sessions_mtls.py +++ b/packages/google-auth/tests/transport/aio/test_sessions_mtls.py @@ -12,6 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. +import http.client as http_client import json import os import ssl @@ -344,3 +345,184 @@ async def test_configure_mtls_channel_close_exception_does_not_abort(self): assert session._is_mtls is True assert session._cached_cert == b"fake_cert_data" await session.close() + + @pytest.mark.asyncio + async def test_cert_rotation_failure_raises_error(self): + mock_creds = mock.AsyncMock(spec=credentials.Credentials) + mock_creds.before_request = mock.AsyncMock(return_value=None) + + mock_resp = mock.Mock() + mock_resp.status_code = http_client.UNAUTHORIZED + mock_auth_req = mock.AsyncMock(return_value=mock_resp) + + session = sessions.AsyncAuthorizedSession( + mock_creds, auth_request=mock_auth_req + ) + session._is_mtls = True + session._cached_cert = b"old_cert" + + new_cert = b"new_cert" + new_key = b"new_key" + + with mock.patch( + "google.auth.transport._mtls_helper.check_parameters_for_unauthorized_response" + ) as mock_check, mock.patch.object( + session, "configure_mtls_channel", new_callable=mock.AsyncMock + ) as mock_conf: + mock_check.return_value = (new_cert, new_key, b"old_fp", b"new_fp") + mock_conf.side_effect = Exception("Failed to reconfigure") + + with pytest.raises(exceptions.MutualTLSChannelError): + await session.request("GET", "https://pubsub.mtls.googleapis.com/test") + + mock_check.assert_called_once() + mock_conf.assert_called_once() + + await session.close() + + @pytest.mark.asyncio + async def test_cert_rotation_check_params_fails(self): + mock_creds = mock.AsyncMock(spec=credentials.Credentials) + mock_creds.before_request = mock.AsyncMock(return_value=None) + + mock_resp = mock.Mock() + mock_resp.status_code = http_client.UNAUTHORIZED + mock_auth_req = mock.AsyncMock(return_value=mock_resp) + + session = sessions.AsyncAuthorizedSession( + mock_creds, auth_request=mock_auth_req + ) + session._is_mtls = True + session._cached_cert = b"old_cert" + + with mock.patch( + "google.auth.transport._mtls_helper.check_parameters_for_unauthorized_response" + ) as mock_check, mock.patch.object( + session, "configure_mtls_channel", new_callable=mock.AsyncMock + ) as mock_conf: + mock_check.side_effect = Exception("Failed to check params") + + resp = await session.request( + "GET", "https://pubsub.mtls.googleapis.com/test" + ) + + assert resp == mock_resp + assert mock_check.call_count >= 1 + mock_conf.assert_not_called() + + await session.close() + + @pytest.mark.asyncio + async def test_no_cert_rotation_when_cert_match_and_mTLS_enabled(self): + mock_creds = mock.AsyncMock(spec=credentials.Credentials) + mock_creds.before_request = mock.AsyncMock(return_value=None) + + mock_resp = mock.Mock() + mock_resp.status_code = http_client.UNAUTHORIZED + mock_auth_req = mock.AsyncMock(return_value=mock_resp) + + session = sessions.AsyncAuthorizedSession( + mock_creds, auth_request=mock_auth_req + ) + session._is_mtls = True + session._cached_cert = b"old_cert" + + new_cert = b"new_cert" + new_key = b"new_key" + + with mock.patch( + "google.auth.transport._mtls_helper.check_parameters_for_unauthorized_response" + ) as mock_check, mock.patch.object( + session, "configure_mtls_channel", new_callable=mock.AsyncMock + ) as mock_conf: + # Matching fingerprints mean no layout rotation is needed + mock_check.return_value = (new_cert, new_key, b"old_fp", b"old_fp") + + resp = await session.request( + "GET", "https://pubsub.mtls.googleapis.com/test" + ) + + assert resp == mock_resp + assert mock_check.call_count >= 1 + mock_conf.assert_not_called() + + await session.close() + + @pytest.mark.asyncio + async def test_cert_rotation_success_and_retry(self): + mock_creds = mock.AsyncMock(spec=credentials.Credentials) + mock_creds.before_request = mock.AsyncMock(return_value=None) + mock_creds.refresh = mock.AsyncMock(return_value=None) + + # Initial request fails natively with 401. Retry succeeds with 200. + mock_resp_401 = mock.Mock() + mock_resp_401.status_code = http_client.UNAUTHORIZED + mock_resp_200 = mock.Mock() + mock_resp_200.status_code = http_client.OK + + # Use side_effect to dynamically yield responses + mock_auth_req = mock.AsyncMock(side_effect=[mock_resp_401, mock_resp_200]) + + session = sessions.AsyncAuthorizedSession( + mock_creds, auth_request=mock_auth_req + ) + session._is_mtls = True + session._cached_cert = b"old_cert" + + new_cert = b"new_cert" + new_key = b"new_key" + + with mock.patch( + "google.auth.transport._mtls_helper.check_parameters_for_unauthorized_response" + ) as mock_check, mock.patch.object( + session, "configure_mtls_channel", new_callable=mock.AsyncMock + ) as mock_conf: + mock_check.return_value = (new_cert, new_key, b"old_fp", b"new_fp") + + resp = await session.request( + "GET", "https://pubsub.mtls.googleapis.com/test" + ) + + # 1. Assert the retried 200 response is successfully returned to the user + assert resp == mock_resp_200 + + # 2. Assert rotation logic correctly executed + mock_check.assert_called_once() + mock_conf.assert_called_once_with(mock.ANY) + + # 3. Assert credentials were explicitly refreshed + mock_creds.refresh.assert_called_once() + + # 4. Assert headers were explicitly rebound on the recursive retry (2 invocations) + assert mock_creds.before_request.call_count == 2 + + await session.close() + + @pytest.mark.asyncio + async def test_non_mtls_url_bypasses_rotation(self): + mock_creds = mock.AsyncMock(spec=credentials.Credentials) + mock_resp_401 = mock.Mock() + mock_resp_401.status_code = http_client.UNAUTHORIZED + mock_auth_req = mock.AsyncMock(return_value=mock_resp_401) + + session = sessions.AsyncAuthorizedSession( + mock_creds, auth_request=mock_auth_req + ) + + # Even if mTLS is enabled globally... + session._is_mtls = True + session._cached_cert = b"old_cert" + + with mock.patch( + "google.auth.transport._mtls_helper.check_parameters_for_unauthorized_response" + ) as mock_check, mock.patch.object( + session, "configure_mtls_channel", new_callable=mock.AsyncMock + ) as mock_conf: + # ...a 401 on a regular domain bypasses checks and just returns the 401 locally + resp = await session.request("GET", "https://pubsub.googleapis.com/test") + + assert resp == mock_resp_401 + mock_check.assert_not_called() + mock_conf.assert_not_called() + + await session.close()