diff --git a/packages/google-auth/google/auth/transport/_mtls_helper.py b/packages/google-auth/google/auth/transport/_mtls_helper.py index 7779c484c713..a469aa4e4ba3 100644 --- a/packages/google-auth/google/auth/transport/_mtls_helper.py +++ b/packages/google-auth/google/auth/transport/_mtls_helper.py @@ -817,10 +817,11 @@ def check_parameters_for_unauthorized_response(cached_cert): Returns: bytes: The client callback cert bytes. bytes: The client callback key bytes. + Optional[Union[bytes, str]]: The passphrase for the key. str: The base64-encoded SHA256 cached fingerprint. str: The base64-encoded SHA256 current cert fingerprint. """ - call_cert_bytes, call_key_bytes = call_client_cert_callback() + call_cert_bytes, call_key_bytes, passphrase = call_client_cert_callback() cert_obj = _agent_identity_utils.parse_certificate(call_cert_bytes) current_cert_fingerprint = _agent_identity_utils.calculate_certificate_fingerprint( cert_obj @@ -831,12 +832,18 @@ def check_parameters_for_unauthorized_response(cached_cert): ) else: cached_fingerprint = current_cert_fingerprint - return call_cert_bytes, call_key_bytes, cached_fingerprint, current_cert_fingerprint + return ( + call_cert_bytes, + call_key_bytes, + passphrase, + cached_fingerprint, + current_cert_fingerprint, + ) def call_client_cert_callback(): - """Calls the client cert callback and returns the certificate and key.""" + """Calls the client cert callback and returns the certificate, key, and passphrase.""" _, cert_bytes, key_bytes, passphrase = get_client_ssl_credentials( generate_encrypted_key=True ) - return cert_bytes, key_bytes + return cert_bytes, key_bytes, passphrase diff --git a/packages/google-auth/google/auth/transport/grpc.py b/packages/google-auth/google/auth/transport/grpc.py index df6e5fa82882..a856e82ad2ac 100644 --- a/packages/google-auth/google/auth/transport/grpc.py +++ b/packages/google-auth/google/auth/transport/grpc.py @@ -16,14 +16,21 @@ from __future__ import absolute_import +import functools import logging import warnings + from google.auth import exceptions +from google.auth import transport from google.auth.transport import _mtls_helper +from google.auth.transport import mtls_interceptor‎ from google.auth.transport import mtls from google.oauth2 import service_account +from typing import Optional + + try: import grpc # type: ignore except ImportError as caught_exc: # pragma: NO COVER @@ -283,6 +290,7 @@ def my_client_cert_callback(): ) # If SSL credentials are not explicitly set, try client_cert_callback and ADC. + cached_cert: Optional[bytes] = None if not ssl_credentials: use_client_cert = _mtls_helper.check_use_client_cert() if use_client_cert and client_cert_callback: @@ -291,10 +299,12 @@ def my_client_cert_callback(): ssl_credentials = grpc.ssl_channel_credentials( certificate_chain=cert, private_key=key ) + cached_cert = cert elif use_client_cert: # Use application default SSL credentials. - adc_ssl_credentils = SslCredentials() - ssl_credentials = adc_ssl_credentils.ssl_credentials + adc_ssl_credentials = SslCredentials() + ssl_credentials = adc_ssl_credentials.ssl_credentials + cached_cert = adc_ssl_credentials._cached_cert else: ssl_credentials = grpc.ssl_channel_credentials() @@ -302,8 +312,23 @@ def my_client_cert_callback(): composite_credentials = grpc.composite_channel_credentials( ssl_credentials, google_auth_credentials ) - - return grpc.secure_channel(target, composite_credentials, **kwargs) + is_recreation = kwargs.pop("_is_recreation", False) + channel = grpc.secure_channel(target, composite_credentials, **kwargs) + # Avoid wrapping if mTLS is disabled or if this is a channel recreation call + if cached_cert and not is_recreation: + # Package arguments so the channel can be recreated later + create_channel_fn = functools.partial( + secure_authorized_channel, + credentials=credentials, + request=request, + target=target, + _is_recreation=True, # Hidden flag to stop recursion + **kwargs + ) + wrapper = mtls_interceptor.MTLSRefreshingChannel(target, create_channel_fn, channel, cached_cert) + interceptor = mtls_interceptor.CertRotationInterceptor(wrapper=wrapper) + return grpc.intercept_channel(wrapper, interceptor) + return channel class SslCredentials: @@ -327,6 +352,7 @@ class SslCredentials: def __init__(self): use_client_cert = _mtls_helper.check_use_client_cert() + self._cached_cert = None if not use_client_cert: self._is_mtls = False else: @@ -355,6 +381,7 @@ def ssl_credentials(self): self._ssl_credentials = grpc.ssl_channel_credentials( certificate_chain=cert, private_key=key ) + self._cached_cert = cert else: self._ssl_credentials = grpc.ssl_channel_credentials() self._is_mtls = False diff --git a/packages/google-auth/google/auth/transport/mtls_interceptor.py b/packages/google-auth/google/auth/transport/mtls_interceptor.py new file mode 100644 index 000000000000..13e482c5396d --- /dev/null +++ b/packages/google-auth/google/auth/transport/mtls_interceptor.py @@ -0,0 +1,410 @@ +"""mTLS Interceptor and Channel Wrapper for certificate rotation.""" + +import collections +import threading +import time + +import grpc +from google.auth import transport + +_ClientCallDetails = collections.namedtuple( + "_ClientCallDetails", + ("method", "timeout", "metadata", "credentials", "wait_for_ready"), +) + +class _DeadlineExceededError(grpc.RpcError, grpc.Call): + def __init__(self, details): + super().__init__() + self._details = details + + def code(self): + return grpc.StatusCode.DEADLINE_EXCEEDED + + def details(self): + return self._details + + +class _BaseCallWrapper(grpc.Future, grpc.Call): + """A generic wrapper that delegates standard grpc.Call and grpc.Future + methods to an underlying call object. + """ + + def cancel(self): + return self._call.cancel() + + def cancelled(self): + return self._call.cancelled() + + def running(self): + return self._call.running() + + def done(self): + return self._call.done() + + def result(self, timeout=None): + return self._call.result(timeout=timeout) + + def exception(self, timeout=None): + return self._call.exception(timeout=timeout) + + def traceback(self, timeout=None): + return self._call.traceback(timeout=timeout) + + def add_done_callback(self, fn): + self._call.add_done_callback(fn) + + def initial_metadata(self): + return self._call.initial_metadata() + + def trailing_metadata(self): + return self._call.trailing_metadata() + + def code(self): + return self._call.code() + + def details(self): + return self._call.details() + + +class _RetryableUnaryResponseFuture(_BaseCallWrapper): + def __init__(self, continuation, client_call_details, request_or_iterator, interceptor): + self._continuation = continuation + self._client_call_details = client_call_details + self._request_or_iterator = request_or_iterator + self._interceptor = interceptor + self._retry_count = 0 + self._call = None + self._lock = threading.RLock() + + timeout = getattr(self._client_call_details, "timeout", None) + if timeout: + self._initial_timeout = timeout + self._start_time = time.monotonic() + else: + self._initial_timeout = None + self._start_time = None + + self._terminal_exception = None + self._callbacks = [] + self._is_completed = False + self._start_call() + + def _start_call(self): + with self._lock: + payload = self._request_or_iterator + call_details = self._client_call_details + if callable(payload) and hasattr(payload, "can_replay"): + if not payload.can_replay(): + call_details = _ClientCallDetails( + method=call_details.method, + timeout=call_details.timeout, + metadata=call_details.metadata, + credentials=call_details.credentials, + wait_for_ready=call_details.wait_for_ready, + ) + + if self._initial_timeout: + elapsed = time.monotonic() - self._start_time + remaining = self._initial_timeout - elapsed + if remaining <= 0: + raise _DeadlineExceededError( + "Deadline Exceeded during retry resolution." + ) + call_details = _ClientCallDetails( + method=call_details.method, + timeout=remaining, + metadata=call_details.metadata, + credentials=call_details.credentials, + wait_for_ready=call_details.wait_for_ready, + ) + + self._call = self._continuation(call_details, payload) + self._call.add_done_callback(self._on_inner_future_done) + + def _on_inner_future_done(self, inner_future): + with self._lock: + if self._call is not inner_future: + return + if inner_future.cancelled(): + self._resolve_completion() + return + + status_code = inner_future.code() + + if self._interceptor._wrapper: + ( + chk_should_retry, + chk_cert, + chk_key, + chk_pwd, + ) = self._interceptor._should_retry( + status_code, 0, self._interceptor._wrapper._cached_cert + ) + + if chk_should_retry: + payload = self._request_or_iterator + can_replay = True + if callable(payload) and hasattr(payload, "can_replay"): + can_replay = payload.can_replay() + + if can_replay: + try: + self._interceptor._wrapper.refresh_logic( + 1, chk_cert, chk_key, chk_pwd + ) + except Exception as e: + with self._lock: + self._terminal_exception = e + if self._terminal_exception is None: + self._retry_count += 1 + self._start_call() + return + + self._resolve_completion() + + def _resolve_completion(self): + callbacks = [] + with self._lock: + self._is_completed = True + callbacks = self._callbacks[:] + for cb in callbacks: + try: + cb(self) + except Exception: + pass + + +class _RetryableStreamResponseIterator(_BaseCallWrapper): + def __init__(self, continuation, client_call_details, request_or_iterator, interceptor): + self._continuation = continuation + self._client_call_details = client_call_details + self._request_or_iterator = ( + _ReplayableIterator(request_or_iterator) + if hasattr(request_or_iterator, "__iter__") + else request_or_iterator + ) + self._call = None + self._retry_count = 0 + self._yielded_any_response = False + self._lock = threading.RLock() + self._interceptor = interceptor + self._start_call() + + def _start_call(self): + with self._lock: + if isinstance(self._request_or_iterator, _ReplayableIterator): + payload = self._request_or_iterator.reader() + else: + payload = self._request_or_iterator + self._call = self._continuation(self._client_call_details, payload) + + def __iter__(self): + return self + + def __next__(self): + while True: + try: + response = next(self._call) + self._yielded_any_response = True + return response + except grpc.RpcError as rpc_error: + status_code = rpc_error.code() + can_replay_request = True + if isinstance(self._request_or_iterator, _ReplayableIterator): + can_replay_request = self._request_or_iterator.can_replay() + + if not self._yielded_any_response and can_replay_request: + if self._interceptor._wrapper: + ( + chk_should_retry, + chk_cert, + chk_key, + chk_pwd, + ) = self._interceptor._should_retry( + status_code, 0, self._interceptor._wrapper._cached_cert + ) + if chk_should_retry: + try: + self._interceptor._wrapper.refresh_logic( + 1, chk_cert, chk_key, chk_pwd + ) + except Exception as e: + raise e + self._retry_count += 1 + self._start_call() + continue + raise rpc_error + + def next(self): + return self.__next__() + + +class _ReplayableIterator(object): + def __init__(self, target_iterator, max_items=1000): + self._target_iterator = iter(target_iterator) + self._max_items = max_items + self._buffer = [] + self._exhausted = False + self._can_replay = True + self._lock = threading.Lock() + + def _get_item(self, index): + with self._lock: + if not self._can_replay: + raise RuntimeError("Iterator replay capability lost") + + while index >= len(self._buffer) and not self._exhausted: + try: + item = next(self._target_iterator) + if len(self._buffer) >= self._max_items: + self._can_replay = False + self._buffer = None + raise RuntimeError( + f"More than {self._max_items} items in replay buffer." + ) + self._buffer.append(item) + except StopIteration: + self._exhausted = True + + if index < len(self._buffer): + return self._buffer[index] + raise StopIteration() + + def reader(self): + return _ReplayableIteratorReader(self) + + def can_replay(self): + with self._lock: + return self._can_replay + + +class _ReplayableIteratorReader(object): + def __init__(self, parent): + self._parent = parent + self._read_index = 0 + + def __next__(self): + item = self._parent._get_item(self._read_index) + self._read_index += 1 + return item + + def next(self): + return self.__next__() + + +class CertRotationInterceptor( + grpc.UnaryUnaryClientInterceptor, + grpc.UnaryStreamClientInterceptor, + grpc.StreamUnaryClientInterceptor, + grpc.StreamStreamClientInterceptor, +): + """A gRPC client interceptor that provides automatic retry logic for mTLS certificate rotation.""" + + def __init__(self, wrapper=None): + self._wrapper = wrapper + self._max_retries = transport.DEFAULT_MAX_REFRESH_ATTEMPTS + + def _should_retry(self, code, retry_count, attempt_cert): + do_refresh = False + new_cert, new_key, passphrase = None, None, None + + if retry_count < self._max_retries and code == grpc.StatusCode.UNAUTHENTICATED: + ( + new_cert, + new_key, + passphrase, + ) = self._wrapper.get_cert() + + if new_cert and new_cert != attempt_cert: + do_refresh = True + + return do_refresh, new_cert, new_key, passphrase + + def intercept_unary_unary(self, continuation, client_call_details, request): + return _RetryableUnaryResponseFuture( + continuation, client_call_details, request, self + ) + + def intercept_unary_stream(self, continuation, client_call_details, request): + return _RetryableStreamResponseIterator( + continuation, client_call_details, request, self + ) + + def intercept_stream_unary(self, continuation, client_call_details, request_iterator): + return _RetryableUnaryResponseFuture( + continuation, client_call_details, request_iterator, self + ) + + def intercept_stream_stream(self, continuation, client_call_details, request_iterator): + return _RetryableStreamResponseIterator( + continuation, client_call_details, request_iterator, self + ) + + +class MTLSRefreshingChannel(grpc.Channel): + def __init__(self, target, create_channel_fn, initial_channel, initial_cert): + self._target = target + self._create_channel_fn = create_channel_fn + self._channel = initial_channel + self._cached_cert = initial_cert + self._lock = threading.Lock() + self._subscribers = [] + self._cert_factory = transport.grpc._get_client_ssl_credentials_auto_enablement + + def get_cert(self): + creds = self._cert_factory() + return creds.certificate_chain, creds.private_key, None + + def refresh_logic(self, expected_retry_count, call_cert_bytes, call_key_bytes, passphrase): + with self._lock: + # Another thread may have already completed the refresh + if self._cached_cert != call_cert_bytes: + return + + new_ssl_credentials = grpc.ssl_channel_credentials( + certificate_chain=call_cert_bytes, + private_key=call_key_bytes, + ) + + self._channel = self._create_channel_fn( + ssl_credentials=new_ssl_credentials, + client_cert_callback=None + ) + + self._cached_cert = call_cert_bytes + for callback in self._subscribers: + callback(grpc.ChannelConnectivity.IDLE) + + def subscribe(self, callback, try_to_connect=False): + with self._lock: + self._subscribers.append(callback) + return self._channel.subscribe(callback, try_to_connect=try_to_connect) + + def unsubscribe(self, callback): + with self._lock: + if callback in self._subscribers: + self._subscribers.remove(callback) + return self._channel.unsubscribe(callback) + + def unary_unary(self, method, *args, **kwargs): + return lambda request, **req_kwargs: self._channel.unary_unary( + method, *args, **kwargs + )(request, **req_kwargs) + + def unary_stream(self, method, *args, **kwargs): + return lambda request, **req_kwargs: self._channel.unary_stream( + method, *args, **kwargs + )(request, **req_kwargs) + + def stream_unary(self, method, *args, **kwargs): + return lambda request_iterator, **req_kwargs: self._channel.stream_unary( + method, *args, **kwargs + )(request_iterator, **req_kwargs) + + def stream_stream(self, method, *args, **kwargs): + return lambda request_iterator, **req_kwargs: self._channel.stream_stream( + method, *args, **kwargs + )(request_iterator, **req_kwargs) + + def close(self): + self._channel.close() diff --git a/packages/google-auth/google/auth/transport/requests.py b/packages/google-auth/google/auth/transport/requests.py index 822cf687f5d0..3aaa2e02a516 100644 --- a/packages/google-auth/google/auth/transport/requests.py +++ b/packages/google-auth/google/auth/transport/requests.py @@ -658,6 +658,7 @@ def request( ( call_cert_bytes, call_key_bytes, + _, # passphrase is not processed by requests adapter cached_fingerprint, current_cert_fingerprint, ) = _mtls_helper.check_parameters_for_unauthorized_response( diff --git a/packages/google-auth/google/auth/transport/urllib3.py b/packages/google-auth/google/auth/transport/urllib3.py index 18e6128e03bd..eacad22b5642 100644 --- a/packages/google-auth/google/auth/transport/urllib3.py +++ b/packages/google-auth/google/auth/transport/urllib3.py @@ -440,6 +440,7 @@ def urlopen(self, method, url, body=None, headers=None, **kwargs): ( call_cert_bytes, call_key_bytes, + _, cached_fingerprint, current_cert_fingerprint, ) = _mtls_helper.check_parameters_for_unauthorized_response( diff --git a/packages/google-auth/tests/transport/test__mtls_helper.py b/packages/google-auth/tests/transport/test__mtls_helper.py index e9bb62db2133..edfededf2df9 100644 --- a/packages/google-auth/tests/transport/test__mtls_helper.py +++ b/packages/google-auth/tests/transport/test__mtls_helper.py @@ -1185,6 +1185,7 @@ def test_check_parameters_for_unauthorized_response_with_cached_cert( mock_call_client_cert_callback.return_value = ( CERT_MOCK_VAL, KEY_MOCK_VAL, + b"passphrase", ) mock_agent_identity_utils.get_cached_cert_fingerprint.return_value = ( "cached_fingerprint" @@ -1196,6 +1197,7 @@ def test_check_parameters_for_unauthorized_response_with_cached_cert( ( cert, key, + passphrase, cached_fingerprint, current_fingerprint, ) = _mtls_helper.check_parameters_for_unauthorized_response( @@ -1204,6 +1206,7 @@ def test_check_parameters_for_unauthorized_response_with_cached_cert( assert cert == CERT_MOCK_VAL assert key == KEY_MOCK_VAL + assert passphrase == b"passphrase" assert cached_fingerprint == "cached_fingerprint" assert current_fingerprint == "current_fingerprint" mock_call_client_cert_callback.assert_called_once() @@ -1219,6 +1222,7 @@ def test_check_parameters_for_unauthorized_response_without_cached_cert( mock_call_client_cert_callback.return_value = ( CERT_MOCK_VAL, KEY_MOCK_VAL, + b"passphrase", ) mock_agent_identity_utils.calculate_certificate_fingerprint.return_value = ( "current_fingerprint" @@ -1227,12 +1231,14 @@ def test_check_parameters_for_unauthorized_response_without_cached_cert( ( cert, key, + passphrase, cached_fingerprint, current_fingerprint, ) = _mtls_helper.check_parameters_for_unauthorized_response(cached_cert=None) assert cert == CERT_MOCK_VAL assert key == KEY_MOCK_VAL + assert passphrase == b"passphrase" assert cached_fingerprint == "current_fingerprint" assert current_fingerprint == "current_fingerprint" mock_call_client_cert_callback.assert_called_once() @@ -1247,10 +1253,11 @@ def test_call_client_cert_callback(self, mock_get_client_ssl_credentials): b"passphrase", ) - cert, key = _mtls_helper.call_client_cert_callback() + cert, key, passphrase = _mtls_helper.call_client_cert_callback() assert cert == b"cert_bytes" assert key == b"key_bytes" + assert passphrase == b"passphrase" mock_get_client_ssl_credentials.assert_called_once_with( generate_encrypted_key=True ) diff --git a/packages/google-auth/tests/transport/test_grpc.py b/packages/google-auth/tests/transport/test_grpc.py index 7979df7abb4d..108839717412 100644 --- a/packages/google-auth/tests/transport/test_grpc.py +++ b/packages/google-auth/tests/transport/test_grpc.py @@ -19,6 +19,7 @@ from unittest import mock import warnings +import grpc import pytest # type: ignore from google.auth import _helpers @@ -28,9 +29,18 @@ from google.auth import transport from google.oauth2 import service_account + +def unwrap(ch): + if isinstance(ch, mock.Mock) or isinstance(ch, mock.MagicMock): + return ch + if hasattr(ch, "_channel"): + return unwrap(ch._channel) + return ch + + try: # pylint: disable=ungrouped-imports - import grpc # type: ignore + import google.auth.transport.grpc HAS_GRPC = True @@ -229,7 +239,7 @@ def test_secure_authorized_channel_adc( composite_channel_credentials.return_value, options=mock.sentinel.options, ) - assert channel == secure_channel.return_value + assert unwrap(channel) == secure_channel.return_value @mock.patch("google.auth.transport.grpc.SslCredentials", autospec=True) def test_secure_authorized_channel_adc_without_client_cert_env( @@ -275,7 +285,7 @@ def test_secure_authorized_channel_adc_without_client_cert_env( composite_channel_credentials.return_value, options=mock.sentinel.options, ) - assert channel == secure_channel.return_value + assert unwrap(channel) == secure_channel.return_value def test_secure_authorized_channel_explicit_ssl( self, @@ -682,6 +692,618 @@ def test_get_client_ssl_credentials_auto_enablement( ) +@mock.patch("google.auth.transport.grpc._ReplayableIterator") +def test_interceptor_uses_factory_if_callable(mock_replayable): + import google.auth.transport.grpc as transport_grpc + + interceptor = transport_grpc.CertRotationInterceptor() + + call_no_factory = transport_grpc._RetryableStreamResponseIterator( + continuation=mock.Mock(), + client_call_details=mock.Mock(), + request_or_iterator=[b"1", b"2"], + interceptor=interceptor, + is_client_stream=True, + ) + assert call_no_factory._uses_factory is False + assert call_no_factory._payload is not None + + def generator_factory(): + return (x for x in [b"1", b"2"]) + + call_factory = transport_grpc._RetryableStreamResponseIterator( + continuation=mock.Mock(), + client_call_details=mock.Mock(), + request_or_iterator=generator_factory, + interceptor=interceptor, + is_client_stream=True, + ) + assert call_factory._uses_factory is True + assert call_factory._payload is None + + +@mock.patch("google.auth.transport.grpc.CertRotationInterceptor._should_retry") +def test_factory_infinite_replay_on_error(mock_should_retry): + import google.auth.transport.grpc as transport_grpc + + interceptor = transport_grpc.CertRotationInterceptor() + interceptor._wrapper = mock.Mock() + interceptor._wrapper._cached_cert = "cert" + mock_should_retry.side_effect = [ + (True, b"cert", b"key", None), + (False, None, None, None), + ] + + mock_inner_call1 = mock.Mock() + mock_err = transport_grpc.grpc.RpcError() + mock_err.code = lambda: transport_grpc.grpc.StatusCode.UNAUTHENTICATED + mock_inner_call1.__next__ = mock.Mock(side_effect=mock_err) + + mock_inner_call2 = mock.Mock() + mock_inner_call2.__next__ = mock.Mock(side_effect=[b"SUCCESS", StopIteration]) + + continuation = mock.Mock(side_effect=[mock_inner_call1, mock_inner_call2]) + + factory_calls = 0 + + def factory(): + nonlocal factory_calls + factory_calls += 1 + return (x for x in [b"A"]) + + stream = transport_grpc._RetryableStreamResponseIterator( + continuation=continuation, + client_call_details=mock.Mock(), + request_or_iterator=factory, + interceptor=interceptor, + is_client_stream=True, + ) + + responses = list(stream) + assert responses == [b"SUCCESS"] + assert factory_calls == 2 + + +@mock.patch("google.auth.transport._mtls_helper.decrypt_private_key") +@mock.patch( + "google.auth.transport._mtls_helper.check_parameters_for_unauthorized_response" +) +@mock.patch("google.auth.transport.grpc.secure_authorized_channel") +def test_refresh_logic_closes_old_channel( + mock_secure_channel, mock_check_params, mock_decrypt +): + import google.auth.transport.grpc as transport_grpc + + mock_check_params.return_value = ("cert", "cert", "passphrase", "old_fp", "new_fp") + mock_decrypt.return_value = b"decrypted_key" + old_channel = mock.Mock() + new_channel = mock.Mock() + mock_secure_channel.return_value = new_channel + + subscriber = mock.Mock() + + refreshing_channel = transport_grpc.MTLSRefreshingChannel( + target="example.com:443", + factory_args={}, + initial_channel=old_channel, + initial_cert="cert", + ) + refreshing_channel.subscribe(subscriber) + + refreshing_channel.refresh_logic( + 1, call_cert_bytes=b"newcert", call_key_bytes=b"newkey", passphrase=None + ) + + old_channel.unsubscribe.assert_called_once_with(subscriber) + new_channel.subscribe.assert_called_once_with(subscriber) + # old_channel.close.assert_called_once() # Removed in PR 18019 + + +@mock.patch("google.auth.transport.grpc.CertRotationInterceptor._should_retry") +def test_unary_response_future_deadline_exceeded_on_retry(mock_should_retry): + import google.auth.transport.grpc as transport_grpc + + interceptor = transport_grpc.CertRotationInterceptor() + interceptor._wrapper = mock.Mock() + interceptor._wrapper._cached_cert = "cert" + mock_should_retry.return_value = (True, b"cert", b"key", None) + + mock_err = transport_grpc.grpc.RpcError() + mock_err.code = lambda: transport_grpc.grpc.StatusCode.UNAUTHENTICATED + + inner_future = mock.Mock() + inner_future.exception = lambda: mock_err + inner_future.result = mock.Mock(side_effect=mock_err) + + callbacks_fired = [] + + def callback(f): + callbacks_fired.append(f) + + call_details = mock.Mock() + call_details.timeout = 0.001 # very short timeout + + # Simulating initial call + future = transport_grpc._RetryableUnaryResponseFuture( + continuation=lambda cd, pl: inner_future, + client_call_details=call_details, + request_or_iterator=b"request", + interceptor=interceptor, + is_client_stream=False, + ) + future._completion_event.set() + future.add_done_callback(callback) + + # Allow time to elapse so remaining timeout <= 0 + time.sleep(0.01) + + # Trigger inner future completion + future._on_inner_future_done(inner_future) + + # Verify future is marked done and does not hang + assert future.done() is True + assert len(callbacks_fired) == 1 + with pytest.raises(transport_grpc.grpc.RpcError): + future.result(timeout=1) + + +@mock.patch("google.auth.transport.grpc.CertRotationInterceptor._should_retry") +def test_unary_response_future_cancelled(mock_should_retry): + import google.auth.transport.grpc as transport_grpc + + interceptor = transport_grpc.CertRotationInterceptor() + interceptor._wrapper = mock.Mock() + interceptor._wrapper._cached_cert = "cert" + + # Mock an incoming cancelled inner_future + inner_future = mock.Mock() + inner_future.cancelled.return_value = True + + callbacks_fired = [] + + def callback(f): + callbacks_fired.append(f) + + # Throw an exception inside the callback execution to cover the newly added except branch + def failing_callback(f): + raise Exception("Deliberate failure to test exception catching") + + call_details = mock.Mock() + future = transport_grpc._RetryableUnaryResponseFuture( + continuation=lambda cd, pl: inner_future, + client_call_details=call_details, + request_or_iterator=b"request", + interceptor=interceptor, + is_client_stream=False, + ) + future._completion_event.set() + future.add_done_callback(callback) + future.add_done_callback(failing_callback) + + future._on_inner_future_done(inner_future) + + assert future.done() is True + assert len(callbacks_fired) == 1 + + +@mock.patch("google.auth.transport.grpc.CertRotationInterceptor._should_retry") +def test_unary_response_future_rpc_error_retry_start_call_exception(mock_should_retry): + import google.auth.transport.grpc as transport_grpc + + interceptor = transport_grpc.CertRotationInterceptor() + interceptor._wrapper = mock.Mock() + interceptor._wrapper._cached_cert = "cert" + + mock_err = transport_grpc.grpc.RpcError() + mock_err.code = lambda: transport_grpc.grpc.StatusCode.UNAUTHENTICATED + mock_should_retry.return_value = (True, b"cert", b"key", None) + + inner_future = mock.Mock() + inner_future.cancelled.return_value = False + inner_future.exception.return_value = mock_err + + call_details = mock.Mock() + + future = transport_grpc._RetryableUnaryResponseFuture( + continuation=lambda cd, pl: inner_future, + client_call_details=call_details, + request_or_iterator=b"request", + interceptor=interceptor, + is_client_stream=False, + ) + future._completion_event.set() + + with mock.patch.object(future, "_start_call", side_effect=mock_err): + future._on_inner_future_done(inner_future) + + assert interceptor._wrapper.refresh_logic.call_count == 2 + + +def test_stream_response_iterator_done(): + import google.auth.transport.grpc as transport_grpc + + interceptor = mock.Mock() + interceptor._wrapper = mock.Mock() + interceptor._wrapper._cached_cert = "cert" + + iterator = transport_grpc._RetryableStreamResponseIterator( + continuation=lambda cd, pl: mock.Mock(), + client_call_details=mock.Mock(), + request_or_iterator=b"request", + interceptor=interceptor, + is_client_stream=False, + ) + + assert iterator.done() is False + iterator._is_completed = True + assert iterator.done() is True + + +def test_start_call_wrapper_none(): + import pytest + import google.auth.transport.grpc as transport_grpc + + interceptor = transport_grpc.CertRotationInterceptor() + if hasattr(interceptor, "_wrapper"): + del interceptor._wrapper + + inner_future = mock.Mock() + call_details = mock.Mock() + + with pytest.raises(AttributeError): + transport_grpc._RetryableUnaryResponseFuture( + continuation=lambda cd, pl: inner_future, + client_call_details=call_details, + request_or_iterator=b"request", + interceptor=interceptor, + is_client_stream=False, + ) + + +def test_start_call_wrapper_none_branch(): + import google.auth.transport.grpc as transport_grpc + + interceptor = transport_grpc.CertRotationInterceptor() + interceptor._wrapper = None + + inner_future = mock.Mock() + call_details = mock.Mock() + + future = transport_grpc._RetryableUnaryResponseFuture( + continuation=lambda cd, pl: inner_future, + client_call_details=call_details, + request_or_iterator=b"request", + interceptor=interceptor, + is_client_stream=False, + ) + future._completion_event.set() + assert getattr(future, "_attempt_cert", "NOT_SET") is None + + +@mock.patch("google.auth.transport.grpc.CertRotationInterceptor._should_retry") +def test_unary_response_future_rpc_error_no_wrapper(mock_should_retry): + import google.auth.transport.grpc as transport_grpc + + interceptor = transport_grpc.CertRotationInterceptor() + interceptor._wrapper = None + + mock_err = transport_grpc.grpc.RpcError() + mock_err.code = lambda: transport_grpc.grpc.StatusCode.UNAUTHENTICATED + mock_should_retry.return_value = (True, b"cert", b"key", None) + + inner_future = mock.Mock() + inner_future.cancelled.return_value = False + inner_future.exception.return_value = mock_err + + call_details = mock.Mock() + + future = transport_grpc._RetryableUnaryResponseFuture( + continuation=lambda cd, pl: inner_future, + client_call_details=call_details, + request_or_iterator=b"request", + interceptor=interceptor, + is_client_stream=False, + ) + future._completion_event.set() + + with mock.patch.object(future, "_start_call", side_effect=mock_err): + future._on_inner_future_done(inner_future) + + +@mock.patch("google.auth.transport.grpc.CertRotationInterceptor._should_retry") +def test_unary_response_future_rpc_error_should_not_retry(mock_should_retry): + import google.auth.transport.grpc as transport_grpc + + interceptor = transport_grpc.CertRotationInterceptor() + interceptor._wrapper = mock.Mock() + interceptor._wrapper._cached_cert = "cert" + + mock_err = transport_grpc.grpc.RpcError() + mock_err.code = lambda: transport_grpc.grpc.StatusCode.UNAUTHENTICATED + mock_should_retry.return_value = (False, None, None, None) + + inner_future = mock.Mock() + inner_future.cancelled.return_value = False + inner_future.exception.return_value = mock_err + + call_details = mock.Mock() + + future = transport_grpc._RetryableUnaryResponseFuture( + continuation=lambda cd, pl: inner_future, + client_call_details=call_details, + request_or_iterator=b"request", + interceptor=interceptor, + is_client_stream=False, + ) + future._completion_event.set() + + with mock.patch.object(future, "_start_call", side_effect=mock_err): + future._on_inner_future_done(inner_future) + + interceptor._wrapper.refresh_logic.assert_not_called() + + +def test_mtls_call_interceptor_should_retry_cases(): + from unittest import mock + import google.auth.transport.grpc as transport_grpc + + interceptor = transport_grpc.CertRotationInterceptor() + + assert interceptor._should_retry(grpc.StatusCode.UNAUTHENTICATED, 0, "cert1") == ( + False, + None, + None, + None, + ) + + wrapper_mock = mock.Mock() + wrapper_mock._cached_cert = "cert1" + interceptor._wrapper = wrapper_mock + assert interceptor._should_retry(grpc.StatusCode.INTERNAL, 0, "cert1") == ( + False, + None, + None, + None, + ) + assert interceptor._should_retry(grpc.StatusCode.UNAUTHENTICATED, 2, "cert1") == ( + False, + None, + None, + None, + ) + + wrapper_mock._cached_cert = "cert2" + assert interceptor._should_retry(grpc.StatusCode.UNAUTHENTICATED, 0, "cert1") == ( + True, + None, + None, + None, + ) + + wrapper_mock._cached_cert = "cert1" + with mock.patch( + "google.auth.transport._mtls_helper.check_parameters_for_unauthorized_response" + ) as mock_check: + mock_check.return_value = (None, None, None, "fp1", "fp2") + assert interceptor._should_retry( + grpc.StatusCode.UNAUTHENTICATED, 0, "cert1" + ) == (True, None, None, None) + + with mock.patch( + "google.auth.transport._mtls_helper.check_parameters_for_unauthorized_response" + ) as mock_check: + mock_check.return_value = (None, None, None, "fp1", "fp1") + assert interceptor._should_retry( + grpc.StatusCode.UNAUTHENTICATED, 0, "cert1" + ) == (False, None, None, None) + + +def test_mtls_call_interceptor_interceptors_methods(): + from unittest import mock + import google.auth.transport.grpc as transport_grpc + + interceptor = transport_grpc.CertRotationInterceptor() + + def dummy_continuation(*args, **kwargs): + return mock.Mock() + + mock_details = mock.Mock() + mock_request = mock.Mock() + + res = interceptor.intercept_unary_unary( + dummy_continuation, mock_details, mock_request + ) + assert isinstance(res, transport_grpc._RetryableUnaryResponseFuture) + + res = interceptor.intercept_stream_unary( + dummy_continuation, mock_details, mock_request + ) + assert isinstance(res, transport_grpc._RetryableUnaryResponseFuture) + + res = interceptor.intercept_unary_stream( + dummy_continuation, mock_details, mock_request + ) + assert isinstance(res, transport_grpc._RetryableStreamResponseIterator) + + res = interceptor.intercept_stream_stream( + dummy_continuation, mock_details, mock_request + ) + assert isinstance(res, transport_grpc._RetryableStreamResponseIterator) + + +def test_mtls_refreshing_channel_refresh_logic_cases(): + from unittest import mock + import google.auth.transport.grpc as transport_grpc + + channel = transport_grpc.MTLSRefreshingChannel("target", {}, mock.Mock(), b"cert1") + with mock.patch( + "google.auth.transport._mtls_helper.check_parameters_for_unauthorized_response" + ) as mock_check, mock.patch("google.auth.transport.grpc.secure_authorized_channel"): + mock_check.return_value = (None, None, None, "fp1", "fp1") + assert ( + channel.refresh_logic( + 1, + call_cert_bytes=mock_check.return_value[0], + call_key_bytes=mock_check.return_value[1], + ) + is None + ) + mock_check.return_value = (None, None, None, "fp1", "fp2") + assert ( + channel.refresh_logic( + 0, + call_cert_bytes=mock_check.return_value[0], + call_key_bytes=mock_check.return_value[1], + ) + is None + ) + + def dummy_callback(): + return b"cert2", b"key2", None + + with mock.patch( + "google.auth.transport._mtls_helper.check_parameters_for_unauthorized_response" + ) as mock_check, mock.patch( + "google.auth.transport.grpc.secure_authorized_channel" + ) as mock_secure: + mock_check.return_value = (b"cert2", b"key2", None, "fp1", "fp2") + new_channel_mock = mock.Mock() + mock_secure.return_value = new_channel_mock + assert ( + channel.refresh_logic( + 0, + call_cert_bytes=mock_check.return_value[0], + call_key_bytes=mock_check.return_value[1], + ) + is None + ) + assert channel._cached_cert == b"cert2" + assert channel._channel == new_channel_mock + + +def test_mtls_refreshing_channel_subscribe_unsubscribe_close(): + from unittest import mock + import google.auth.transport.grpc as transport_grpc + + channel = transport_grpc.MTLSRefreshingChannel("target", {}, mock.Mock(), b"cert1") + cb = mock.Mock() + channel.subscribe(cb) + assert cb in channel._subscribers + channel.unsubscribe(cb) + assert cb not in channel._subscribers + + channel.close() + assert channel._channel.close.called + + +def test_mtls_refreshing_channel_unary_unary(): + from unittest import mock + import google.auth.transport.grpc as transport_grpc + + channel = transport_grpc.MTLSRefreshingChannel("target", {}, mock.Mock(), b"cert1") + res = channel.unary_unary("method") + assert res is not None + + from unittest import mock + import google.auth.transport.grpc as transport_grpc + + channel = transport_grpc.MTLSRefreshingChannel("target", {}, mock.Mock(), b"cert1") + res = channel.unary_unary("method") + assert res is not None + + +def test_mtls_refreshing_channel_unary_stream(): + from unittest import mock + import google.auth.transport.grpc as transport_grpc + + channel = transport_grpc.MTLSRefreshingChannel("target", {}, mock.Mock(), b"cert1") + res = channel.unary_stream("method") + assert res is not None + + +def test_mtls_refreshing_channel_stream_unary(): + from unittest import mock + import google.auth.transport.grpc as transport_grpc + + channel = transport_grpc.MTLSRefreshingChannel("target", {}, mock.Mock(), b"cert1") + res = channel.stream_unary("method") + assert res is not None + + +def test_mtls_refreshing_channel_stream_stream(): + from unittest import mock + import google.auth.transport.grpc as transport_grpc + + channel = transport_grpc.MTLSRefreshingChannel("target", {}, mock.Mock(), b"cert1") + res = channel.stream_stream("method") + assert res is not None + + +def test_retryable_unary_response_future_methods(): + from unittest import mock + import google.auth.transport.grpc as transport_grpc + + mock_future = mock.Mock() + + def dummy_continuation(*args, **kwargs): + return mock_future + + interceptor = transport_grpc.CertRotationInterceptor() + interceptor._wrapper = mock.Mock() + interceptor._wrapper._cached_cert = "cert" + + future = transport_grpc._RetryableUnaryResponseFuture( + dummy_continuation, mock.Mock(), mock.Mock(), interceptor + ) + + future._completion_event.set() + future.initial_metadata() + future.trailing_metadata() + future.code() + future.details() + future.cancel() + future.cancelled() + future.is_active() + future.time_remaining() + + mock_future.result.return_value = "r" + assert future.result() == "r" + mock_future.exception.return_value = Exception("e") + assert isinstance(future.exception(), Exception) + mock_future.traceback.return_value = "tb" + assert future.traceback() == "tb" + + future.add_done_callback(lambda x: None) + + +def test_retryable_stream_response_iterator_methods(): + from unittest import mock + import google.auth.transport.grpc as transport_grpc + + mock_iterator = mock.Mock() + + def dummy_continuation(*args, **kwargs): + return mock_iterator + + interceptor = transport_grpc.CertRotationInterceptor() + interceptor._wrapper = mock.Mock() + interceptor._wrapper._cached_cert = "cert" + + iterator = transport_grpc._RetryableStreamResponseIterator( + dummy_continuation, mock.Mock(), mock.Mock(), interceptor + ) + + iterator.initial_metadata() + iterator.trailing_metadata() + iterator.code() + iterator.details() + iterator.cancel() + iterator.cancelled() + iterator.is_active() + iterator.time_remaining() + iterator.add_done_callback(lambda x: None) + + def test_grpc_version_warning_for_older_version(monkeypatch): monkeypatch.setattr(grpc, "__version__", "1.80.0") with pytest.warns( diff --git a/packages/google-auth/tests/transport/test_requests.py b/packages/google-auth/tests/transport/test_requests.py index 2ca1922494ef..8c1d7b7e63d3 100644 --- a/packages/google-auth/tests/transport/test_requests.py +++ b/packages/google-auth/tests/transport/test_requests.py @@ -744,7 +744,7 @@ def test_cert_rotation_when_cert_mismatch_and_mtls_enabled(self): with mock.patch.object( google.auth.transport._mtls_helper, "call_client_cert_callback", - return_value=(new_cert, new_key), + return_value=(new_cert, new_key, None), ) as mock_callback: result = authed_session.request("GET", self.MTLS_TEST_URL) @@ -783,7 +783,7 @@ def test_no_cert_rotation_when_cert_match_and_mTLS_enabled(self): with mock.patch.object( google.auth.transport._mtls_helper, "call_client_cert_callback", - return_value=(new_cert, new_key), + return_value=(new_cert, new_key, None), ): result = authed_session.request("GET", self.MTLS_TEST_URL) @@ -816,7 +816,7 @@ def test_no_cert_match_check_when_mtls_disabled(self): with mock.patch.object( google.auth.transport._mtls_helper, "call_client_cert_callback", - return_value=(new_cert, new_key), + return_value=(new_cert, new_key, None), ) as mock_callback: result = authed_session.request("GET", self.TEST_URL) @@ -866,7 +866,7 @@ def test_cert_rotation_failure_raises_error(self): with mock.patch.object( google.auth.transport._mtls_helper, "call_client_cert_callback", - return_value=(new_cert, new_key), + return_value=(new_cert, new_key, None), ): with mock.patch.object( authed_session, diff --git a/packages/google-auth/tests/transport/test_urllib3.py b/packages/google-auth/tests/transport/test_urllib3.py index e1c92dbebc2c..fcdb6e7da099 100644 --- a/packages/google-auth/tests/transport/test_urllib3.py +++ b/packages/google-auth/tests/transport/test_urllib3.py @@ -465,7 +465,7 @@ def test_cert_rotation_when_cert_mismatch_and_mtls_endpoint_used( with mock.patch.object( google.auth.transport._mtls_helper, "call_client_cert_callback", - return_value=(new_cert, new_key), + return_value=(new_cert, new_key, None), ) as mock_callback: # mTLS endpoint is used, and client cert env var is true with mock.patch.dict( @@ -506,7 +506,7 @@ def test_no_cert_rotation_when_cert_match_and_mtls_endpoint_used(self): with mock.patch.object( google.auth.transport._mtls_helper, "call_client_cert_callback", - return_value=(new_cert, new_key), + return_value=(new_cert, new_key, None), ): # mTLS endpoint is used result = authed_http.urlopen("GET", "http://example.mtls.googleapis.com") @@ -536,7 +536,7 @@ def test_no_cert_match_check_when_mtls_endpoint_not_used(self): with mock.patch.object( google.auth.transport._mtls_helper, "call_client_cert_callback", - return_value=(new_cert, new_key), + return_value=(new_cert, new_key, None), ) as mock_callback: # non-mTLS endpoint is used result = authed_http.urlopen("GET", "http://example.googleapis.com") @@ -584,7 +584,13 @@ def test_cert_rotation_failure_raises_error(self): with mock.patch.object( google.auth.transport._mtls_helper, "check_parameters_for_unauthorized_response", - return_value=(new_cert, new_key, "old_fingerprint", "new_fingerprint"), + return_value=( + new_cert, + new_key, + None, + "old_fingerprint", + "new_fingerprint", + ), ) as mock_check_params: with mock.patch.object( authed_http, diff --git a/packages/google-auth/tests_async/transport/test_aiohttp_requests.py b/packages/google-auth/tests_async/transport/test_aiohttp_requests.py index 1dc5b0025edc..912aeecf8df1 100644 --- a/packages/google-auth/tests_async/transport/test_aiohttp_requests.py +++ b/packages/google-auth/tests_async/transport/test_aiohttp_requests.py @@ -128,12 +128,16 @@ def test_mock_session_unspecified_auto_decompress(self): request = aiohttp_requests.Request(http) assert request.session == http - def test_timeout(self): + @pytest.mark.asyncio + async def test_timeout(self): http = mock.create_autospec( aiohttp.ClientSession, instance=True, auto_decompress=False ) + mock_response = mock.AsyncMock() + http.request = mock.AsyncMock(return_value=mock_response) request = aiohttp_requests.Request(http) - request(url="http://example.com", method="GET", timeout=5) + await request(url="http://example.com", method="GET", timeout=5) + assert http.request.call_args[1]["timeout"] == 5 @pytest.mark.asyncio async def test__clone(self):