From fa24e36e8f248ada2ef8bdc88565268d76dfe7ff Mon Sep 17 00:00:00 2001 From: agrawalradhika-cell Date: Thu, 6 Aug 2026 13:39:55 -0700 Subject: [PATCH 01/10] feat: Add passphrase handling to client cert callback feat: Add passphrase handling to client cert callback --- .../google/auth/transport/_mtls_helper.py | 20 ++++++++++++++----- 1 file changed, 15 insertions(+), 5 deletions(-) diff --git a/packages/google-auth/google/auth/transport/_mtls_helper.py b/packages/google-auth/google/auth/transport/_mtls_helper.py index eb0600740c0d..d0ae71434b46 100644 --- a/packages/google-auth/google/auth/transport/_mtls_helper.py +++ b/packages/google-auth/google/auth/transport/_mtls_helper.py @@ -609,7 +609,10 @@ def get_client_ssl_credentials( """ # 1. Attempt to retrieve X.509 Workload cert and key. - cert, key = _get_workload_cert_and_key(certificate_config_path) + try: + cert, key = _get_workload_cert_and_key(certificate_config_path) + except exceptions.ClientCertError: + cert, key = None, None if cert and key: return True, cert, key, None @@ -784,10 +787,11 @@ def check_parameters_for_unauthorized_response(cached_cert): Returns: bytes: The client callback cert bytes. bytes: The client callback key bytes. + 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 @@ -798,12 +802,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 From 49a141260815c45af978dadd2c7f6f8c3a58d044 Mon Sep 17 00:00:00 2001 From: agrawalradhika-cell Date: Thu, 6 Aug 2026 13:41:48 -0700 Subject: [PATCH 02/10] feat: Add cert rotation handling support feat: Add cert rotation handling support --- .../google-auth/google/auth/transport/grpc.py | 685 +++++++++++++++++- 1 file changed, 681 insertions(+), 4 deletions(-) diff --git a/packages/google-auth/google/auth/transport/grpc.py b/packages/google-auth/google/auth/transport/grpc.py index 5930a07f1d08..f6e8fdc72dc9 100644 --- a/packages/google-auth/google/auth/transport/grpc.py +++ b/packages/google-auth/google/auth/transport/grpc.py @@ -16,7 +16,11 @@ from __future__ import absolute_import +import collections +import concurrent.futures import logging +import threading +import time from google.auth import exceptions from google.auth.transport import _mtls_helper @@ -262,6 +266,7 @@ def my_client_cert_callback(): ) # If SSL credentials are not explicitly set, try client_cert_callback and ADC. + cached_cert = None if not ssl_credentials: use_client_cert = _mtls_helper.check_use_client_cert() if use_client_cert and client_cert_callback: @@ -270,10 +275,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() @@ -281,8 +288,27 @@ 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_retry = kwargs.pop("_is_retry", False) + channel = grpc.secure_channel(target, composite_credentials, **kwargs) + # Check if we are already inside a retry to avoid infinite recursion + if cached_cert and not is_retry: + # Package arguments to recreate the channel if rotation occurs + factory_args = { + "credentials": credentials, + "request": request, + "target": target, + "ssl_credentials": None, + "client_cert_callback": client_cert_callback, + "_is_retry": True, # Hidden flag to stop recursion + **kwargs, + } + interceptor = _MTLSCallInterceptor() + + wrapper = _MTLSRefreshingChannel(target, factory_args, channel, cached_cert) + + interceptor._wrapper = wrapper + return grpc.intercept_channel(wrapper, interceptor) + return channel class SslCredentials: @@ -306,6 +332,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: @@ -334,6 +361,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 @@ -349,3 +377,652 @@ def ssl_credentials(self): def is_mtls(self): """Indicates if the created SSL channel credentials is mutual TLS.""" return self._is_mtls + + +class _MTLSCallInterceptor( + grpc.UnaryUnaryClientInterceptor, + grpc.UnaryStreamClientInterceptor, + grpc.StreamUnaryClientInterceptor, + grpc.StreamStreamClientInterceptor, +): + def __init__(self): + self._wrapper = None + self._max_retries = 2 # Set your desired limit here + + def _should_retry(self, code, retry_count, attempt_cert): + if code != grpc.StatusCode.UNAUTHENTICATED or not self._wrapper: + return False + + if retry_count >= self._max_retries: + _LOGGER.debug( + "Max retries reached (%d/%d).", retry_count, self._max_retries + ) + return False + + # If the wrapper has already rotated to a new cert, we can retry immediately + if attempt_cert != self._wrapper._cached_cert: + return True + + # Fingerprint check logic + ( + _, + _, + _, + cached_fp, + current_fp, + ) = _mtls_helper.check_parameters_for_unauthorized_response(attempt_cert) + return cached_fp != current_fp + + 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, factory_args, initial_channel, initial_cert): + self._target = target + self._factory_args = factory_args + self._channel = initial_channel + self._cached_cert = initial_cert + self._lock = threading.Lock() + self._subscribers = set() + + def refresh_logic(self, count): + with self._lock: + # Re-check inside lock to prevent race conditions + ( + call_cert_bytes, + call_key_bytes, + passphrase, + cached_fp, + current_fp, + ) = _mtls_helper.check_parameters_for_unauthorized_response( + self._cached_cert + ) + if cached_fp != current_fp: + _LOGGER.debug( + "Wrapper: Refreshing mTLS channel. Retry count: %d", count + ) + old_channel = self._channel + + # Consume the exact credential bytes fetched during the fingerprint check + self._cached_cert = call_cert_bytes + + # Support encrypted keys + if passphrase is not None: + call_key_bytes = _mtls_helper.decrypt_private_key( + call_key_bytes, passphrase + ) + + # The factory args must use the new credentials exactly to build the rotation channel + factory_args = self._factory_args.copy() + factory_args["client_cert_callback"] = None + factory_args["ssl_credentials"] = grpc.ssl_channel_credentials( + certificate_chain=call_cert_bytes, + private_key=call_key_bytes, + ) + + self._channel = secure_authorized_channel(**factory_args) + + for callback in self._subscribers: + try: + old_channel.unsubscribe(callback) + except Exception: + pass + self._channel.subscribe(callback) + + try: + old_channel.close() + except Exception: + pass + + 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 = 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) > getattr( + self._parent, "_max_items", 10000 + ): + 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 _RetryableUnaryResponseFuture(grpc.Future, grpc.Call): + 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._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 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 grpc.RpcError("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._target_future = self._continuation(call_details, payload) + self._target_future.add_done_callback(self._on_inner_future_done) + + def _on_inner_future_done(self, inner_future): + with self._lock: + if self._target_future is not inner_future: + 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) + ) + + if can_replay and self._interceptor._should_retry( + status_code, self._retry_count, getattr(self, "_attempt_cert", None) + ): + if getattr(self._interceptor, "_wrapper", None): + self._interceptor._wrapper.refresh_logic(1) + + with self._lock: + self._retry_count += 1 + try: + self._start_call() + return + except Exception as e: + self._terminal_exception = e + + if getattr(self._interceptor, "_wrapper", None): + if self._interceptor._should_retry( + status_code, 0, getattr(self, "_attempt_cert", None) + ): + self._interceptor._wrapper.refresh_logic(1) + + 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: + pass + + 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._target_future + 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._target_future.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._target_future.traceback() + + def initial_metadata(self): + self._completion_event.wait() + with self._lock: + if self._terminal_exception is not None: + return None + return self._target_future.initial_metadata() + + def trailing_metadata(self): + self._completion_event.wait() + with self._lock: + if self._terminal_exception is not None: + return None + return self._target_future.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._target_future.code() + + def details(self): + self._completion_event.wait() + with self._lock: + if hasattr(self._terminal_exception, "details"): + return self._terminal_exception.details() + return self._target_future.details() + + def cancel(self): + with self._lock: + return self._target_future.cancel() + + def cancelled(self): + with self._lock: + return self._target_future.cancelled() + + def running(self): + with self._lock: + return self._target_future.running() + + def done(self): + with self._lock: + return self._completion_event.is_set() + + def is_active(self): + with self._lock: + return self._target_future.is_active() + + def time_remaining(self): + with self._lock: + return self._target_future.time_remaining() + + def add_callback(self, callback): + with self._lock: + return self._target_future.add_callback(callback) + + +class _RetryableStreamResponseIterator(grpc.Call): + 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._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 grpc.RpcError("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 + 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 + ) + ) + + if ( + not self._yielded_any_response + and can_replay + and self._interceptor._should_retry( + status_code, + self._retry_count, + getattr(self, "_attempt_cert", None), + ) + ): + if getattr(self._interceptor, "_wrapper", None): + self._interceptor._wrapper.refresh_logic(1) + + with self._lock: + self._retry_count += 1 + try: + self._start_call() + except Exception as timeout_e: + self._trigger_callbacks() + raise timeout_e + continue + else: + if getattr(self._interceptor, "_wrapper", None): + if self._interceptor._should_retry( + status_code, 0, getattr(self, "_attempt_cert", None) + ): + self._interceptor._wrapper.refresh_logic(1) + 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 + + def cancel(self): + with self._lock: + return self._call.cancel() + + def cancelled(self): + with self._lock: + return self._call.cancelled() + + def running(self): + with self._lock: + return self._call.running() + + def done(self): + with self._lock: + return getattr(self, "_is_completed", False) + + def initial_metadata(self): + with self._lock: + return self._call.initial_metadata() + + def trailing_metadata(self): + with self._lock: + return self._call.trailing_metadata() + + def code(self): + with self._lock: + return self._call.code() + + def details(self): + with self._lock: + return self._call.details() + + def is_active(self): + return self._call.is_active() + + def time_remaining(self): + return self._call.time_remaining() + + def add_callback(self, callback): + self._call.add_callback(callback) From c5020061e2b533bda902ec732b9b7da2f19f8176 Mon Sep 17 00:00:00 2001 From: agrawalradhika-cell Date: Thu, 6 Aug 2026 13:44:02 -0700 Subject: [PATCH 03/10] chore: Add passphrase in _mtls_helper for requests chore: Add passphrase in _mtls_helper for requests --- packages/google-auth/google/auth/transport/requests.py | 1 + 1 file changed, 1 insertion(+) diff --git a/packages/google-auth/google/auth/transport/requests.py b/packages/google-auth/google/auth/transport/requests.py index 822cf687f5d0..4ac32461c467 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( From 12169ffddf0e5864796775ac1eff8b3e501d1cec Mon Sep 17 00:00:00 2001 From: agrawalradhika-cell Date: Thu, 6 Aug 2026 13:44:53 -0700 Subject: [PATCH 04/10] chore: Modify _mtls_helper call to include additional variable passphrase chore: Modify _mtls_helper call to include additional variable passphrase --- packages/google-auth/google/auth/transport/urllib3.py | 1 + 1 file changed, 1 insertion(+) 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( From c90c0b1b2f8ef50428d268e2132e3fe20871b0b8 Mon Sep 17 00:00:00 2001 From: agrawalradhika-cell Date: Thu, 6 Aug 2026 13:46:48 -0700 Subject: [PATCH 05/10] chore: Modify mock return values in test_urllib3.py Updated mock return values in test cases to include None for additional parameters. --- packages/google-auth/tests/transport/test_urllib3.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/packages/google-auth/tests/transport/test_urllib3.py b/packages/google-auth/tests/transport/test_urllib3.py index e1c92dbebc2c..0705aa7cb7df 100644 --- a/packages/google-auth/tests/transport/test_urllib3.py +++ b/packages/google-auth/tests/transport/test_urllib3.py @@ -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,7 @@ 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, From 3cde59e177240583f28d150b1e4d259148d2bd05 Mon Sep 17 00:00:00 2001 From: agrawalradhika-cell Date: Thu, 6 Aug 2026 13:47:52 -0700 Subject: [PATCH 06/10] chore: Modify mock callback return value in tests all Updated mock callback to return an additional None value. --- packages/google-auth/tests/transport/test_urllib3.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/packages/google-auth/tests/transport/test_urllib3.py b/packages/google-auth/tests/transport/test_urllib3.py index 0705aa7cb7df..a01242f05b9f 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") From 96b41613a3f39369d28fdb0fae80f429214e956a Mon Sep 17 00:00:00 2001 From: agrawalradhika-cell Date: Thu, 6 Aug 2026 13:48:32 -0700 Subject: [PATCH 07/10] chore: Modify mock return value in test_requests.py Updated mock call_client_cert_callback to include a None value in the return tuple. --- packages/google-auth/tests/transport/test_requests.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) 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, From ba1c85a6fc207d8b0993cae44fd5c97a7571d37e Mon Sep 17 00:00:00 2001 From: agrawalradhika-cell Date: Thu, 6 Aug 2026 13:50:14 -0700 Subject: [PATCH 08/10] chore: Add unit tests for cert rotation handling for grpc chore: Add unit tests for cert rotation handling for grpc --- .../google-auth/tests/transport/test_grpc.py | 163 +++++++++++++++++- 1 file changed, 161 insertions(+), 2 deletions(-) diff --git a/packages/google-auth/tests/transport/test_grpc.py b/packages/google-auth/tests/transport/test_grpc.py index e5f9b7945a39..f78c13ce848a 100644 --- a/packages/google-auth/tests/transport/test_grpc.py +++ b/packages/google-auth/tests/transport/test_grpc.py @@ -26,9 +26,19 @@ 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 @@ -227,7 +237,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( @@ -273,7 +283,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, @@ -678,3 +688,152 @@ def test_get_client_ssl_credentials_auto_enablement( mock_ssl_channel_credentials.assert_called_once_with( certificate_chain=PUBLIC_CERT_BYTES, private_key=PRIVATE_KEY_BYTES ) + + +@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._MTLSCallInterceptor() + + 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._MTLSCallInterceptor._should_retry") +def test_factory_infinite_replay_on_error(mock_should_retry): + import google.auth.transport.grpc as transport_grpc + + interceptor = transport_grpc._MTLSCallInterceptor() + interceptor._wrapper = mock.Mock() + interceptor._wrapper._cached_cert = "cert" + mock_should_retry.side_effect = [True, False] + + 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) + + old_channel.unsubscribe.assert_called_once_with(subscriber) + new_channel.subscribe.assert_called_once_with(subscriber) + old_channel.close.assert_called_once() + + +@mock.patch("google.auth.transport.grpc._MTLSCallInterceptor._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._MTLSCallInterceptor() + interceptor._wrapper = mock.Mock() + interceptor._wrapper._cached_cert = "cert" + mock_should_retry.return_value = True + + 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.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) From b4271282e7052fc6d1a1d5bbee95816ec3310af1 Mon Sep 17 00:00:00 2001 From: agrawalradhika-cell Date: Thu, 6 Aug 2026 13:51:20 -0700 Subject: [PATCH 09/10] chore: Update unit tests for grpc cert rotation handling compatibility --- .../tests_async/transport/test_aiohttp_requests.py | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) 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): From a9e45afdaf427cb9cdf8bdb6a6fc19d2eeaa1902 Mon Sep 17 00:00:00 2001 From: agrawalradhika-cell Date: Thu, 6 Aug 2026 14:13:28 -0700 Subject: [PATCH 10/10] fix: Refactor gRPC call handling and state management fix: Refactor gRPC call handling and state management --- .../google-auth/google/auth/transport/grpc.py | 19 +++++++++++++++---- 1 file changed, 15 insertions(+), 4 deletions(-) diff --git a/packages/google-auth/google/auth/transport/grpc.py b/packages/google-auth/google/auth/transport/grpc.py index f6e8fdc72dc9..1ba2251cc9c0 100644 --- a/packages/google-auth/google/auth/transport/grpc.py +++ b/packages/google-auth/google/auth/transport/grpc.py @@ -655,7 +655,7 @@ def __init__( def _start_call(self): self._attempt_cert = ( self._interceptor._wrapper._cached_cert - if getattr(self._interceptor, "_wrapper", None) + if self._interceptor._wrapper else None ) @@ -689,6 +689,17 @@ def _on_inner_future_done(self, inner_future): if self._target_future 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() @@ -713,9 +724,9 @@ def _on_inner_future_done(self, inner_future): except Exception as e: self._terminal_exception = e - if getattr(self._interceptor, "_wrapper", None): + if self._interceptor._wrapper: if self._interceptor._should_retry( - status_code, 0, getattr(self, "_attempt_cert", None) + status_code, 0, self._attempt_cert ): self._interceptor._wrapper.refresh_logic(1) @@ -1000,7 +1011,7 @@ def running(self): def done(self): with self._lock: - return getattr(self, "_is_completed", False) + return self._is_completed def initial_metadata(self): with self._lock: