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..e8a43d4d4f49 --- /dev/null +++ b/packages/google-auth/google/auth/transport/mtls_interceptor.py @@ -0,0 +1,723 @@ +# Copyright 2016 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""mTLS Interceptor and Channel Wrapper for certificate rotation.""" + +import collections +import logging +import threading +import time + +import grpc +from google.auth import transport +from google.auth.transport import _mtls_helper + +_LOGGER = logging.getLogger(__name__) + +class CertRotationInterceptor( + grpc.UnaryUnaryClientInterceptor, + grpc.UnaryStreamClientInterceptor, + grpc.StreamUnaryClientInterceptor, + grpc.StreamStreamClientInterceptor, +): + """A gRPC client interceptor that provides automatic retry logic for mTLS certificate rotation. + + This interceptor wraps all gRPC client calls (unary and streaming) with retryable + futures or iterators. Its primary role is to monitor responses for `UNAUTHENTICATED` + errors. When an authentication failure occurs, it uses `_should_retry()` to check + if a new mTLS certificate is available. If a new certificate is found, it signals + its associated `MTLSRefreshingChannel` wrapper to refresh the underlying gRPC + channel's credentials and automatically replays the failed RPC. + """ + + 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): + """Determines if the RPC should be retried due to a certificate rotation. + + Returns a tuple: (should_retry, call_cert_bytes, call_key_bytes, passphrase). + """ + if code != grpc.StatusCode.UNAUTHENTICATED or not self._wrapper: + return False, None, None, None + + if retry_count >= self._max_retries: + _LOGGER.debug( + "Max retries reached (%d/%d) for channel recreation.", + retry_count, + self._max_retries, + ) + return False, None, None, None + + # If another thread already refreshed the channel with an updated cert, retry immediately + if attempt_cert != self._wrapper._cached_cert: + return True, None, None, None + + # Check if the certificate on disk or callback has changed since this request was attempted + ( + call_cert_bytes, + call_key_bytes, + passphrase, + cached_fingerprint, + current_cert_fingerprint, + ) = _mtls_helper.check_parameters_for_unauthorized_response(attempt_cert) + should_retry = cached_fingerprint != current_cert_fingerprint + return should_retry, call_cert_bytes, call_key_bytes, passphrase + + def intercept_unary_unary(self, continuation, client_call_details, request): + return _RetryableUnaryResponseFuture( + continuation, client_call_details, request, self, is_client_stream=False + ) + + def intercept_stream_unary( + self, continuation, client_call_details, request_iterator + ): + return _RetryableUnaryResponseFuture( + continuation, + client_call_details, + request_iterator, + self, + is_client_stream=True, + ) + + def intercept_unary_stream(self, continuation, client_call_details, request): + return _RetryableStreamResponseIterator( + continuation, client_call_details, request, self, is_client_stream=False + ) + + def intercept_stream_stream( + self, continuation, client_call_details, request_iterator + ): + return _RetryableStreamResponseIterator( + continuation, + client_call_details, + request_iterator, + self, + is_client_stream=True, + ) + + +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 = set() + + def refresh_logic( + self, count, call_cert_bytes=None, call_key_bytes=None, passphrase=None + ): + with self._lock: + if not call_cert_bytes or self._cached_cert == call_cert_bytes: + return + + _LOGGER.debug("Wrapper: Refreshing mTLS channel. Retry count: %d", count) + old_channel = self._channel + + if passphrase is not None: + call_key_bytes = _mtls_helper.decrypt_private_key( + call_key_bytes, passphrase + ) + + # Call the partial, overriding only the cert-related arguments + 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: + try: + old_channel.unsubscribe(callback) + except Exception: + pass + self._channel.subscribe(callback) + + def unary_unary(self, method, *args, **kwargs): + # Always return a callable from the CURRENT channel + return self._channel.unary_unary(method, *args, **kwargs) + + # Mandatory passthroughs + def unary_stream(self, method, *args, **kwargs): + return self._channel.unary_stream(method, *args, **kwargs) + + def stream_unary(self, method, *args, **kwargs): + return self._channel.stream_unary(method, *args, **kwargs) + + def stream_stream(self, method, *args, **kwargs): + return self._channel.stream_stream(method, *args, **kwargs) + + def subscribe(self, callback, try_to_connect=False): + with self._lock: + self._subscribers.add(callback) + return self._channel.subscribe(callback, try_to_connect=try_to_connect) + + def unsubscribe(self, callback): + with self._lock: + self._subscribers.discard(callback) + return self._channel.unsubscribe(callback) + + def close(self): + self._channel.close() + + +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() + self._consumer_lock = threading.Lock() + self._active_reader = None + + def __iter__(self): + reader = _ReplayableIteratorReader(self) + with self._lock: + self._active_reader = reader + return reader + + 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): + while True: + with self._parent._lock: + if self._read_index < len(self._parent._buffer): + val = self._parent._buffer[self._read_index] + self._read_index += 1 + return val + + if self._parent._exhausted: + raise StopIteration() + + if self._parent._active_reader is not self: + raise StopIteration() + + with self._parent._consumer_lock: + with self._parent._lock: + if self._read_index < len(self._parent._buffer): + continue + if self._parent._active_reader is not self: + raise StopIteration() + + try: + val = next(self._parent._target_iterator) + except StopIteration: + with self._parent._lock: + if self._parent._active_reader is self: + self._parent._exhausted = True + raise + + with self._parent._lock: + if self._parent._active_reader is not self: + if self._parent._can_replay: + self._parent._buffer.append(val) + raise StopIteration() + + if self._parent._can_replay: + self._parent._buffer.append(val) + if len(self._parent._buffer) > self._parent._max_items: + self._parent._buffer.clear() + self._parent._can_replay = False + + self._read_index += 1 + return val + + +_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 _RetryableUnaryResponseFuture(_BaseCallWrapper): + def __init__( + self, + continuation, + client_call_details, + request_or_iterator, + interceptor, + is_client_stream=False, + ): + self._continuation = continuation + self._client_call_details = client_call_details + self._is_client_stream = is_client_stream + self._source_request = request_or_iterator + self._interceptor = interceptor + + self._uses_factory = is_client_stream and callable(request_or_iterator) + self._payload = ( + None + if self._uses_factory + else ( + _ReplayableIterator(request_or_iterator) + if is_client_stream + else request_or_iterator + ) + ) + + self._retry_count = 0 + self._call = None + self._lock = threading.RLock() + + timeout = getattr(self._client_call_details, "timeout", None) + self._initial_timeout = timeout if isinstance(timeout, (int, float)) else None + self._start_time = time.monotonic() if self._initial_timeout else None + + self._completion_event = threading.Event() + self._done_callbacks = [] + self._terminal_exception = None + + self._start_call() + + def _start_call(self): + self._attempt_cert = ( + self._interceptor._wrapper._cached_cert + if self._interceptor._wrapper + else None + ) + + with self._lock: + if self._uses_factory: + payload = self._source_request() + else: + payload = ( + iter(self._payload) if self._is_client_stream else self._payload + ) + + call_details = self._client_call_details + if self._start_time and 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(): + with self._lock: + self._completion_event.set() + callbacks_to_fire = list(self._done_callbacks) + for fn in callbacks_to_fire: + try: + fn(self) + except Exception as e: + _LOGGER.warning("Callback failed: %s", e) + return + + exc = inner_future.exception() + if isinstance(exc, grpc.RpcError): + status_code = exc.code() + + can_replay = ( + True + if self._uses_factory + else (self._payload.can_replay() if self._is_client_stream else True) + ) + + should_retry, call_cert, call_key, pwd = self._interceptor._should_retry( + status_code, self._retry_count, getattr(self, "_attempt_cert", None) + ) + if can_replay and should_retry: + if getattr(self._interceptor, "_wrapper", None): + try: + self._interceptor._wrapper.refresh_logic( + 1, call_cert, call_key, pwd + ) + except Exception as e: + with self._lock: + self._terminal_exception = e + self._completion_event.set() + return + with self._lock: + self._retry_count += 1 + try: + self._start_call() + return + except Exception as e: + self._terminal_exception = e + + if self._interceptor._wrapper: + ( + chk_should_retry, + chk_cert, + chk_key, + chk_pwd, + ) = self._interceptor._should_retry(status_code, 0, self._attempt_cert) + if chk_should_retry: + try: + self._interceptor._wrapper.refresh_logic( + 1, chk_cert, chk_key, chk_pwd + ) + except Exception: + pass + with self._lock: + self._completion_event.set() + callbacks_to_fire = list(self._done_callbacks) + + for fn in callbacks_to_fire: + try: + fn(self) + except Exception as e: + _LOGGER.warning("Callback failed: %s", e) + + def add_done_callback(self, fn): + with self._lock: + if self._completion_event.is_set(): + fire_now = True + else: + self._done_callbacks.append(fn) + fire_now = False + + if fire_now: + try: + fn(self) + except Exception: + pass + + def result(self, timeout=None): + if not self._completion_event.wait(timeout): + raise grpc.FutureTimeoutError() + with self._lock: + if self._terminal_exception is not None: + raise self._terminal_exception + current_future = self._call + return current_future.result() + + def exception(self, timeout=None): + if not self._completion_event.wait(timeout): + raise grpc.FutureTimeoutError() + with self._lock: + if self._terminal_exception is not None: + return self._terminal_exception + return self._call.exception() + + def traceback(self, timeout=None): + if not self._completion_event.wait(timeout): + raise grpc.FutureTimeoutError() + with self._lock: + if self._terminal_exception is not None: + return self._terminal_exception.__traceback__ + return self._call.traceback() + + def initial_metadata(self): + self._completion_event.wait() + with self._lock: + if self._terminal_exception is not None: + return None + return self._call.initial_metadata() + + def trailing_metadata(self): + self._completion_event.wait() + with self._lock: + if self._terminal_exception is not None: + return None + return self._call.trailing_metadata() + + def code(self): + self._completion_event.wait() + with self._lock: + if hasattr(self._terminal_exception, "code"): + return self._terminal_exception.code() + return self._call.code() + + def details(self): + self._completion_event.wait() + with self._lock: + if hasattr(self._terminal_exception, "details"): + return self._terminal_exception.details() + return self._call.details() + +class _RetryableStreamResponseIterator(_BaseCallWrapper): + def __init__( + self, + continuation, + client_call_details, + request_or_iterator, + interceptor, + is_client_stream=False, + ): + self._continuation = continuation + self._client_call_details = client_call_details + self._is_client_stream = is_client_stream + self._source_request = request_or_iterator + self._interceptor = interceptor + + self._uses_factory = is_client_stream and callable(request_or_iterator) + self._payload = ( + None + if self._uses_factory + else ( + _ReplayableIterator(request_or_iterator) + if is_client_stream + else request_or_iterator + ) + ) + self._call = None + self._retry_count = 0 + self._yielded_any_response = False + self._lock = threading.RLock() + + timeout = getattr(self._client_call_details, "timeout", None) + self._initial_timeout = timeout if isinstance(timeout, (int, float)) else None + self._start_time = time.monotonic() if self._initial_timeout else None + + self._is_completed = False + self._done_callbacks = [] + + self._start_call() + + def _start_call(self): + self._attempt_cert = ( + self._interceptor._wrapper._cached_cert + if getattr(self._interceptor, "_wrapper", None) + else None + ) + with self._lock: + if self._uses_factory: + payload = self._source_request() + else: + payload = ( + iter(self._payload) if self._is_client_stream else self._payload + ) + + call_details = self._client_call_details + if self._start_time and 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_call_done) + + def _trigger_callbacks(self): + with self._lock: + if self._is_completed: + return + self._is_completed = True + callbacks = list(self._done_callbacks) + + for fn in callbacks: + try: + fn(self) + except Exception: + pass + + def _on_inner_call_done(self, inner_call): + with self._lock: + if self._call is not inner_call: + return + # Intercept and suppress premature callbacks for UNAUTHENTICATED. + # __next__ inherently handles this error and manages triggering callbacks + # later if retriies are exhausted. + if ( + callable(getattr(inner_call, "code", None)) + and inner_call.code() == grpc.StatusCode.UNAUTHENTICATED + ): + return + self._trigger_callbacks() + + def __iter__(self): + return self + + def __next__(self): + while True: + with self._lock: + current_call = self._call + + try: + response = next(current_call) + self._yielded_any_response = True + return response + except StopIteration: + self._trigger_callbacks() + raise + except grpc.RpcError as e: + status_code = getattr(e, "code", lambda: None)() + with self._lock: + if self._call is not current_call: + continue + + can_replay = ( + True + if self._uses_factory + else ( + self._payload.can_replay() if self._is_client_stream else True + ) + ) + + ( + should_retry, + call_cert, + call_key, + pwd, + ) = self._interceptor._should_retry( + status_code, + self._retry_count, + getattr(self, "_attempt_cert", None), + ) + + if not self._yielded_any_response and can_replay and should_retry: + try: + if getattr(self._interceptor, "_wrapper", None): + self._interceptor._wrapper.refresh_logic( + 1, call_cert, call_key, pwd + ) + + with self._lock: + self._retry_count += 1 + self._start_call() + + except Exception as fallback_e: + self._trigger_callbacks() + raise fallback_e + + continue + else: + # Non-retryable error, check if another rotation happened while we were finishing + if getattr(self._interceptor, "_wrapper", None): + ( + chk_should_retry, + chk_cert, + chk_key, + chk_pwd, + ) = self._interceptor._should_retry( + status_code, 0, getattr(self, "_attempt_cert", None) + ) + if chk_should_retry: + try: + self._interceptor._wrapper.refresh_logic( + 1, chk_cert, chk_key, chk_pwd + ) + except Exception: + pass # Terminal anyway + + self._trigger_callbacks() + raise e + + def add_done_callback(self, fn): + with self._lock: + if getattr(self, "_is_completed", False): + fire_now = True + else: + self._done_callbacks.append(fn) + fire_now = False + + if fire_now: + try: + fn(self) + except Exception: + pass + +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() + + def time_remaining(self): + return self._call.time_remaining() + + def add_callback(self, callback): + self._call.add_callback(callback) 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..316d641c6281 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,15 +1253,15 @@ 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 ) - class TestSecureCertKeyPaths(object): def test_tier1_pass_through(self): with _mtls_helper.secure_cert_key_paths( diff --git a/packages/google-auth/tests/transport/test_grpc.py b/packages/google-auth/tests/transport/test_grpc.py index 7979df7abb4d..af8e820d1481 100644 --- a/packages/google-auth/tests/transport/test_grpc.py +++ b/packages/google-auth/tests/transport/test_grpc.py @@ -13,12 +13,11 @@ # limitations under the License. import datetime -import importlib import os import time from unittest import mock -import warnings +import grpc import pytest # type: ignore from google.auth import _helpers @@ -28,9 +27,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 @@ -134,35 +142,6 @@ def test__get_authorization_headers_with_service_account_and_default_host(self): "https://{}/".format(default_host) ) - def test_suppress_metrics_header(self): - credentials = mock.create_autospec(service_account.Credentials) - - # Mock credentials before_request that adds metric and authorization - def mock_before_request(request, method, url, headers): - headers["x-goog-api-client"] = "foo" - headers["authorization"] = "Bearer token" - - credentials.before_request.side_effect = mock_before_request - request = mock.create_autospec(transport.Request) - - # By default, suppress_metrics_header=False - plugin = google.auth.transport.grpc.AuthMetadataPlugin(credentials, request) - context = mock.create_autospec(grpc.AuthMetadataContext, instance=True) - context.method_name = "methodName" - context.service_url = "https://pubsub.googleapis.com/methodName" - - headers = dict(plugin._get_authorization_headers(context)) - assert "x-goog-api-client" in headers - assert headers["x-goog-api-client"] == "foo" - - # With suppress_metrics_header=True - plugin_suppressed = google.auth.transport.grpc.AuthMetadataPlugin( - credentials, request, suppress_metrics_header=True - ) - headers_suppressed = dict(plugin_suppressed._get_authorization_headers(context)) - assert "x-goog-api-client" not in headers_suppressed - assert headers_suppressed["authorization"] == "Bearer token" - @mock.patch( "google.auth.transport._mtls_helper.get_client_ssl_credentials", autospec=True @@ -229,7 +208,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 +254,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, @@ -679,26 +658,3 @@ def test_get_client_ssl_credentials_auto_enablement( mock_get_client_ssl_credentials.assert_called_once() mock_ssl_channel_credentials.assert_called_once_with( certificate_chain=PUBLIC_CERT_BYTES, private_key=PRIVATE_KEY_BYTES - ) - - -def test_grpc_version_warning_for_older_version(monkeypatch): - monkeypatch.setattr(grpc, "__version__", "1.80.0") - with pytest.warns( - FutureWarning, match="does not support Post-Quantum Cryptography" - ): - importlib.reload(google.auth.transport.grpc) - - -def test_grpc_version_warning_not_emitted_for_supported_version(monkeypatch): - monkeypatch.setattr(grpc, "__version__", "1.83.0") - with warnings.catch_warnings(): - warnings.simplefilter("error", FutureWarning) - importlib.reload(google.auth.transport.grpc) - - -def test_grpc_version_warning_not_emitted_when_no_version(monkeypatch): - monkeypatch.delattr(grpc, "__version__", raising=False) - with warnings.catch_warnings(): - warnings.simplefilter("error", FutureWarning) - importlib.reload(google.auth.transport.grpc) diff --git a/packages/google-auth/tests/transport/test_mtls_interceptor.py b/packages/google-auth/tests/transport/test_mtls_interceptor.py new file mode 100644 index 000000000000..ca42097f0ea4 --- /dev/null +++ b/packages/google-auth/tests/transport/test_mtls_interceptor.py @@ -0,0 +1,59 @@ +# Copyright 2024 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import time +from unittest import mock +import grpc +import pytest +import google.auth.transport.grpc as transport_grpc +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]) 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):