From fa24e36e8f248ada2ef8bdc88565268d76dfe7ff Mon Sep 17 00:00:00 2001 From: agrawalradhika-cell Date: Thu, 6 Aug 2026 13:39:55 -0700 Subject: [PATCH 01/26] 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/26] 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/26] 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/26] 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/26] 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/26] 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/26] 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/26] 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/26] 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/26] 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: From 524247bc318700c98cf67f2b79b4775a6024dc4d Mon Sep 17 00:00:00 2001 From: Radhika Agrawal Date: Mon, 10 Aug 2026 20:46:23 +0000 Subject: [PATCH 11/26] chore: format files with black --- packages/google-auth/google/auth/transport/grpc.py | 4 +--- packages/google-auth/google/auth/transport/requests.py | 2 +- .../google-auth/tests/transport/test__mtls_helper.py | 9 ++++++++- packages/google-auth/tests/transport/test_urllib3.py | 8 +++++++- 4 files changed, 17 insertions(+), 6 deletions(-) diff --git a/packages/google-auth/google/auth/transport/grpc.py b/packages/google-auth/google/auth/transport/grpc.py index 1ba2251cc9c0..2ed12ec3d460 100644 --- a/packages/google-auth/google/auth/transport/grpc.py +++ b/packages/google-auth/google/auth/transport/grpc.py @@ -725,9 +725,7 @@ def _on_inner_future_done(self, inner_future): self._terminal_exception = e if self._interceptor._wrapper: - if self._interceptor._should_retry( - status_code, 0, self._attempt_cert - ): + if self._interceptor._should_retry(status_code, 0, self._attempt_cert): self._interceptor._wrapper.refresh_logic(1) with self._lock: diff --git a/packages/google-auth/google/auth/transport/requests.py b/packages/google-auth/google/auth/transport/requests.py index 4ac32461c467..3aaa2e02a516 100644 --- a/packages/google-auth/google/auth/transport/requests.py +++ b/packages/google-auth/google/auth/transport/requests.py @@ -658,7 +658,7 @@ def request( ( call_cert_bytes, call_key_bytes, - _, # passphrase is not processed by requests adapter + _, # 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/tests/transport/test__mtls_helper.py b/packages/google-auth/tests/transport/test__mtls_helper.py index 537ef47e7295..1bf5b9abf276 100644 --- a/packages/google-auth/tests/transport/test__mtls_helper.py +++ b/packages/google-auth/tests/transport/test__mtls_helper.py @@ -972,6 +972,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" @@ -983,6 +984,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( @@ -991,6 +993,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() @@ -1006,6 +1009,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" @@ -1014,12 +1018,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() @@ -1034,10 +1040,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_urllib3.py b/packages/google-auth/tests/transport/test_urllib3.py index a01242f05b9f..fcdb6e7da099 100644 --- a/packages/google-auth/tests/transport/test_urllib3.py +++ b/packages/google-auth/tests/transport/test_urllib3.py @@ -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, None, "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 ea7082ba959c13cf538386bea5e16f1e57b4e75b Mon Sep 17 00:00:00 2001 From: Radhika Agrawal Date: Tue, 11 Aug 2026 01:19:41 +0000 Subject: [PATCH 12/26] chore: add coverage and fix lint errors --- .../google-auth/google/auth/transport/grpc.py | 1 - .../google-auth/tests/transport/test_grpc.py | 38 +++++++++++++++++++ 2 files changed, 38 insertions(+), 1 deletion(-) diff --git a/packages/google-auth/google/auth/transport/grpc.py b/packages/google-auth/google/auth/transport/grpc.py index 2ed12ec3d460..81249dc744be 100644 --- a/packages/google-auth/google/auth/transport/grpc.py +++ b/packages/google-auth/google/auth/transport/grpc.py @@ -17,7 +17,6 @@ from __future__ import absolute_import import collections -import concurrent.futures import logging import threading import time diff --git a/packages/google-auth/tests/transport/test_grpc.py b/packages/google-auth/tests/transport/test_grpc.py index f78c13ce848a..215c91739f27 100644 --- a/packages/google-auth/tests/transport/test_grpc.py +++ b/packages/google-auth/tests/transport/test_grpc.py @@ -837,3 +837,41 @@ def callback(f): assert len(callbacks_fired) == 1 with pytest.raises(transport_grpc.grpc.RpcError): future.result(timeout=1) + + +@mock.patch("google.auth.transport.grpc._MTLSCallInterceptor._should_retry") +def test_unary_response_future_cancelled(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 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.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 From 60a606abc780089934c9777c5c9bf13ffc90383b Mon Sep 17 00:00:00 2001 From: Radhika Agrawal Date: Tue, 11 Aug 2026 05:32:01 +0000 Subject: [PATCH 13/26] chore: add coverage for grpc.py --- .../google-auth/tests/transport/test_grpc.py | 35 +++++++++++++++++++ 1 file changed, 35 insertions(+) diff --git a/packages/google-auth/tests/transport/test_grpc.py b/packages/google-auth/tests/transport/test_grpc.py index 215c91739f27..77a8a7ede654 100644 --- a/packages/google-auth/tests/transport/test_grpc.py +++ b/packages/google-auth/tests/transport/test_grpc.py @@ -875,3 +875,38 @@ def failing_callback(f): assert future.done() is True assert len(callbacks_fired) == 1 + +@mock.patch("google.auth.transport.grpc._MTLSCallInterceptor._should_retry") +def test_unary_response_future_no_retry_but_refresh(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 RpcError exception for the future + mock_err = transport_grpc.grpc.RpcError() + mock_err.code = lambda: transport_grpc.grpc.StatusCode.UNAUTHENTICATED + + inner_future = mock.Mock() + inner_future.cancelled.return_value = False + inner_future.exception.return_value = mock_err + + mock_should_retry.return_value = True + + # Set can_replay=False by mocking the payload's can_replay logic + payload = mock.Mock() + payload.can_replay.return_value = False + + call_details = mock.Mock() + future = transport_grpc._RetryableUnaryResponseFuture( + continuation=lambda cd, pl: inner_future, + client_call_details=call_details, + request_or_iterator=payload, + interceptor=interceptor, + is_client_stream=True, + ) + + future._on_inner_future_done(inner_future) + + interceptor._wrapper.refresh_logic.assert_called_once_with(1) From 90555fccaf7e956d881eb790e257915253bf70f3 Mon Sep 17 00:00:00 2001 From: Radhika Agrawal Date: Tue, 11 Aug 2026 15:33:11 +0000 Subject: [PATCH 14/26] test: Updating tests for file coverage --- .../google-auth/tests/transport/test_grpc.py | 372 +++++++++++++++++- 1 file changed, 362 insertions(+), 10 deletions(-) diff --git a/packages/google-auth/tests/transport/test_grpc.py b/packages/google-auth/tests/transport/test_grpc.py index 77a8a7ede654..b55a5ce1aa4e 100644 --- a/packages/google-auth/tests/transport/test_grpc.py +++ b/packages/google-auth/tests/transport/test_grpc.py @@ -17,6 +17,7 @@ import time from unittest import mock +import grpc import pytest # type: ignore from google.auth import _helpers @@ -37,7 +38,6 @@ def unwrap(ch): try: # pylint: disable=ungrouped-imports - import grpc # type: ignore import google.auth.transport.grpc @@ -824,6 +824,7 @@ def callback(f): interceptor=interceptor, is_client_stream=False, ) + future._completion_event.set() future.add_done_callback(callback) # Allow time to elapse so remaining timeout <= 0 @@ -868,6 +869,7 @@ def failing_callback(f): interceptor=interceptor, is_client_stream=False, ) + future._completion_event.set() future.add_done_callback(callback) future.add_done_callback(failing_callback) @@ -876,37 +878,387 @@ def failing_callback(f): assert future.done() is True assert len(callbacks_fired) == 1 + @mock.patch("google.auth.transport.grpc._MTLSCallInterceptor._should_retry") -def test_unary_response_future_no_retry_but_refresh(mock_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._MTLSCallInterceptor() interceptor._wrapper = mock.Mock() interceptor._wrapper._cached_cert = "cert" - # Mock RpcError exception for the future mock_err = transport_grpc.grpc.RpcError() mock_err.code = lambda: transport_grpc.grpc.StatusCode.UNAUTHENTICATED + mock_should_retry.return_value = True 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._MTLSCallInterceptor() + 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._MTLSCallInterceptor() + 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._MTLSCallInterceptor._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._MTLSCallInterceptor() + interceptor._wrapper = None + + mock_err = transport_grpc.grpc.RpcError() + mock_err.code = lambda: transport_grpc.grpc.StatusCode.UNAUTHENTICATED mock_should_retry.return_value = True - # Set can_replay=False by mocking the payload's can_replay logic - payload = mock.Mock() - payload.can_replay.return_value = False + 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=payload, + request_or_iterator=b"request", interceptor=interceptor, - is_client_stream=True, + is_client_stream=False, ) + future._completion_event.set() - future._on_inner_future_done(inner_future) + with mock.patch.object(future, "_start_call", side_effect=mock_err): + future._on_inner_future_done(inner_future) + + +@mock.patch("google.auth.transport.grpc._MTLSCallInterceptor._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._MTLSCallInterceptor() + 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 + + 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._MTLSCallInterceptor() + + assert ( + interceptor._should_retry(grpc.StatusCode.UNAUTHENTICATED, 0, "cert1") is False + ) + + wrapper_mock = mock.Mock() + wrapper_mock._cached_cert = "cert1" + interceptor._wrapper = wrapper_mock + assert interceptor._should_retry(grpc.StatusCode.INTERNAL, 0, "cert1") is False + assert ( + interceptor._should_retry(grpc.StatusCode.UNAUTHENTICATED, 2, "cert1") is False + ) + + wrapper_mock._cached_cert = "cert2" + assert ( + interceptor._should_retry(grpc.StatusCode.UNAUTHENTICATED, 0, "cert1") is True + ) + + 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") + is True + ) + + 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") + is False + ) + + +def test_mtls_call_interceptor_interceptors_methods(): + from unittest import mock + import google.auth.transport.grpc as transport_grpc + + interceptor = transport_grpc._MTLSCallInterceptor() + + 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) is None + mock_check.return_value = (None, None, None, "fp1", "fp2") + assert channel.refresh_logic(0) 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) 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._MTLSCallInterceptor() + 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._MTLSCallInterceptor() + interceptor._wrapper = mock.Mock() + interceptor._wrapper._cached_cert = "cert" + + iterator = transport_grpc._RetryableStreamResponseIterator( + dummy_continuation, mock.Mock(), mock.Mock(), interceptor + ) - interceptor._wrapper.refresh_logic.assert_called_once_with(1) + 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) From 3139e83675aff3f5900393d2d60e044195ba390d Mon Sep 17 00:00:00 2001 From: Radhika Agrawal Date: Wed, 12 Aug 2026 20:35:27 +0000 Subject: [PATCH 15/26] fix: fix lint errors post resolving merge conflicts Signed-off-by: Radhika Agrawal --- packages/google-auth/tests/transport/test_grpc.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/packages/google-auth/tests/transport/test_grpc.py b/packages/google-auth/tests/transport/test_grpc.py index d1c5bd4c43d5..7990c9ae838f 100644 --- a/packages/google-auth/tests/transport/test_grpc.py +++ b/packages/google-auth/tests/transport/test_grpc.py @@ -691,6 +691,7 @@ def test_get_client_ssl_credentials_auto_enablement( 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 @@ -1264,6 +1265,7 @@ def dummy_continuation(*args, **kwargs): 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( @@ -1284,4 +1286,3 @@ def test_grpc_version_warning_not_emitted_when_no_version(monkeypatch): with warnings.catch_warnings(): warnings.simplefilter("error", FutureWarning) importlib.reload(google.auth.transport.grpc) - From f344dcf4aa52b3cfed1632273f73d41eef06a1ad Mon Sep 17 00:00:00 2001 From: agrawalradhika-cell Date: Thu, 20 Aug 2026 13:25:17 -0700 Subject: [PATCH 16/26] chore: Update Refactor mTLS gRPC client interceptor logic based on comments and fix doctrings chore: Update Refactor mTLS gRPC client interceptor logic based on comments and fix doctrings --- .../google-auth/google/auth/transport/grpc.py | 195 ++++++++++-------- 1 file changed, 110 insertions(+), 85 deletions(-) diff --git a/packages/google-auth/google/auth/transport/grpc.py b/packages/google-auth/google/auth/transport/grpc.py index 82bd4d85499c..524d7c20c508 100644 --- a/packages/google-auth/google/auth/transport/grpc.py +++ b/packages/google-auth/google/auth/transport/grpc.py @@ -24,6 +24,7 @@ from google.auth import exceptions +from google.auth import transport from google.auth.transport import _mtls_helper from google.auth.transport import mtls from google.oauth2 import service_account @@ -309,18 +310,18 @@ def my_client_cert_callback(): composite_credentials = grpc.composite_channel_credentials( ssl_credentials, google_auth_credentials ) - is_retry = kwargs.pop("_is_retry", False) + is_recreation = kwargs.pop("_is_recreation", 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 + # 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 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 + "_is_recreation": True, # Hidden flag to stop recursion **kwargs, } interceptor = _MTLSCallInterceptor() @@ -406,33 +407,48 @@ class _MTLSCallInterceptor( 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): self._wrapper = None - self._max_retries = 2 # Set your desired limit here + 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 + return False, None, None, None if retry_count >= self._max_retries: _LOGGER.debug( - "Max retries reached (%d/%d).", retry_count, self._max_retries + "Max retries reached (%d/%d) for channel recreation.", retry_count, self._max_retries ) - return False + return False, None, None, None - # If the wrapper has already rotated to a new cert, we can retry immediately + # If another thread already refreshed the channel with an updated cert, retry immediately if attempt_cert != self._wrapper._cached_cert: - return True + return True, None, None, None - # Fingerprint check logic + # Check if the certificate on disk or callback has changed since this request was attempted ( - _, - _, - _, - cached_fp, - current_fp, + call_cert_bytes, + call_key_bytes, + passphrase, + cached_fingerprint, + current_cert_fingerprint, ) = _mtls_helper.check_parameters_for_unauthorized_response(attempt_cert) - return cached_fp != current_fp + 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( @@ -476,54 +492,37 @@ def __init__(self, target, factory_args, initial_channel, initial_cert): self._lock = threading.Lock() self._subscribers = set() - def refresh_logic(self, count): + def refresh_logic(self, count, call_cert_bytes=None, call_key_bytes=None, passphrase=None): 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 + if not call_cert_bytes or self._cached_cert == call_cert_bytes: + return - # Support encrypted keys - if passphrase is not None: - call_key_bytes = _mtls_helper.decrypt_private_key( - call_key_bytes, passphrase - ) + _LOGGER.debug( + "Wrapper: Refreshing mTLS channel. Retry count: %d", count + ) + old_channel = self._channel - # 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, + if passphrase is not None: + call_key_bytes = _mtls_helper.decrypt_private_key( + call_key_bytes, passphrase ) - self._channel = secure_authorized_channel(**factory_args) - - for callback in self._subscribers: - try: - old_channel.unsubscribe(callback) - except Exception: - pass - self._channel.subscribe(callback) + # 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) + self._cached_cert = call_cert_bytes + for callback in self._subscribers: try: - old_channel.close() + 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 @@ -555,7 +554,7 @@ def close(self): class _ReplayableIterator(object): def __init__(self, target_iterator, max_items=1000): - self._target_iterator = target_iterator + self._target_iterator = iter(target_iterator) self._max_items = max_items self._buffer = [] self._exhausted = False @@ -731,12 +730,18 @@ def _on_inner_future_done(self, inner_future): else (self._payload.can_replay() if self._is_client_stream else True) ) - if can_replay and self._interceptor._should_retry( + 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): - self._interceptor._wrapper.refresh_logic(1) - + 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: @@ -746,18 +751,23 @@ def _on_inner_future_done(self, inner_future): self._terminal_exception = e if self._interceptor._wrapper: - if self._interceptor._should_retry(status_code, 0, self._attempt_cert): - self._interceptor._wrapper.refresh_logic(1) - + 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: - pass + except Exception as e: + _LOGGER.warning("Callback failed: %s", e) def add_done_callback(self, fn): with self._lock: @@ -942,6 +952,11 @@ 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): @@ -973,32 +988,42 @@ def __next__(self): ) ) + 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 self._interceptor._should_retry( - status_code, - self._retry_count, - getattr(self, "_attempt_cert", None), - ) + and should_retry ): - if getattr(self._interceptor, "_wrapper", None): - self._interceptor._wrapper.refresh_logic(1) - - with self._lock: - self._retry_count += 1 - try: + 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 timeout_e: - self._trigger_callbacks() - raise timeout_e + + 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): - if self._interceptor._should_retry( + chk_should_retry, chk_cert, chk_key, chk_pwd = self._interceptor._should_retry( status_code, 0, getattr(self, "_attempt_cert", None) - ): - self._interceptor._wrapper.refresh_logic(1) + ) + 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 From b7db29e26875cdbf7fa6d256ab2d0cb7eb10259b Mon Sep 17 00:00:00 2001 From: agrawalradhika-cell Date: Thu, 20 Aug 2026 13:40:15 -0700 Subject: [PATCH 17/26] chore: Refactor error handling for certificate retrieval and update type hint for passphrase. chore: Refactor error handling for certificate retrieval and update type hint for passphrase. --- packages/google-auth/google/auth/transport/_mtls_helper.py | 7 ++----- 1 file changed, 2 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 03ca956d7f2d..a469aa4e4ba3 100644 --- a/packages/google-auth/google/auth/transport/_mtls_helper.py +++ b/packages/google-auth/google/auth/transport/_mtls_helper.py @@ -642,10 +642,7 @@ def get_client_ssl_credentials( """ # 1. Attempt to retrieve X.509 Workload cert and key. - try: - cert, key = _get_workload_cert_and_key(certificate_config_path) - except exceptions.ClientCertError: - cert, key = None, None + cert, key = _get_workload_cert_and_key(certificate_config_path) if cert and key: return True, cert, key, None @@ -820,7 +817,7 @@ 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. + Optional[Union[bytes, str]]: The passphrase for the key. str: The base64-encoded SHA256 cached fingerprint. str: The base64-encoded SHA256 current cert fingerprint. """ From 91570a309207171c6455100ed6116a5e229a8b5e Mon Sep 17 00:00:00 2001 From: agrawalradhika-cell Date: Thu, 20 Aug 2026 13:54:06 -0700 Subject: [PATCH 18/26] fix: Refactor deadline error handling in grpc.py fix: Refactor deadline error handling in grpc.py --- .../google-auth/google/auth/transport/grpc.py | 16 +++++++++++----- 1 file changed, 11 insertions(+), 5 deletions(-) diff --git a/packages/google-auth/google/auth/transport/grpc.py b/packages/google-auth/google/auth/transport/grpc.py index 524d7c20c508..4bb91e926611 100644 --- a/packages/google-auth/google/auth/transport/grpc.py +++ b/packages/google-auth/google/auth/transport/grpc.py @@ -617,9 +617,7 @@ def __next__(self): if self._parent._can_replay: self._parent._buffer.append(val) - if len(self._parent._buffer) > getattr( - self._parent, "_max_items", 10000 - ): + if len(self._parent._buffer) > self._parent._max_items: self._parent._buffer.clear() self._parent._can_replay = False @@ -632,6 +630,14 @@ def __next__(self): ("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(grpc.Future, grpc.Call): def __init__( @@ -692,7 +698,7 @@ def _start_call(self): elapsed = time.monotonic() - self._start_time remaining = self._initial_timeout - elapsed if remaining <= 0: - raise grpc.RpcError("Deadline Exceeded during retry resolution.") + raise _DeadlineExceededError("Deadline Exceeded during retry resolution.") call_details = _ClientCallDetails( method=call_details.method, timeout=remaining, @@ -923,7 +929,7 @@ def _start_call(self): elapsed = time.monotonic() - self._start_time remaining = self._initial_timeout - elapsed if remaining <= 0: - raise grpc.RpcError("Deadline Exceeded during retry resolution.") + raise _DeadlineExceededError("Deadline Exceeded during retry resolution.") call_details = _ClientCallDetails( method=call_details.method, timeout=remaining, From d8050edfcfb0b7b9cdb9d563b0f6133d3b8f131c Mon Sep 17 00:00:00 2001 From: Radhika Agrawal Date: Thu, 20 Aug 2026 22:14:18 +0000 Subject: [PATCH 19/26] fix: Fix lint and unit tests Signed-off-by: Radhika Agrawal --- .../google-auth/google/auth/transport/grpc.py | 116 ++++++++------ .../google-auth/tests/transport/test_grpc.py | 142 +++++++++++------- 2 files changed, 163 insertions(+), 95 deletions(-) diff --git a/packages/google-auth/google/auth/transport/grpc.py b/packages/google-auth/google/auth/transport/grpc.py index 4bb91e926611..3451eb8292d7 100644 --- a/packages/google-auth/google/auth/transport/grpc.py +++ b/packages/google-auth/google/auth/transport/grpc.py @@ -24,7 +24,7 @@ from google.auth import exceptions -from google.auth import transport +from google.auth import transport from google.auth.transport import _mtls_helper from google.auth.transport import mtls from google.oauth2 import service_account @@ -324,9 +324,9 @@ def my_client_cert_callback(): "_is_recreation": True, # Hidden flag to stop recursion **kwargs, } - interceptor = _MTLSCallInterceptor() + interceptor = CertRotationInterceptor() - wrapper = _MTLSRefreshingChannel(target, factory_args, channel, cached_cert) + wrapper = MTLSRefreshingChannel(target, factory_args, channel, cached_cert) interceptor._wrapper = wrapper return grpc.intercept_channel(wrapper, interceptor) @@ -401,19 +401,19 @@ def is_mtls(self): return self._is_mtls -class _MTLSCallInterceptor( +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 + + 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. """ @@ -423,15 +423,17 @@ def __init__(self): 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). - """ + + 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 + "Max retries reached (%d/%d) for channel recreation.", + retry_count, + self._max_retries, ) return False, None, None, None @@ -483,7 +485,7 @@ def intercept_stream_stream( ) -class _MTLSRefreshingChannel(grpc.Channel): +class MTLSRefreshingChannel(grpc.Channel): def __init__(self, target, factory_args, initial_channel, initial_cert): self._target = target self._factory_args = factory_args @@ -492,14 +494,14 @@ def __init__(self, target, factory_args, initial_channel, initial_cert): self._lock = threading.Lock() self._subscribers = set() - def refresh_logic(self, count, call_cert_bytes=None, call_key_bytes=None, passphrase=None): + 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 - ) + _LOGGER.debug("Wrapper: Refreshing mTLS channel. Retry count: %d", count) old_channel = self._channel if passphrase is not None: @@ -630,15 +632,19 @@ def __next__(self): ("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(grpc.Future, grpc.Call): def __init__( self, @@ -698,7 +704,9 @@ def _start_call(self): elapsed = time.monotonic() - self._start_time remaining = self._initial_timeout - elapsed if remaining <= 0: - raise _DeadlineExceededError("Deadline Exceeded during retry resolution.") + raise _DeadlineExceededError( + "Deadline Exceeded during retry resolution." + ) call_details = _ClientCallDetails( method=call_details.method, timeout=remaining, @@ -742,7 +750,9 @@ def _on_inner_future_done(self, inner_future): if can_replay and should_retry: if getattr(self._interceptor, "_wrapper", None): try: - self._interceptor._wrapper.refresh_logic(1, call_cert, call_key, pwd) + self._interceptor._wrapper.refresh_logic( + 1, call_cert, call_key, pwd + ) except Exception as e: with self._lock: self._terminal_exception = e @@ -757,18 +767,23 @@ def _on_inner_future_done(self, inner_future): 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 - ) + ( + 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) + 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) @@ -929,7 +944,9 @@ def _start_call(self): elapsed = time.monotonic() - self._start_time remaining = self._initial_timeout - elapsed if remaining <= 0: - raise _DeadlineExceededError("Deadline Exceeded during retry resolution.") + raise _DeadlineExceededError( + "Deadline Exceeded during retry resolution." + ) call_details = _ClientCallDetails( method=call_details.method, timeout=remaining, @@ -959,9 +976,12 @@ def _on_inner_call_done(self, inner_call): if self._call is not inner_call: return # Intercept and suppress premature callbacks for UNAUTHENTICATED. - # __next__ inherently handles this error and manages triggering callbacks + # __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: + if ( + callable(getattr(inner_call, "code", None)) + and inner_call.code() == grpc.StatusCode.UNAUTHENTICATED + ): return self._trigger_callbacks() @@ -994,42 +1014,52 @@ def __next__(self): ) ) - should_retry, call_cert, call_key, pwd = self._interceptor._should_retry( + ( + 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 - ): + 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) - + 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( + ( + 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) + self._interceptor._wrapper.refresh_logic( + 1, chk_cert, chk_key, chk_pwd + ) except Exception: - pass # Terminal anyway - + pass # Terminal anyway + self._trigger_callbacks() raise e diff --git a/packages/google-auth/tests/transport/test_grpc.py b/packages/google-auth/tests/transport/test_grpc.py index 7990c9ae838f..108839717412 100644 --- a/packages/google-auth/tests/transport/test_grpc.py +++ b/packages/google-auth/tests/transport/test_grpc.py @@ -696,7 +696,7 @@ def test_get_client_ssl_credentials_auto_enablement( def test_interceptor_uses_factory_if_callable(mock_replayable): import google.auth.transport.grpc as transport_grpc - interceptor = transport_grpc._MTLSCallInterceptor() + interceptor = transport_grpc.CertRotationInterceptor() call_no_factory = transport_grpc._RetryableStreamResponseIterator( continuation=mock.Mock(), @@ -722,14 +722,17 @@ def generator_factory(): assert call_factory._payload is None -@mock.patch("google.auth.transport.grpc._MTLSCallInterceptor._should_retry") +@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._MTLSCallInterceptor() + interceptor = transport_grpc.CertRotationInterceptor() interceptor._wrapper = mock.Mock() interceptor._wrapper._cached_cert = "cert" - mock_should_retry.side_effect = [True, False] + 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() @@ -779,7 +782,7 @@ def test_refresh_logic_closes_old_channel( subscriber = mock.Mock() - refreshing_channel = transport_grpc._MTLSRefreshingChannel( + refreshing_channel = transport_grpc.MTLSRefreshingChannel( target="example.com:443", factory_args={}, initial_channel=old_channel, @@ -787,21 +790,23 @@ def test_refresh_logic_closes_old_channel( ) refreshing_channel.subscribe(subscriber) - refreshing_channel.refresh_logic(1) + 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() + # old_channel.close.assert_called_once() # Removed in PR 18019 -@mock.patch("google.auth.transport.grpc._MTLSCallInterceptor._should_retry") +@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._MTLSCallInterceptor() + interceptor = transport_grpc.CertRotationInterceptor() interceptor._wrapper = mock.Mock() interceptor._wrapper._cached_cert = "cert" - mock_should_retry.return_value = True + 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 @@ -842,11 +847,11 @@ def callback(f): future.result(timeout=1) -@mock.patch("google.auth.transport.grpc._MTLSCallInterceptor._should_retry") +@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._MTLSCallInterceptor() + interceptor = transport_grpc.CertRotationInterceptor() interceptor._wrapper = mock.Mock() interceptor._wrapper._cached_cert = "cert" @@ -881,17 +886,17 @@ def failing_callback(f): assert len(callbacks_fired) == 1 -@mock.patch("google.auth.transport.grpc._MTLSCallInterceptor._should_retry") +@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._MTLSCallInterceptor() + 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 + mock_should_retry.return_value = (True, b"cert", b"key", None) inner_future = mock.Mock() inner_future.cancelled.return_value = False @@ -938,7 +943,7 @@ def test_start_call_wrapper_none(): import pytest import google.auth.transport.grpc as transport_grpc - interceptor = transport_grpc._MTLSCallInterceptor() + interceptor = transport_grpc.CertRotationInterceptor() if hasattr(interceptor, "_wrapper"): del interceptor._wrapper @@ -958,7 +963,7 @@ def test_start_call_wrapper_none(): def test_start_call_wrapper_none_branch(): import google.auth.transport.grpc as transport_grpc - interceptor = transport_grpc._MTLSCallInterceptor() + interceptor = transport_grpc.CertRotationInterceptor() interceptor._wrapper = None inner_future = mock.Mock() @@ -975,16 +980,16 @@ def test_start_call_wrapper_none_branch(): assert getattr(future, "_attempt_cert", "NOT_SET") is None -@mock.patch("google.auth.transport.grpc._MTLSCallInterceptor._should_retry") +@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._MTLSCallInterceptor() + 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 + mock_should_retry.return_value = (True, b"cert", b"key", None) inner_future = mock.Mock() inner_future.cancelled.return_value = False @@ -1005,17 +1010,17 @@ def test_unary_response_future_rpc_error_no_wrapper(mock_should_retry): future._on_inner_future_done(inner_future) -@mock.patch("google.auth.transport.grpc._MTLSCallInterceptor._should_retry") +@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._MTLSCallInterceptor() + 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 + mock_should_retry.return_value = (False, None, None, None) inner_future = mock.Mock() inner_future.cancelled.return_value = False @@ -1042,23 +1047,37 @@ def test_mtls_call_interceptor_should_retry_cases(): from unittest import mock import google.auth.transport.grpc as transport_grpc - interceptor = transport_grpc._MTLSCallInterceptor() + interceptor = transport_grpc.CertRotationInterceptor() - assert ( - interceptor._should_retry(grpc.StatusCode.UNAUTHENTICATED, 0, "cert1") is False + 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") is False - assert ( - interceptor._should_retry(grpc.StatusCode.UNAUTHENTICATED, 2, "cert1") is False + 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") is True + assert interceptor._should_retry(grpc.StatusCode.UNAUTHENTICATED, 0, "cert1") == ( + True, + None, + None, + None, ) wrapper_mock._cached_cert = "cert1" @@ -1066,26 +1085,24 @@ def test_mtls_call_interceptor_should_retry_cases(): "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") - is True - ) + 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") - is False - ) + 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._MTLSCallInterceptor() + interceptor = transport_grpc.CertRotationInterceptor() def dummy_continuation(*args, **kwargs): return mock.Mock() @@ -1118,14 +1135,28 @@ 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") + 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) is None + 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) is None + 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 @@ -1138,7 +1169,14 @@ def dummy_callback(): 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) is None + 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 @@ -1147,7 +1185,7 @@ 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") + channel = transport_grpc.MTLSRefreshingChannel("target", {}, mock.Mock(), b"cert1") cb = mock.Mock() channel.subscribe(cb) assert cb in channel._subscribers @@ -1162,14 +1200,14 @@ 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") + 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") + channel = transport_grpc.MTLSRefreshingChannel("target", {}, mock.Mock(), b"cert1") res = channel.unary_unary("method") assert res is not None @@ -1178,7 +1216,7 @@ 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") + channel = transport_grpc.MTLSRefreshingChannel("target", {}, mock.Mock(), b"cert1") res = channel.unary_stream("method") assert res is not None @@ -1187,7 +1225,7 @@ 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") + channel = transport_grpc.MTLSRefreshingChannel("target", {}, mock.Mock(), b"cert1") res = channel.stream_unary("method") assert res is not None @@ -1196,7 +1234,7 @@ 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") + channel = transport_grpc.MTLSRefreshingChannel("target", {}, mock.Mock(), b"cert1") res = channel.stream_stream("method") assert res is not None @@ -1210,7 +1248,7 @@ def test_retryable_unary_response_future_methods(): def dummy_continuation(*args, **kwargs): return mock_future - interceptor = transport_grpc._MTLSCallInterceptor() + interceptor = transport_grpc.CertRotationInterceptor() interceptor._wrapper = mock.Mock() interceptor._wrapper._cached_cert = "cert" @@ -1247,7 +1285,7 @@ def test_retryable_stream_response_iterator_methods(): def dummy_continuation(*args, **kwargs): return mock_iterator - interceptor = transport_grpc._MTLSCallInterceptor() + interceptor = transport_grpc.CertRotationInterceptor() interceptor._wrapper = mock.Mock() interceptor._wrapper._cached_cert = "cert" From 02fa5ea217155e3f3e9f8b9f9b5a51b1bd0925c9 Mon Sep 17 00:00:00 2001 From: agrawalradhika-cell Date: Tue, 25 Aug 2026 11:50:33 -0700 Subject: [PATCH 20/26] chore: Refactor gRPC call handling with base wrapper class chore: Refactor gRPC call handling with base wrapper class --- .../google-auth/google/auth/transport/grpc.py | 105 +++++++----------- 1 file changed, 43 insertions(+), 62 deletions(-) diff --git a/packages/google-auth/google/auth/transport/grpc.py b/packages/google-auth/google/auth/transport/grpc.py index 3451eb8292d7..45f165055e7c 100644 --- a/packages/google-auth/google/auth/transport/grpc.py +++ b/packages/google-auth/google/auth/transport/grpc.py @@ -29,6 +29,9 @@ 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 @@ -288,7 +291,7 @@ def my_client_cert_callback(): ) # If SSL credentials are not explicitly set, try client_cert_callback and ADC. - cached_cert = None + 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: @@ -645,7 +648,7 @@ def details(self): return self._details -class _RetryableUnaryResponseFuture(grpc.Future, grpc.Call): +class _RetryableUnaryResponseFuture(_BaseCallWrapper): def __init__( self, continuation, @@ -672,6 +675,7 @@ def __init__( ) self._retry_count = 0 + self._call = None self._lock = threading.RLock() timeout = getattr(self._client_call_details, "timeout", None) @@ -715,12 +719,12 @@ def _start_call(self): 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) + 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._target_future is not inner_future: + if self._call is not inner_future: return if inner_future.cancelled(): @@ -810,7 +814,7 @@ def result(self, timeout=None): with self._lock: if self._terminal_exception is not None: raise self._terminal_exception - current_future = self._target_future + current_future = self._call return current_future.result() def exception(self, timeout=None): @@ -819,7 +823,7 @@ def exception(self, timeout=None): with self._lock: if self._terminal_exception is not None: return self._terminal_exception - return self._target_future.exception() + return self._call.exception() def traceback(self, timeout=None): if not self._completion_event.wait(timeout): @@ -827,66 +831,37 @@ def traceback(self, timeout=None): with self._lock: if self._terminal_exception is not None: return self._terminal_exception.__traceback__ - return self._target_future.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._target_future.initial_metadata() + 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._target_future.trailing_metadata() + 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._target_future.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._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) - + return self._call.details() -class _RetryableStreamResponseIterator(grpc.Call): +class _RetryableStreamResponseIterator(_BaseCallWrapper): def __init__( self, continuation, @@ -911,7 +886,7 @@ def __init__( else request_or_iterator ) ) - + self._call = None self._retry_count = 0 self._yielded_any_response = False self._lock = threading.RLock() @@ -1077,40 +1052,46 @@ def add_done_callback(self, fn): 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): - with self._lock: - return self._call.cancel() + return self._call.cancel() def cancelled(self): - with self._lock: - return self._call.cancelled() + return self._call.cancelled() def running(self): - with self._lock: - return self._call.running() + return self._call.running() def done(self): - with self._lock: - return self._is_completed + 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): - with self._lock: - return self._call.initial_metadata() + return self._call.initial_metadata() def trailing_metadata(self): - with self._lock: - return self._call.trailing_metadata() + return self._call.trailing_metadata() def code(self): - with self._lock: - return self._call.code() + return self._call.code() def details(self): - with self._lock: - return self._call.details() - - def is_active(self): - return self._call.is_active() + return self._call.details() def time_remaining(self): return self._call.time_remaining() From 9b01e859ea508a500756ae63b2f9004958b4ca80 Mon Sep 17 00:00:00 2001 From: agrawalradhika-cell Date: Tue, 25 Aug 2026 15:02:08 -0700 Subject: [PATCH 21/26] chore: Include Wrapper in CertRotationInterceptor initialization chore: Include Wrapper in CertRotationInterceptor initialization --- packages/google-auth/google/auth/transport/grpc.py | 9 +++------ 1 file changed, 3 insertions(+), 6 deletions(-) diff --git a/packages/google-auth/google/auth/transport/grpc.py b/packages/google-auth/google/auth/transport/grpc.py index 45f165055e7c..873ae3598db0 100644 --- a/packages/google-auth/google/auth/transport/grpc.py +++ b/packages/google-auth/google/auth/transport/grpc.py @@ -327,11 +327,8 @@ def my_client_cert_callback(): "_is_recreation": True, # Hidden flag to stop recursion **kwargs, } - interceptor = CertRotationInterceptor() - wrapper = MTLSRefreshingChannel(target, factory_args, channel, cached_cert) - - interceptor._wrapper = wrapper + interceptor = CertRotationInterceptor(wrapper=wrapper) return grpc.intercept_channel(wrapper, interceptor) return channel @@ -420,8 +417,8 @@ class CertRotationInterceptor( channel's credentials and automatically replays the failed RPC. """ - def __init__(self): - self._wrapper = None + 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): From 07ab346f37225dad7de3d2ff892ec77ab534c9a4 Mon Sep 17 00:00:00 2001 From: agrawalradhika-cell Date: Tue, 25 Aug 2026 21:47:19 -0700 Subject: [PATCH 22/26] chore: Refactor MTLS channel creation to use a partial function for channel factory arguments. chore: Refactor MTLS channel creation to use a partial function for channel factory arguments. --- .../google-auth/google/auth/transport/grpc.py | 34 +++++++++++-------- 1 file changed, 19 insertions(+), 15 deletions(-) diff --git a/packages/google-auth/google/auth/transport/grpc.py b/packages/google-auth/google/auth/transport/grpc.py index 873ae3598db0..2f49abb9c59b 100644 --- a/packages/google-auth/google/auth/transport/grpc.py +++ b/packages/google-auth/google/auth/transport/grpc.py @@ -17,6 +17,7 @@ from __future__ import absolute_import import collections +import functools import logging import threading import time @@ -318,16 +319,15 @@ def my_client_cert_callback(): # 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 - factory_args = { - "credentials": credentials, - "request": request, - "target": target, - "ssl_credentials": None, - "client_cert_callback": client_cert_callback, - "_is_recreation": True, # Hidden flag to stop recursion - **kwargs, - } - wrapper = MTLSRefreshingChannel(target, factory_args, channel, cached_cert) + create_channel_fn = functools.partial( + secure_authorized_channel, + credentials=credentials, + request=request, + target=target, + _is_recreation=True, # Hidden flag to stop recursion + **kwargs + ) + wrapper = MTLSRefreshingChannel(target, create_channel_fn, channel, cached_cert) interceptor = CertRotationInterceptor(wrapper=wrapper) return grpc.intercept_channel(wrapper, interceptor) return channel @@ -488,7 +488,7 @@ def intercept_stream_stream( class MTLSRefreshingChannel(grpc.Channel): def __init__(self, target, factory_args, initial_channel, initial_cert): self._target = target - self._factory_args = factory_args + self._create_channel_fn = create_channel_fn self._channel = initial_channel self._cached_cert = initial_cert self._lock = threading.Lock() @@ -509,14 +509,18 @@ def refresh_logic( 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( + # 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._channel = secure_authorized_channel(**factory_args) self._cached_cert = call_cert_bytes for callback in self._subscribers: From 4a0063180c56f559c4e5d90580ff912083a40be6 Mon Sep 17 00:00:00 2001 From: agrawalradhika-cell Date: Tue, 25 Aug 2026 22:09:04 -0700 Subject: [PATCH 23/26] chore: Add mTLS Interceptor and Wrapper in separate file for certificate rotation This file implements an mTLS Interceptor and Channel Wrapper that supports automatic retry logic for certificate rotation in gRPC calls. --- .../auth/transport/_mtls_interceptor.py | 410 ++++++++++++++++++ 1 file changed, 410 insertions(+) create mode 100644 packages/google-auth/google/auth/transport/_mtls_interceptor.py 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() From 9091ec39b62269b76429f9e56fd2af799ee51e6f Mon Sep 17 00:00:00 2001 From: agrawalradhika-cell Date: Wed, 26 Aug 2026 07:49:25 -0700 Subject: [PATCH 24/26] chore: Refactor gRPC transport to use mtls_interceptor Updated the gRPC transport module to use mtls_interceptor for mutual TLS support. Removed the old CertRotationInterceptor and MTLSRefreshingChannel classes. --- .../google-auth/google/auth/transport/grpc.py | 708 +----------------- 1 file changed, 3 insertions(+), 705 deletions(-) diff --git a/packages/google-auth/google/auth/transport/grpc.py b/packages/google-auth/google/auth/transport/grpc.py index 2f49abb9c59b..3d831ab4efdd 100644 --- a/packages/google-auth/google/auth/transport/grpc.py +++ b/packages/google-auth/google/auth/transport/grpc.py @@ -16,17 +16,15 @@ from __future__ import absolute_import -import collections import functools import logging -import threading -import time 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.py‎ from google.auth.transport import mtls from google.oauth2 import service_account @@ -327,8 +325,8 @@ def my_client_cert_callback(): _is_recreation=True, # Hidden flag to stop recursion **kwargs ) - wrapper = MTLSRefreshingChannel(target, create_channel_fn, channel, cached_cert) - interceptor = CertRotationInterceptor(wrapper=wrapper) + 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 @@ -399,703 +397,3 @@ def ssl_credentials(self): def is_mtls(self): """Indicates if the created SSL channel credentials is mutual TLS.""" return self._is_mtls - - -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, factory_args, 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._channel = secure_authorized_channel(**factory_args) - 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) From 4305ab754fe171aa7d26c6403c13f30c3a421412 Mon Sep 17 00:00:00 2001 From: agrawalradhika-cell Date: Wed, 26 Aug 2026 07:50:03 -0700 Subject: [PATCH 25/26] chore: Rename _mtls_interceptor.py to mtls_interceptor.py to make public chore: Rename _mtls_interceptor.py to mtls_interceptor.py to make public --- .../auth/transport/{_mtls_interceptor.py => mtls_interceptor.py} | 0 1 file changed, 0 insertions(+), 0 deletions(-) rename packages/google-auth/google/auth/transport/{_mtls_interceptor.py => mtls_interceptor.py} (100%) diff --git a/packages/google-auth/google/auth/transport/_mtls_interceptor.py b/packages/google-auth/google/auth/transport/mtls_interceptor.py similarity index 100% rename from packages/google-auth/google/auth/transport/_mtls_interceptor.py rename to packages/google-auth/google/auth/transport/mtls_interceptor.py From f7a2ca3f2832a3b242ae6cbdc1647cf204f03941 Mon Sep 17 00:00:00 2001 From: agrawalradhika-cell Date: Wed, 26 Aug 2026 07:50:58 -0700 Subject: [PATCH 26/26] fix: Fix import statement for mtls_interceptor Fix import statement for mtls_interceptor --- packages/google-auth/google/auth/transport/grpc.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/packages/google-auth/google/auth/transport/grpc.py b/packages/google-auth/google/auth/transport/grpc.py index 3d831ab4efdd..a856e82ad2ac 100644 --- a/packages/google-auth/google/auth/transport/grpc.py +++ b/packages/google-auth/google/auth/transport/grpc.py @@ -24,7 +24,7 @@ from google.auth import exceptions from google.auth import transport from google.auth.transport import _mtls_helper -from google.auth.transport import mtls_interceptor.py‎ +from google.auth.transport import mtls_interceptor‎ from google.auth.transport import mtls from google.oauth2 import service_account