From d7afecc31455dac717b9603e30707cb910206039 Mon Sep 17 00:00:00 2001 From: Radhika Agrawal Date: Fri, 2 Oct 2026 21:20:16 +0000 Subject: [PATCH 01/12] feat(auth): add CertRotationInterceptor and MTLSRefreshingChannel for gRPC mTLS Signed-off-by: Radhika Agrawal --- .../google/auth/transport/mtls_interceptor.py | 724 ++++++++++++ .../tests/transport/test_mtls_interceptor.py | 1039 +++++++++++++++++ 2 files changed, 1763 insertions(+) create mode 100644 packages/google-auth/google/auth/transport/mtls_interceptor.py create mode 100644 packages/google-auth/tests/transport/test_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..b39b45d9b8fa --- /dev/null +++ b/packages/google-auth/google/auth/transport/mtls_interceptor.py @@ -0,0 +1,724 @@ +# Copyright 2016 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""mTLS Interceptor and Channel Wrapper for certificate rotation.""" + +import collections +import logging +import threading +import time + +import grpc + +from google.auth import transport +from google.auth.transport import _mtls_helper + +_LOGGER = logging.getLogger(__name__) + + +class CertRotationInterceptor( + grpc.UnaryUnaryClientInterceptor, + grpc.UnaryStreamClientInterceptor, + grpc.StreamUnaryClientInterceptor, + grpc.StreamStreamClientInterceptor, +): + """A gRPC client interceptor that provides automatic retry logic for mTLS certificate rotation. + + This interceptor wraps all gRPC client calls (unary and streaming) with retryable + futures or iterators. Its primary role is to monitor responses for `UNAUTHENTICATED` + errors. When an authentication failure occurs, it uses `_should_retry()` to check + if a new mTLS certificate is available. If a new certificate is found, it signals + its associated `MTLSRefreshingChannel` wrapper to refresh the underlying gRPC + channel's credentials and automatically replays the failed RPC. + """ + + def __init__(self, wrapper=None): + self._wrapper = wrapper + self._max_retries = transport.DEFAULT_MAX_REFRESH_ATTEMPTS + + def _should_retry(self, code, retry_count, attempt_cert): + """Determines if the RPC should be retried due to a certificate rotation. + + Returns a tuple: (should_retry, call_cert_bytes, call_key_bytes). + """ + if code != grpc.StatusCode.UNAUTHENTICATED or not self._wrapper: + return False, 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 + + # If another thread already refreshed the channel with an updated cert, retry immediately + if attempt_cert != self._wrapper._cached_cert: + return True, None, None + + # Check if the certificate on disk or callback has changed since this request was attempted + ( + call_cert_bytes, + call_key_bytes, + 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 + + def intercept_unary_unary(self, continuation, client_call_details, request): + return _RetryableUnaryResponseFuture( + continuation, client_call_details, request, self, is_client_stream=False + ) + + def intercept_stream_unary( + self, continuation, client_call_details, request_iterator + ): + return _RetryableUnaryResponseFuture( + continuation, + client_call_details, + request_iterator, + self, + is_client_stream=True, + ) + + def intercept_unary_stream(self, continuation, client_call_details, request): + return _RetryableStreamResponseIterator( + continuation, client_call_details, request, self, is_client_stream=False + ) + + def intercept_stream_stream( + self, continuation, client_call_details, request_iterator + ): + return _RetryableStreamResponseIterator( + continuation, + client_call_details, + request_iterator, + self, + is_client_stream=True, + ) + + +class MTLSRefreshingChannel(grpc.Channel): + def __init__(self, target, create_channel_fn, initial_channel, initial_cert): + self._target = target + self._create_channel_fn = create_channel_fn + self._channel = initial_channel + self._cached_cert = initial_cert + self._lock = threading.Lock() + self._subscribers = set() + + def refresh_logic(self, count, call_cert_bytes=None, call_key_bytes=None): + 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 + + # Call the partial, overriding only the cert-related arguments + new_ssl_credentials = grpc.ssl_channel_credentials( + certificate_chain=call_cert_bytes, + private_key=call_key_bytes, + ) + + self._channel = self._create_channel_fn( + ssl_credentials=new_ssl_credentials, client_cert_callback=None + ) + + self._cached_cert = call_cert_bytes + for callback in self._subscribers: + try: + old_channel.unsubscribe(callback) + except Exception: + pass + self._channel.subscribe(callback) + + def unary_unary(self, method, *args, **kwargs): + # Always return a callable from the CURRENT channel + return self._channel.unary_unary(method, *args, **kwargs) + + # Mandatory passthroughs + def unary_stream(self, method, *args, **kwargs): + return self._channel.unary_stream(method, *args, **kwargs) + + def stream_unary(self, method, *args, **kwargs): + return self._channel.stream_unary(method, *args, **kwargs) + + def stream_stream(self, method, *args, **kwargs): + return self._channel.stream_stream(method, *args, **kwargs) + + def subscribe(self, callback, try_to_connect=False): + with self._lock: + self._subscribers.add(callback) + return self._channel.subscribe(callback, try_to_connect=try_to_connect) + + def unsubscribe(self, callback): + with self._lock: + self._subscribers.discard(callback) + return self._channel.unsubscribe(callback) + + def close(self): + self._channel.close() + + +class _ReplayableIterator(object): + def __init__(self, target_iterator, max_items=1000): + self._target_iterator = iter(target_iterator) + self._max_items = max_items + self._buffer = [] + self._exhausted = False + self._can_replay = True + + self._lock = threading.Lock() + self._consumer_lock = threading.Lock() + self._active_reader = None + + def __iter__(self): + reader = _ReplayableIteratorReader(self) + with self._lock: + self._active_reader = reader + return reader + + def can_replay(self): + with self._lock: + return self._can_replay + + +class _ReplayableIteratorReader(object): + def __init__(self, parent): + self._parent = parent + self._read_index = 0 + + def __iter__(self): + return self + + 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 _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 is_active(self): + return self._call.is_active() + + def add_callback(self, callback): + self._call.add_callback(callback) + + +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 + + status_code = None + 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 = 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 + ) + 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, + ) = 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) + 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, + ) = 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 + ) + + 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, + ) = 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 + ) + 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 diff --git a/packages/google-auth/tests/transport/test_mtls_interceptor.py b/packages/google-auth/tests/transport/test_mtls_interceptor.py new file mode 100644 index 000000000000..a5c24eb32467 --- /dev/null +++ b/packages/google-auth/tests/transport/test_mtls_interceptor.py @@ -0,0 +1,1039 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from unittest import mock + +import pytest # type: ignore + +from google.auth import transport + +try: + import grpc # type: ignore + + from google.auth.transport import mtls_interceptor + + HAS_GRPC = True +except ImportError: # pragma: NO COVER + HAS_GRPC = False + +pytestmark = pytest.mark.skipif(not HAS_GRPC, reason="gRPC is unavailable.") + +CHECK_PARAMS = ( + "google.auth.transport._mtls_helper.check_parameters_for_unauthorized_response" +) +SHOULD_RETRY = ( + "google.auth.transport.mtls_interceptor.CertRotationInterceptor._should_retry" +) + + +# --------------------------------------------------------------------------- +# Test doubles +# --------------------------------------------------------------------------- + + +class FakeRpcError(grpc.RpcError, grpc.Call, grpc.Future): + """An RpcError that, like grpc's _InactiveRpcError, is also a finished call.""" + + def __init__(self, code, details="error"): + super().__init__() + self._code = code + self._details = details + + def code(self): + return self._code + + def details(self): + return self._details + + def initial_metadata(self): + return None + + def trailing_metadata(self): + return None + + def is_active(self): + return False + + def time_remaining(self): + return None + + def add_callback(self, callback): + return False + + def cancel(self): + return False + + def cancelled(self): + return False + + def running(self): + return False + + def done(self): + return True + + def result(self, timeout=None): + raise self + + def exception(self, timeout=None): + return self + + def traceback(self, timeout=None): + return None + + def add_done_callback(self, fn): + fn(self) + + +class CompletedFuture(object): + """An inner unary future that is already finished. + + This mirrors what grpc's blocking ``__call__`` / ``with_call`` path hands to + an interceptor: ``add_done_callback`` runs the callback synchronously. + """ + + def __init__(self, result=None, exception=None, cancelled=False): + self._result = result + self._exception = exception + self._cancelled = cancelled + + def add_done_callback(self, fn): + fn(self) + + def cancelled(self): + return self._cancelled + + def exception(self, timeout=None): + return self._exception + + def result(self, timeout=None): + if self._exception is not None: + raise self._exception + return self._result + + def code(self): + if self._exception is not None and hasattr(self._exception, "code"): + return self._exception.code() + return grpc.StatusCode.OK + + def details(self): + return "details" + + def initial_metadata(self): + return ("initial", "md") + + def trailing_metadata(self): + return ("trailing", "md") + + def traceback(self, timeout=None): + return None + + +class FakeStreamCall(object): + """An inner response-streaming call backed by a list of items/exceptions.""" + + def __init__(self, items, code=grpc.StatusCode.OK): + self._items = list(items) + self._code = code + self.done_callbacks = [] + + def __iter__(self): + return self + + def __next__(self): + if not self._items: + raise StopIteration() + item = self._items.pop(0) + if isinstance(item, BaseException): + raise item + return item + + def add_done_callback(self, fn): + self.done_callbacks.append(fn) + + def fire_done(self): + for fn in self.done_callbacks: + fn(self) + + def code(self): + return self._code + + +def make_interceptor(cached_cert=b"old-cert"): + wrapper = mock.Mock(spec=["_cached_cert", "refresh_logic"]) + wrapper._cached_cert = cached_cert + return mtls_interceptor.CertRotationInterceptor(wrapper=wrapper), wrapper + + +def call_details(timeout=None): + return mtls_interceptor._ClientCallDetails( + method="/svc/Method", + timeout=timeout, + metadata=None, + credentials=None, + wait_for_ready=None, + ) + + +def unauthenticated(): + return FakeRpcError(grpc.StatusCode.UNAUTHENTICATED, "cert mismatch") + + +# --------------------------------------------------------------------------- +# CertRotationInterceptor +# --------------------------------------------------------------------------- + + +class TestCertRotationInterceptor(object): + def test_init_defaults(self): + interceptor = mtls_interceptor.CertRotationInterceptor() + assert interceptor._wrapper is None + assert interceptor._max_retries == transport.DEFAULT_MAX_REFRESH_ATTEMPTS + + def test_init_with_wrapper(self): + wrapper = mock.Mock() + interceptor = mtls_interceptor.CertRotationInterceptor(wrapper=wrapper) + assert interceptor._wrapper is wrapper + + def test_should_retry_without_wrapper(self): + interceptor = mtls_interceptor.CertRotationInterceptor() + assert interceptor._should_retry( + grpc.StatusCode.UNAUTHENTICATED, 0, b"cert" + ) == (False, None, None) + + @pytest.mark.parametrize( + "code", [None, grpc.StatusCode.OK, grpc.StatusCode.INTERNAL] + ) + def test_should_retry_non_unauthenticated(self, code): + interceptor, _ = make_interceptor() + with mock.patch(CHECK_PARAMS) as check: + assert interceptor._should_retry(code, 0, b"old-cert") == ( + False, + None, + None, + ) + check.assert_not_called() + + def test_should_retry_max_retries_reached(self): + interceptor, _ = make_interceptor() + with mock.patch(CHECK_PARAMS) as check: + assert interceptor._should_retry( + grpc.StatusCode.UNAUTHENTICATED, + transport.DEFAULT_MAX_REFRESH_ATTEMPTS, + b"old-cert", + ) == (False, None, None) + check.assert_not_called() + + def test_should_retry_channel_already_refreshed_by_other_thread(self): + interceptor, _ = make_interceptor(cached_cert=b"new-cert") + with mock.patch(CHECK_PARAMS) as check: + assert interceptor._should_retry( + grpc.StatusCode.UNAUTHENTICATED, 0, b"old-cert" + ) == (True, None, None) + check.assert_not_called() + + def test_should_retry_cert_rotated_returns_new_material(self): + interceptor, _ = make_interceptor() + with mock.patch(CHECK_PARAMS) as check: + check.return_value = (b"new-cert", b"new-key", "fp1", "fp2") + assert interceptor._should_retry( + grpc.StatusCode.UNAUTHENTICATED, 0, b"old-cert" + ) == (True, b"new-cert", b"new-key") + check.assert_called_once_with(b"old-cert") + + def test_should_retry_cert_not_rotated(self): + interceptor, _ = make_interceptor() + with mock.patch(CHECK_PARAMS) as check: + check.return_value = (b"old-cert", b"key", "fp1", "fp1") + should_retry, *_ = interceptor._should_retry( + grpc.StatusCode.UNAUTHENTICATED, 0, b"old-cert" + ) + assert should_retry is False + + def test_intercept_methods_return_wrappers(self): + interceptor = mtls_interceptor.CertRotationInterceptor() + + def unary_continuation(details, request): + return CompletedFuture(result="ok") + + def stream_continuation(details, request): + return FakeStreamCall([]) + + assert isinstance( + interceptor.intercept_unary_unary(unary_continuation, call_details(), "r"), + mtls_interceptor._RetryableUnaryResponseFuture, + ) + assert isinstance( + interceptor.intercept_stream_unary( + unary_continuation, call_details(), iter(["r"]) + ), + mtls_interceptor._RetryableUnaryResponseFuture, + ) + assert isinstance( + interceptor.intercept_unary_stream( + stream_continuation, call_details(), "r" + ), + mtls_interceptor._RetryableStreamResponseIterator, + ) + assert isinstance( + interceptor.intercept_stream_stream( + stream_continuation, call_details(), iter(["r"]) + ), + mtls_interceptor._RetryableStreamResponseIterator, + ) + + +# --------------------------------------------------------------------------- +# MTLSRefreshingChannel +# --------------------------------------------------------------------------- + + +class TestMTLSRefreshingChannel(object): + def make_channel(self): + old_channel = mock.Mock() + new_channel = mock.Mock() + create_channel_fn = mock.Mock(return_value=new_channel) + channel = mtls_interceptor.MTLSRefreshingChannel( + "example.com:443", create_channel_fn, old_channel, b"old-cert" + ) + return channel, create_channel_fn, old_channel, new_channel + + @pytest.mark.parametrize("cert", [None, b"", b"old-cert"]) + def test_refresh_logic_noop_without_new_cert(self, cert): + channel, create_channel_fn, old_channel, _ = self.make_channel() + channel.refresh_logic(1, call_cert_bytes=cert, call_key_bytes=b"key") + create_channel_fn.assert_not_called() + assert channel._channel is old_channel + assert channel._cached_cert == b"old-cert" + + @mock.patch("grpc.ssl_channel_credentials", autospec=True) + def test_refresh_logic_rebuilds_channel(self, ssl_channel_credentials): + channel, create_channel_fn, old_channel, new_channel = self.make_channel() + subscriber = mock.Mock() + channel.subscribe(subscriber) + + channel.refresh_logic(1, call_cert_bytes=b"new-cert", call_key_bytes=b"new-key") + + ssl_channel_credentials.assert_called_once_with( + certificate_chain=b"new-cert", private_key=b"new-key" + ) + create_channel_fn.assert_called_once_with( + ssl_credentials=ssl_channel_credentials.return_value, + client_cert_callback=None, + ) + assert channel._channel is new_channel + assert channel._cached_cert == b"new-cert" + # Subscribers move to the new channel. + old_channel.unsubscribe.assert_called_once_with(subscriber) + new_channel.subscribe.assert_called_once_with(subscriber) + # The old channel is not closed so in-flight RPCs on it can finish. + old_channel.close.assert_not_called() + + @mock.patch("grpc.ssl_channel_credentials", autospec=True) + def test_refresh_logic_ignores_unsubscribe_errors(self, ssl_channel_credentials): + channel, _, old_channel, new_channel = self.make_channel() + subscriber = mock.Mock() + channel.subscribe(subscriber) + old_channel.unsubscribe.side_effect = ValueError("already gone") + + channel.refresh_logic(1, call_cert_bytes=b"new-cert", call_key_bytes=b"k") + + new_channel.subscribe.assert_called_once_with(subscriber) + + @mock.patch("grpc.ssl_channel_credentials", autospec=True) + def test_refresh_logic_failure_keeps_cached_cert(self, ssl_channel_credentials): + channel, create_channel_fn, old_channel, _ = self.make_channel() + create_channel_fn.side_effect = RuntimeError("cannot build channel") + + with pytest.raises(RuntimeError): + channel.refresh_logic(1, call_cert_bytes=b"new-cert", call_key_bytes=b"k") + + # State is not half-updated, so a later 401 can still trigger a refresh. + assert channel._cached_cert == b"old-cert" + assert channel._channel is old_channel + + @pytest.mark.parametrize( + "method", ["unary_unary", "unary_stream", "stream_unary", "stream_stream"] + ) + def test_multicallables_use_current_channel(self, method): + channel, _, old_channel, new_channel = self.make_channel() + assert ( + getattr(channel, method)("/svc/M") + is getattr(old_channel, method).return_value + ) + + channel._channel = new_channel + assert ( + getattr(channel, method)("/svc/M", "extra", kw=1) + is getattr(new_channel, method).return_value + ) + getattr(new_channel, method).assert_called_once_with("/svc/M", "extra", kw=1) + + def test_subscribe_unsubscribe_close(self): + channel, _, old_channel, _ = self.make_channel() + callback = mock.Mock() + + channel.subscribe(callback, try_to_connect=True) + assert callback in channel._subscribers + old_channel.subscribe.assert_called_once_with(callback, try_to_connect=True) + + channel.unsubscribe(callback) + assert callback not in channel._subscribers + old_channel.unsubscribe.assert_called_once_with(callback) + + channel.close() + old_channel.close.assert_called_once_with() + + +# --------------------------------------------------------------------------- +# _ReplayableIterator +# --------------------------------------------------------------------------- + + +class TestReplayableIterator(object): + def test_accepts_plain_iterables(self): + replayable = mtls_interceptor._ReplayableIterator([b"a", b"b"]) + assert list(iter(replayable)) == [b"a", b"b"] + + def test_replays_buffered_items_on_new_reader(self): + replayable = mtls_interceptor._ReplayableIterator(x for x in [b"a", b"b"]) + first = iter(replayable) + assert next(first) == b"a" + + second = iter(replayable) + assert list(second) == [b"a", b"b"] + assert replayable.can_replay() + + def test_stale_reader_stops(self): + replayable = mtls_interceptor._ReplayableIterator([b"a", b"b"]) + first = iter(replayable) + iter(replayable) # A newer reader takes over. + with pytest.raises(StopIteration): + next(first) + + def test_disables_replay_after_max_items(self): + replayable = mtls_interceptor._ReplayableIterator(range(5), max_items=2) + assert list(iter(replayable)) == [0, 1, 2, 3, 4] + assert replayable.can_replay() is False + + +# --------------------------------------------------------------------------- +# _DeadlineExceededError +# --------------------------------------------------------------------------- + + +def test_deadline_exceeded_error_implements_call(): + error = mtls_interceptor._DeadlineExceededError("too late") + assert isinstance(error, grpc.RpcError) + assert isinstance(error, grpc.Call) + assert error.code() == grpc.StatusCode.DEADLINE_EXCEEDED + assert error.details() == "too late" + + +# --------------------------------------------------------------------------- +# _RetryableUnaryResponseFuture +# --------------------------------------------------------------------------- + + +class TestRetryableUnaryResponseFuture(object): + def test_success_with_already_completed_call(self): + # Regression: status_code was unbound on success, raising UnboundLocalError + # out of intercept_unary_unary on grpc's blocking call path. + interceptor, wrapper = make_interceptor() + callback = mock.Mock() + + future = interceptor.intercept_unary_unary( + lambda details, request: CompletedFuture(result="response"), + call_details(), + "request", + ) + future.add_done_callback(callback) + + assert future.result(timeout=1) == "response" + assert future.exception(timeout=1) is None + assert future.code() == grpc.StatusCode.OK + callback.assert_called_once_with(future) + wrapper.refresh_logic.assert_not_called() + + def test_non_rpc_error_is_surfaced(self): + interceptor, wrapper = make_interceptor() + error = ValueError("boom") + + future = interceptor.intercept_unary_unary( + lambda details, request: CompletedFuture(exception=error), + call_details(), + "request", + ) + + with pytest.raises(ValueError): + future.result(timeout=1) + assert future.exception(timeout=1) is error + wrapper.refresh_logic.assert_not_called() + + def test_cancelled_call_completes_and_fires_callbacks(self): + interceptor, _ = make_interceptor() + inner = CompletedFuture(cancelled=True) + calls = [] + future = mtls_interceptor._RetryableUnaryResponseFuture.__new__( + mtls_interceptor._RetryableUnaryResponseFuture + ) + # Build normally, but with a call whose callback we trigger manually. + pending = mock.Mock() + future = mtls_interceptor._RetryableUnaryResponseFuture( + lambda details, request: pending, call_details(), "request", interceptor + ) + future.add_done_callback(calls.append) + future.add_done_callback(mock.Mock(side_effect=RuntimeError("ignored"))) + future._call = inner + + future._on_inner_future_done(inner) + + assert future._completion_event.is_set() + assert calls == [future] + + def test_callback_for_stale_call_is_ignored(self): + interceptor, _ = make_interceptor() + pending = mock.Mock() + future = mtls_interceptor._RetryableUnaryResponseFuture( + lambda details, request: pending, call_details(), "request", interceptor + ) + + future._on_inner_future_done(CompletedFuture(result="stale")) + + assert not future._completion_event.is_set() + + def test_retries_after_cert_rotation(self): + interceptor, wrapper = make_interceptor() + continuation = mock.Mock( + side_effect=[ + CompletedFuture(exception=unauthenticated()), + CompletedFuture(result="response"), + ] + ) + + with mock.patch(CHECK_PARAMS) as check: + check.return_value = (b"new-cert", b"new-key", "fp1", "fp2") + future = interceptor.intercept_unary_unary( + continuation, call_details(), "request" + ) + + assert future.result(timeout=1) == "response" + assert continuation.call_count == 2 + assert future._retry_count == 1 + # The cert material from the fingerprint check is reused, not re-fetched. + wrapper.refresh_logic.assert_called_once_with(1, b"new-cert", b"new-key") + check.assert_called_once() + + def test_no_retry_when_cert_unchanged(self): + interceptor, wrapper = make_interceptor() + error = unauthenticated() + continuation = mock.Mock(return_value=CompletedFuture(exception=error)) + + with mock.patch(CHECK_PARAMS) as check: + check.return_value = (b"old-cert", b"key", "fp1", "fp1") + future = interceptor.intercept_unary_unary( + continuation, call_details(), "request" + ) + + with pytest.raises(FakeRpcError): + future.result(timeout=1) + assert future.code() == grpc.StatusCode.UNAUTHENTICATED + assert continuation.call_count == 1 + wrapper.refresh_logic.assert_not_called() + + def test_stops_after_max_retries(self): + interceptor, wrapper = make_interceptor() + continuation = mock.Mock( + side_effect=lambda details, request: CompletedFuture( + exception=unauthenticated() + ) + ) + + with mock.patch(SHOULD_RETRY, autospec=True) as should_retry: + should_retry.side_effect = lambda self, code, count, cert: ( + (count < 2, b"c", b"k") + ) + future = interceptor.intercept_unary_unary( + continuation, call_details(), "request" + ) + + with pytest.raises(FakeRpcError): + future.result(timeout=1) + assert continuation.call_count == 3 + + def test_refresh_failure_becomes_terminal_exception(self): + interceptor, wrapper = make_interceptor() + wrapper.refresh_logic.side_effect = RuntimeError("refresh failed") + + with mock.patch(SHOULD_RETRY) as should_retry: + should_retry.return_value = (True, b"c", b"k") + future = interceptor.intercept_unary_unary( + lambda details, request: CompletedFuture(exception=unauthenticated()), + call_details(), + "request", + ) + + assert future._completion_event.is_set() + with pytest.raises(RuntimeError, match="refresh failed"): + future.result(timeout=1) + assert isinstance(future.exception(timeout=1), RuntimeError) + assert future.initial_metadata() is None + assert future.trailing_metadata() is None + + def test_deadline_exceeded_during_retry(self): + interceptor, wrapper = make_interceptor() + continuation = mock.Mock( + return_value=CompletedFuture(exception=unauthenticated()) + ) + + with ( + mock.patch(SHOULD_RETRY) as should_retry, + mock.patch.object( + mtls_interceptor.time, "monotonic", side_effect=[1000.0, 1000.0, 2000.0] + ), + ): + should_retry.side_effect = [ + (True, b"c", b"k"), + (False, None, None), + ] + future = interceptor.intercept_unary_unary( + continuation, call_details(timeout=5.0), "request" + ) + + with pytest.raises(mtls_interceptor._DeadlineExceededError): + future.result(timeout=1) + assert future.code() == grpc.StatusCode.DEADLINE_EXCEEDED + assert future.details() == "Deadline Exceeded during retry resolution." + assert continuation.call_count == 1 + + def test_retry_uses_remaining_timeout(self): + interceptor, _ = make_interceptor() + continuation = mock.Mock( + side_effect=[ + CompletedFuture(exception=unauthenticated()), + CompletedFuture(result="response"), + ] + ) + + with ( + mock.patch(SHOULD_RETRY) as should_retry, + mock.patch.object( + mtls_interceptor.time, "monotonic", side_effect=[1000.0, 1000.0, 1002.0] + ), + ): + should_retry.return_value = (True, b"c", b"k") + future = interceptor.intercept_unary_unary( + continuation, call_details(timeout=5.0), "request" + ) + + assert future.result(timeout=1) == "response" + retry_details = continuation.call_args_list[1][0][0] + assert retry_details.timeout == pytest.approx(3.0) + + def test_client_stream_is_replayed_on_retry(self): + interceptor, _ = make_interceptor() + seen = [] + + def continuation(details, request_iterator): + seen.append(list(request_iterator)) + if len(seen) == 1: + return CompletedFuture(exception=unauthenticated()) + return CompletedFuture(result="response") + + with mock.patch(SHOULD_RETRY) as should_retry: + should_retry.return_value = (True, b"c", b"k") + future = interceptor.intercept_stream_unary( + continuation, call_details(), iter([b"a", b"b"]) + ) + + assert future.result(timeout=1) == "response" + assert seen == [[b"a", b"b"], [b"a", b"b"]] + + def test_client_stream_not_retried_when_replay_disabled(self): + interceptor, wrapper = make_interceptor() + continuation = mock.Mock( + return_value=CompletedFuture(exception=unauthenticated()) + ) + + with ( + mock.patch.object( + mtls_interceptor._ReplayableIterator, "can_replay", return_value=False + ), + mock.patch(SHOULD_RETRY) as should_retry, + ): + should_retry.return_value = (True, b"c", b"k") + future = interceptor.intercept_stream_unary( + continuation, call_details(), iter([b"a"]) + ) + + with pytest.raises(FakeRpcError): + future.result(timeout=1) + assert continuation.call_count == 1 + + def test_request_factory_is_called_per_attempt(self): + interceptor, _ = make_interceptor() + factory = mock.Mock(side_effect=lambda: iter([b"a"])) + continuation = mock.Mock( + side_effect=[ + CompletedFuture(exception=unauthenticated()), + CompletedFuture(result="response"), + ] + ) + + with mock.patch(SHOULD_RETRY) as should_retry: + should_retry.return_value = (True, b"c", b"k") + future = interceptor.intercept_stream_unary( + continuation, call_details(), factory + ) + + assert future._uses_factory is True + assert future.result(timeout=1) == "response" + assert factory.call_count == 2 + + def test_without_wrapper_does_not_retry(self): + interceptor = mtls_interceptor.CertRotationInterceptor() + future = interceptor.intercept_unary_unary( + lambda details, request: CompletedFuture(exception=unauthenticated()), + call_details(), + "request", + ) + + assert future._attempt_cert is None + with pytest.raises(FakeRpcError): + future.result(timeout=1) + + def test_result_times_out_while_pending(self): + interceptor, _ = make_interceptor() + future = interceptor.intercept_unary_unary( + lambda details, request: mock.Mock(), call_details(), "request" + ) + with pytest.raises(grpc.FutureTimeoutError): + future.result(timeout=0.01) + with pytest.raises(grpc.FutureTimeoutError): + future.exception(timeout=0.01) + with pytest.raises(grpc.FutureTimeoutError): + future.traceback(timeout=0.01) + + def test_add_done_callback_after_completion_fires_immediately(self): + interceptor, _ = make_interceptor() + future = interceptor.intercept_unary_unary( + lambda details, request: CompletedFuture(result="response"), + call_details(), + "request", + ) + callback = mock.Mock() + future.add_done_callback(callback) + callback.assert_called_once_with(future) + # Errors raised by late callbacks are swallowed. + future.add_done_callback(mock.Mock(side_effect=RuntimeError("ignored"))) + + def test_call_methods_delegate_to_current_call(self): + interceptor, _ = make_interceptor() + future = interceptor.intercept_unary_unary( + lambda details, request: CompletedFuture(result="response"), + call_details(), + "request", + ) + assert future.initial_metadata() == ("initial", "md") + assert future.trailing_metadata() == ("trailing", "md") + assert future.details() == "details" + assert future.traceback(timeout=1) is None + + +# --------------------------------------------------------------------------- +# _RetryableStreamResponseIterator +# --------------------------------------------------------------------------- + + +class TestRetryableStreamResponseIterator(object): + def test_yields_responses_and_fires_callbacks_once(self): + interceptor, _ = make_interceptor() + inner = FakeStreamCall([b"r1", b"r2"]) + stream = interceptor.intercept_unary_stream( + lambda details, request: inner, call_details(), "request" + ) + callback = mock.Mock() + stream.add_done_callback(callback) + + assert list(stream) == [b"r1", b"r2"] + inner.fire_done() + + callback.assert_called_once_with(stream) + # Late callbacks fire immediately. + late = mock.Mock() + stream.add_done_callback(late) + late.assert_called_once_with(stream) + + def test_retries_unauthenticated_before_first_response(self): + interceptor, wrapper = make_interceptor() + first = FakeStreamCall( + [unauthenticated()], code=grpc.StatusCode.UNAUTHENTICATED + ) + second = FakeStreamCall([b"r1"]) + continuation = mock.Mock(side_effect=[first, second]) + + with mock.patch(CHECK_PARAMS) as check: + check.return_value = (b"new-cert", b"new-key", "fp1", "fp2") + stream = interceptor.intercept_unary_stream( + continuation, call_details(), "request" + ) + assert list(stream) == [b"r1"] + + assert continuation.call_count == 2 + wrapper.refresh_logic.assert_called_once_with(1, b"new-cert", b"new-key") + + def test_unauthenticated_done_callback_is_suppressed(self): + # A 401 on the first attempt must not complete the outer stream while a + # retry may still follow. + interceptor, _ = make_interceptor() + first = FakeStreamCall( + [unauthenticated()], code=grpc.StatusCode.UNAUTHENTICATED + ) + stream = interceptor.intercept_unary_stream( + lambda details, request: first, call_details(), "request" + ) + callback = mock.Mock() + stream.add_done_callback(callback) + + first.fire_done() + + callback.assert_not_called() + assert stream._is_completed is False + + def test_non_unauthenticated_done_callback_fires(self): + interceptor, _ = make_interceptor() + inner = FakeStreamCall([], code=grpc.StatusCode.INTERNAL) + stream = interceptor.intercept_unary_stream( + lambda details, request: inner, call_details(), "request" + ) + callback = mock.Mock() + stream.add_done_callback(callback) + + inner.fire_done() + + callback.assert_called_once_with(stream) + + def test_no_retry_after_a_response_was_yielded(self): + interceptor, wrapper = make_interceptor() + error = unauthenticated() + inner = FakeStreamCall([b"r1", error]) + continuation = mock.Mock(return_value=inner) + + with mock.patch(SHOULD_RETRY) as should_retry: + should_retry.return_value = (True, b"c", b"k") + stream = interceptor.intercept_unary_stream( + continuation, call_details(), "request" + ) + assert next(stream) == b"r1" + with pytest.raises(FakeRpcError): + next(stream) + + assert continuation.call_count == 1 + + def test_terminal_error_fires_callbacks_and_raises(self): + interceptor, wrapper = make_interceptor() + error = FakeRpcError(grpc.StatusCode.INTERNAL) + stream = interceptor.intercept_unary_stream( + lambda details, request: FakeStreamCall([error]), + call_details(), + "request", + ) + callback = mock.Mock() + stream.add_done_callback(callback) + + with pytest.raises(FakeRpcError): + next(stream) + + callback.assert_called_once_with(stream) + wrapper.refresh_logic.assert_not_called() + + def test_refresh_failure_is_raised(self): + interceptor, wrapper = make_interceptor() + wrapper.refresh_logic.side_effect = RuntimeError("refresh failed") + + with mock.patch(SHOULD_RETRY) as should_retry: + should_retry.return_value = (True, b"c", b"k") + stream = interceptor.intercept_unary_stream( + lambda details, request: FakeStreamCall([unauthenticated()]), + call_details(), + "request", + ) + callback = mock.Mock() + stream.add_done_callback(callback) + with pytest.raises(RuntimeError, match="refresh failed"): + next(stream) + + callback.assert_called_once_with(stream) + + def test_deadline_exceeded_during_retry(self): + interceptor, _ = make_interceptor() + continuation = mock.Mock(return_value=FakeStreamCall([unauthenticated()])) + + with ( + mock.patch(SHOULD_RETRY) as should_retry, + mock.patch.object( + mtls_interceptor.time, "monotonic", side_effect=[1000.0, 1000.0, 2000.0] + ), + ): + should_retry.return_value = (True, b"c", b"k") + stream = interceptor.intercept_unary_stream( + continuation, call_details(timeout=5.0), "request" + ) + with pytest.raises(mtls_interceptor._DeadlineExceededError) as excinfo: + next(stream) + + assert excinfo.value.code() == grpc.StatusCode.DEADLINE_EXCEEDED + + def test_bidi_stream_replays_requests(self): + interceptor, _ = make_interceptor() + seen = [] + + def continuation(details, request_iterator): + seen.append(list(request_iterator)) + if len(seen) == 1: + return FakeStreamCall([unauthenticated()]) + return FakeStreamCall([b"r1"]) + + with mock.patch(SHOULD_RETRY) as should_retry: + should_retry.return_value = (True, b"c", b"k") + stream = interceptor.intercept_stream_stream( + continuation, call_details(), iter([b"a", b"b"]) + ) + assert list(stream) == [b"r1"] + + assert seen == [[b"a", b"b"], [b"a", b"b"]] + + def test_request_factory_is_called_per_attempt(self): + interceptor, _ = make_interceptor() + factory = mock.Mock(side_effect=lambda: iter([b"a"])) + continuation = mock.Mock( + side_effect=[FakeStreamCall([unauthenticated()]), FakeStreamCall([b"r1"])] + ) + + with mock.patch(SHOULD_RETRY) as should_retry: + should_retry.return_value = (True, b"c", b"k") + stream = interceptor.intercept_stream_stream( + continuation, call_details(), factory + ) + assert list(stream) == [b"r1"] + + assert stream._uses_factory is True + assert factory.call_count == 2 + + +# --------------------------------------------------------------------------- +# _BaseCallWrapper +# --------------------------------------------------------------------------- + + +@pytest.mark.parametrize( + "method, args", + [ + ("cancel", ()), + ("cancelled", ()), + ("running", ()), + ("done", ()), + ("initial_metadata", ()), + ("trailing_metadata", ()), + ("code", ()), + ("details", ()), + ("time_remaining", ()), + ("add_callback", (mock.sentinel.callback,)), + ("add_done_callback", (mock.sentinel.callback,)), + ], +) +def test_base_call_wrapper_delegates(method, args): + wrapper = mtls_interceptor._BaseCallWrapper() + wrapper._call = mock.Mock() + result = getattr(wrapper, method)(*args) + getattr(wrapper._call, method).assert_called_once_with(*args) + if method not in ("add_callback", "add_done_callback"): + assert result is getattr(wrapper._call, method).return_value + + +@pytest.mark.parametrize("method", ["result", "exception", "traceback"]) +def test_base_call_wrapper_delegates_with_timeout(method): + wrapper = mtls_interceptor._BaseCallWrapper() + wrapper._call = mock.Mock() + assert ( + getattr(wrapper, method)(timeout=3) + is getattr(wrapper._call, method).return_value + ) + getattr(wrapper._call, method).assert_called_once_with(timeout=3) + + +# --------------------------------------------------------------------------- +# End to end through grpc.intercept_channel +# --------------------------------------------------------------------------- + + +class _FakeUnaryMultiCallable(object): + def __init__(self, outcome): + self._outcome = outcome + self.calls = 0 + + def with_call(self, request, **kwargs): + self.calls += 1 + if isinstance(self._outcome, BaseException): + raise self._outcome + return self._outcome, CompletedFuture(result=self._outcome) + + +class _FakeChannel(object): + def __init__(self, outcome): + self.multicallable = _FakeUnaryMultiCallable(outcome) + + def unary_unary(self, method, *args, **kwargs): + return self.multicallable + + +@mock.patch("grpc.ssl_channel_credentials", autospec=True) +def test_blocking_call_recovers_from_cert_rotation(ssl_channel_credentials): + old_channel = _FakeChannel(unauthenticated()) + new_channel = _FakeChannel("response") + create_channel_fn = mock.Mock(return_value=new_channel) + refreshing = mtls_interceptor.MTLSRefreshingChannel( + "example.com:443", create_channel_fn, old_channel, b"old-cert" + ) + channel = grpc.intercept_channel( + refreshing, mtls_interceptor.CertRotationInterceptor(wrapper=refreshing) + ) + + with mock.patch(CHECK_PARAMS) as check: + check.return_value = (b"new-cert", b"new-key", "fp1", "fp2") + response = channel.unary_unary("/svc/Method")(b"request") + + assert response == "response" + assert old_channel.multicallable.calls == 1 + assert new_channel.multicallable.calls == 1 + assert refreshing._cached_cert == b"new-cert" + create_channel_fn.assert_called_once_with( + ssl_credentials=ssl_channel_credentials.return_value, + client_cert_callback=None, + ) + + +def test_blocking_call_success_without_rotation(): + old_channel = _FakeChannel("response") + refreshing = mtls_interceptor.MTLSRefreshingChannel( + "example.com:443", mock.Mock(), old_channel, b"old-cert" + ) + channel = grpc.intercept_channel( + refreshing, mtls_interceptor.CertRotationInterceptor(wrapper=refreshing) + ) + + with mock.patch(CHECK_PARAMS) as check: + assert channel.unary_unary("/svc/Method")(b"request") == "response" + + check.assert_not_called() From 5e48581acb811820debd4b19a72eeae8f479b445 Mon Sep 17 00:00:00 2001 From: agrawalradhika-cell Date: Mon, 5 Oct 2026 11:14:57 -0700 Subject: [PATCH 02/12] Update packages/google-auth/google/auth/transport/mtls_interceptor.py Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com> --- packages/google-auth/google/auth/transport/mtls_interceptor.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/packages/google-auth/google/auth/transport/mtls_interceptor.py b/packages/google-auth/google/auth/transport/mtls_interceptor.py index b39b45d9b8fa..a1eac7c5b952 100644 --- a/packages/google-auth/google/auth/transport/mtls_interceptor.py +++ b/packages/google-auth/google/auth/transport/mtls_interceptor.py @@ -422,7 +422,7 @@ def _on_inner_future_done(self, inner_future): ) should_retry, call_cert, call_key = self._interceptor._should_retry( - status_code, self._retry_count, getattr(self, "_attempt_cert", None) + status_code, self._retry_count, self._attempt_cert ) if can_replay and should_retry: if getattr(self._interceptor, "_wrapper", None): From fa3336a0dd9a1f70a9763ded9a33e346ba82ba2c Mon Sep 17 00:00:00 2001 From: agrawalradhika-cell Date: Mon, 5 Oct 2026 11:15:08 -0700 Subject: [PATCH 03/12] Update packages/google-auth/google/auth/transport/mtls_interceptor.py Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com> --- packages/google-auth/google/auth/transport/mtls_interceptor.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/packages/google-auth/google/auth/transport/mtls_interceptor.py b/packages/google-auth/google/auth/transport/mtls_interceptor.py index a1eac7c5b952..8c1e34fb90ea 100644 --- a/packages/google-auth/google/auth/transport/mtls_interceptor.py +++ b/packages/google-auth/google/auth/transport/mtls_interceptor.py @@ -696,7 +696,7 @@ def __next__(self): chk_cert, chk_key, ) = self._interceptor._should_retry( - status_code, 0, getattr(self, "_attempt_cert", None) + status_code, 0, self._attempt_cert ) if chk_should_retry: try: From 7f6a1e186f2eaff22e5d116746fb54987c75dc03 Mon Sep 17 00:00:00 2001 From: agrawalradhika-cell Date: Mon, 5 Oct 2026 11:15:24 -0700 Subject: [PATCH 04/12] Update packages/google-auth/google/auth/transport/mtls_interceptor.py Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com> --- packages/google-auth/google/auth/transport/mtls_interceptor.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/packages/google-auth/google/auth/transport/mtls_interceptor.py b/packages/google-auth/google/auth/transport/mtls_interceptor.py index 8c1e34fb90ea..4aa632003473 100644 --- a/packages/google-auth/google/auth/transport/mtls_interceptor.py +++ b/packages/google-auth/google/auth/transport/mtls_interceptor.py @@ -574,7 +574,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 ) with self._lock: From 402470c8aa4d543347088bb75be433736a68e40b Mon Sep 17 00:00:00 2001 From: agrawalradhika-cell Date: Mon, 5 Oct 2026 11:15:51 -0700 Subject: [PATCH 05/12] Update packages/google-auth/google/auth/transport/mtls_interceptor.py Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com> --- packages/google-auth/google/auth/transport/mtls_interceptor.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/packages/google-auth/google/auth/transport/mtls_interceptor.py b/packages/google-auth/google/auth/transport/mtls_interceptor.py index 4aa632003473..deab9521f94c 100644 --- a/packages/google-auth/google/auth/transport/mtls_interceptor.py +++ b/packages/google-auth/google/auth/transport/mtls_interceptor.py @@ -669,7 +669,7 @@ def __next__(self): ) = self._interceptor._should_retry( status_code, self._retry_count, - getattr(self, "_attempt_cert", None), + self._attempt_cert, ) if not self._yielded_any_response and can_replay and should_retry: From bb998ef90f2294f9088b086e9222b4a4454cf0a5 Mon Sep 17 00:00:00 2001 From: agrawalradhika-cell Date: Mon, 5 Oct 2026 11:16:14 -0700 Subject: [PATCH 06/12] Update packages/google-auth/google/auth/transport/mtls_interceptor.py Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com> --- packages/google-auth/google/auth/transport/mtls_interceptor.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/packages/google-auth/google/auth/transport/mtls_interceptor.py b/packages/google-auth/google/auth/transport/mtls_interceptor.py index deab9521f94c..843a60d6e689 100644 --- a/packages/google-auth/google/auth/transport/mtls_interceptor.py +++ b/packages/google-auth/google/auth/transport/mtls_interceptor.py @@ -711,7 +711,7 @@ def __next__(self): def add_done_callback(self, fn): with self._lock: - if getattr(self, "_is_completed", False): + if self._is_completed: fire_now = True else: self._done_callbacks.append(fn) From 96ec8765c6bbc551bdff36e0c4e5e22ca0fffb7c Mon Sep 17 00:00:00 2001 From: agrawalradhika-cell Date: Mon, 5 Oct 2026 11:43:23 -0700 Subject: [PATCH 07/12] chore: Implement additional methods in mtls_interceptor Added methods for initial and trailing metadata, time remaining, and active status --- .../google/auth/transport/mtls_interceptor.py | 21 ++++++++++++++++--- 1 file changed, 18 insertions(+), 3 deletions(-) diff --git a/packages/google-auth/google/auth/transport/mtls_interceptor.py b/packages/google-auth/google/auth/transport/mtls_interceptor.py index 843a60d6e689..460ca67b978c 100644 --- a/packages/google-auth/google/auth/transport/mtls_interceptor.py +++ b/packages/google-auth/google/auth/transport/mtls_interceptor.py @@ -266,6 +266,18 @@ def code(self): def details(self): return self._details + def initial_metadata(self): + return None + + def trailing_metadata(self): + return None + + def time_remaining(self): + return 0.0 + + def is_active(self): + return False + class _BaseCallWrapper(grpc.Future, grpc.Call): """A generic wrapper that delegates standard grpc.Call and grpc.Future @@ -475,7 +487,8 @@ def add_done_callback(self, fn): if fire_now: try: fn(self) - except Exception: + except Exception as e: + _LOGGER.warning("Callback failed: %s", e) pass def result(self, timeout=None): @@ -614,7 +627,8 @@ def _trigger_callbacks(self): for fn in callbacks: try: fn(self) - except Exception: + except Exception as e: + _LOGGER.warning("Callback failed: %s", e) pass def _on_inner_call_done(self, inner_call): @@ -720,5 +734,6 @@ def add_done_callback(self, fn): if fire_now: try: fn(self) - except Exception: + except Exception as e: + _LOGGER.warning("Callback failed: %s", e) pass From 068ff10064afb584ebf3437d74da3112d496edd9 Mon Sep 17 00:00:00 2001 From: agrawalradhika-cell Date: Mon, 5 Oct 2026 13:46:28 -0700 Subject: [PATCH 08/12] chore: Add exception handling for channel subscription Handle exceptions during channel subscription to avoid crashes. --- .../google-auth/google/auth/transport/mtls_interceptor.py | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/packages/google-auth/google/auth/transport/mtls_interceptor.py b/packages/google-auth/google/auth/transport/mtls_interceptor.py index 460ca67b978c..ae9700e8a0fb 100644 --- a/packages/google-auth/google/auth/transport/mtls_interceptor.py +++ b/packages/google-auth/google/auth/transport/mtls_interceptor.py @@ -143,7 +143,10 @@ def refresh_logic(self, count, call_cert_bytes=None, call_key_bytes=None): old_channel.unsubscribe(callback) except Exception: pass - self._channel.subscribe(callback) + try: + self._channel.subscribe(callback) + except Exception: + pass def unary_unary(self, method, *args, **kwargs): # Always return a callable from the CURRENT channel @@ -637,7 +640,7 @@ def _on_inner_call_done(self, inner_call): return # Intercept and suppress premature callbacks for UNAUTHENTICATED. # __next__ inherently handles this error and manages triggering callbacks - # later if retriies are exhausted. + # later if retries are exhausted. if ( callable(getattr(inner_call, "code", None)) and inner_call.code() == grpc.StatusCode.UNAUTHENTICATED From ae1633a10883843167b78c6e9975305912ce5c5c Mon Sep 17 00:00:00 2001 From: agrawalradhika-cell Date: Mon, 5 Oct 2026 14:16:23 -0700 Subject: [PATCH 09/12] chore: Clean up exception handling by removing pass Removed redundant pass statements in exception handling. --- packages/google-auth/google/auth/transport/mtls_interceptor.py | 3 --- 1 file changed, 3 deletions(-) diff --git a/packages/google-auth/google/auth/transport/mtls_interceptor.py b/packages/google-auth/google/auth/transport/mtls_interceptor.py index ae9700e8a0fb..00f8b0501cce 100644 --- a/packages/google-auth/google/auth/transport/mtls_interceptor.py +++ b/packages/google-auth/google/auth/transport/mtls_interceptor.py @@ -492,7 +492,6 @@ def add_done_callback(self, fn): fn(self) except Exception as e: _LOGGER.warning("Callback failed: %s", e) - pass def result(self, timeout=None): if not self._completion_event.wait(timeout): @@ -632,7 +631,6 @@ def _trigger_callbacks(self): fn(self) except Exception as e: _LOGGER.warning("Callback failed: %s", e) - pass def _on_inner_call_done(self, inner_call): with self._lock: @@ -739,4 +737,3 @@ def add_done_callback(self, fn): fn(self) except Exception as e: _LOGGER.warning("Callback failed: %s", e) - pass From 01d064839194b7ba50640b6a7d156baa49cb328e Mon Sep 17 00:00:00 2001 From: agrawalradhika-cell Date: Mon, 5 Oct 2026 14:48:58 -0700 Subject: [PATCH 10/12] chore: handle callbacks on terminal exception handling Added logic to fire callbacks after handling exceptions. --- .../google-auth/google/auth/transport/mtls_interceptor.py | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/packages/google-auth/google/auth/transport/mtls_interceptor.py b/packages/google-auth/google/auth/transport/mtls_interceptor.py index 00f8b0501cce..ba439368dd5a 100644 --- a/packages/google-auth/google/auth/transport/mtls_interceptor.py +++ b/packages/google-auth/google/auth/transport/mtls_interceptor.py @@ -449,6 +449,10 @@ def _on_inner_future_done(self, inner_future): with self._lock: self._terminal_exception = e self._completion_event.set() + callbacks_to_fire = list(self._done_callbacks) + self._done_callbacks.clear() + for callback in callbacks_to_fire: + callback(self) return with self._lock: self._retry_count += 1 From ac0f7e3908f4d87595b06b6a8f64e8e0a1a05dd8 Mon Sep 17 00:00:00 2001 From: agrawalradhika-cell Date: Mon, 5 Oct 2026 15:11:21 -0700 Subject: [PATCH 11/12] Add cancel and add_callback methods to interceptor to avoid AttributeError Add cancel and add_callback methods to interceptor to avoid AttributeError --- .../google-auth/google/auth/transport/mtls_interceptor.py | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/packages/google-auth/google/auth/transport/mtls_interceptor.py b/packages/google-auth/google/auth/transport/mtls_interceptor.py index ba439368dd5a..a928a21d6de8 100644 --- a/packages/google-auth/google/auth/transport/mtls_interceptor.py +++ b/packages/google-auth/google/auth/transport/mtls_interceptor.py @@ -281,6 +281,11 @@ def time_remaining(self): def is_active(self): return False + def cancel(self): + return False + + def add_callback(self, callback): + return False class _BaseCallWrapper(grpc.Future, grpc.Call): """A generic wrapper that delegates standard grpc.Call and grpc.Future From 72a335e71bad7e8021760b36f1d4d39319dc5800 Mon Sep 17 00:00:00 2001 From: agrawalradhika-cell Date: Tue, 6 Oct 2026 10:41:13 -0700 Subject: [PATCH 12/12] fix: Add empty line before _BaseCallWrapper class for lint --- packages/google-auth/google/auth/transport/mtls_interceptor.py | 1 + 1 file changed, 1 insertion(+) diff --git a/packages/google-auth/google/auth/transport/mtls_interceptor.py b/packages/google-auth/google/auth/transport/mtls_interceptor.py index a928a21d6de8..149a5cfd53ab 100644 --- a/packages/google-auth/google/auth/transport/mtls_interceptor.py +++ b/packages/google-auth/google/auth/transport/mtls_interceptor.py @@ -287,6 +287,7 @@ def cancel(self): def add_callback(self, callback): return False + class _BaseCallWrapper(grpc.Future, grpc.Call): """A generic wrapper that delegates standard grpc.Call and grpc.Future methods to an underlying call object.