From 89428d630d77d310187c7d204de85fb537e03dfa Mon Sep 17 00:00:00 2001 From: Matt Castelaz Date: Wed, 5 Aug 2026 01:29:26 +0000 Subject: [PATCH 01/29] feat(gax): support transparent retries during mTLS certificate rotations - Add CertificateBasedAccess and WorkloadCertificateUtils for SPIFFE and custom certificate loading - Implement RefreshingHttpJsonChannel and ChannelPool mTLS certificate fingerprint tracking and rotation - Enable transparent retries for retryable UnauthenticatedExceptions in ApiResultRetryAlgorithm and AttemptCallable - Add override delegation for getEndpoint, getHttpTransport, and getExecutor to preserve SLF4J MDC logging in Showcase tests --- .../com/google/api/gax/grpc/ChannelPool.java | 131 ++++++- .../google/api/gax/grpc/GrpcCallContext.java | 71 +++- .../api/gax/grpc/GrpcTransportChannel.java | 17 + .../InstantiatingGrpcChannelProvider.java | 22 +- .../google/api/gax/grpc/ChannelPoolTest.java | 113 +++++- .../api/gax/grpc/GrpcCallContextTest.java | 21 + .../api/gax/grpc/GrpcClientCallsTest.java | 1 + .../api/gax/httpjson/HttpJsonCallContext.java | 75 +++- .../httpjson/HttpJsonTransportChannel.java | 10 + .../InstantiatingHttpJsonChannelProvider.java | 76 ++-- .../gax/httpjson/ManagedHttpJsonChannel.java | 45 ++- .../ManagedHttpJsonInterceptorChannel.java | 29 ++ .../httpjson/RefreshingHttpJsonChannel.java | 369 ++++++++++++++++++ ...tantiatingHttpJsonChannelProviderTest.java | 72 ++-- .../RefreshingHttpJsonChannelTest.java | 327 ++++++++++++++++ .../google/api/gax/rpc/ApiCallContext.java | 14 +- .../api/gax/rpc/ApiResultRetryAlgorithm.java | 7 +- .../google/api/gax/rpc/AttemptCallable.java | 28 +- .../api/gax/rpc/BidiStreamingCallable.java | 43 +- .../api/gax/rpc/ClientStreamingCallable.java | 38 +- .../rpc/ServerStreamingAttemptCallable.java | 21 + .../google/api/gax/rpc/TransportChannel.java | 18 +- .../gax/rpc/mtls/CertificateBasedAccess.java | 171 +++++++- .../rpc/mtls/WorkloadCertificateUtils.java | 68 ++++ .../api/gax/rpc/AttemptCallableTest.java | 50 +++ .../api/gax/rpc/EndpointContextTest.java | 48 ++- .../ServerStreamingAttemptCallableTest.java | 41 ++ .../api/gax/rpc/StreamingCallableTest.java | 4 +- .../AbstractMtlsTransportChannelTest.java | 18 +- .../rpc/mtls/CertificateBasedAccessTest.java | 216 +++++++++- .../api/gax/rpc/testing/FakeCallContext.java | 5 + .../api/gax/rpc/testing/FakeChannel.java | 22 +- .../gax/rpc/testing/FakeTransportChannel.java | 27 ++ 33 files changed, 1988 insertions(+), 230 deletions(-) create mode 100644 sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java create mode 100644 sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java create mode 100644 sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/WorkloadCertificateUtils.java diff --git a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java index fab73a55dccf..5c583654f43b 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java @@ -31,6 +31,7 @@ import com.google.api.core.InternalApi; import com.google.api.gax.core.FixedExecutorProvider; +import com.google.api.gax.rpc.mtls.WorkloadCertificateUtils; import com.google.common.annotations.VisibleForTesting; import com.google.common.base.Preconditions; import com.google.common.collect.ImmutableList; @@ -72,18 +73,36 @@ @NullMarked class ChannelPool extends ManagedChannel { static final String CHANNEL_POOL_CONSECUTIVE_RESIZING_WARNING = - "The gRPC ChannelPool used in the client has been flagged to be repeatedly resizing (5+ times). See https://github.com/googleapis/google-cloud-java/blob/main/docs/grpc_channel_pool_guide.md for more information about this behavior."; + "The gRPC ChannelPool used in the client has been flagged to be repeatedly resizing (5+" + + " times). See" + + " https://github.com/googleapis/google-cloud-java/blob/main/docs/grpc_channel_pool_guide.md" + + " for more information about this behavior."; @VisibleForTesting static final Logger LOG = Logger.getLogger(ChannelPool.class.getName()); private static final java.time.Duration REFRESH_PERIOD = java.time.Duration.ofMinutes(50); private final ChannelPoolSettings settings; private final ChannelFactory channelFactory; private final FixedExecutorProvider backgroundExecutorProvider; + private final String workloadCertPath; private @Nullable ScheduledFuture refreshFuture = null; private @Nullable ScheduledFuture resizeFuture = null; + private static class DiskCheckResult { + final String fingerprint; + final long timestampNanos; + + DiskCheckResult(String fingerprint, long timestampNanos) { + this.fingerprint = fingerprint; + this.timestampNanos = timestampNanos; + } + } + + private volatile DiskCheckResult lastDiskCheck = null; + private final java.util.concurrent.locks.ReentrantLock diskCheckLock = + new java.util.concurrent.locks.ReentrantLock(); private final Object entryWriteLock = new Object(); + private volatile String activeCertFingerprint = ""; @VisibleForTesting final AtomicReference> entries = new AtomicReference<>(); private final AtomicInteger indexTicker = new AtomicInteger(); private final String authority; @@ -100,14 +119,15 @@ class ChannelPool extends ManagedChannel { static ChannelPool create( ChannelPoolSettings settings, ChannelFactory channelFactory, - @Nullable ScheduledExecutorService backgroundExecutor) + @Nullable ScheduledExecutorService backgroundExecutor, + @Nullable String workloadCertPath) throws IOException { FixedExecutorProvider executorProvider = backgroundExecutor == null ? FixedExecutorProvider.create(Executors.newSingleThreadScheduledExecutor(), true) : FixedExecutorProvider.create(backgroundExecutor, false); - return new ChannelPool(settings, channelFactory, executorProvider); + return new ChannelPool(settings, channelFactory, executorProvider, workloadCertPath); } /** @@ -121,11 +141,13 @@ static ChannelPool create( ChannelPool( ChannelPoolSettings settings, ChannelFactory channelFactory, - FixedExecutorProvider executorProvider) + FixedExecutorProvider executorProvider, + @Nullable String workloadCertPath) throws IOException { this.settings = settings; this.channelFactory = channelFactory; this.backgroundExecutorProvider = executorProvider; + this.workloadCertPath = workloadCertPath; ImmutableList.Builder initialListBuilder = ImmutableList.builder(); @@ -136,6 +158,11 @@ static ChannelPool create( entries.set(initialListBuilder.build()); authority = entries.get().get(0).channel.authority(); + if (workloadCertPath != null) { + this.activeCertFingerprint = + WorkloadCertificateUtils.getCertificateFingerprint(workloadCertPath); + } + if (!settings.isStaticSize()) { resizeFuture = backgroundExecutorProvider @@ -421,12 +448,54 @@ private void expand(int desiredSize) { private void refreshSafely() { try { - refresh(); + synchronized (entryWriteLock) { + if (workloadCertPath != null) { + String currentDiskFingerprint = getOrUpdateDiskFingerprint(workloadCertPath); + if (!currentDiskFingerprint.isEmpty()) { + this.activeCertFingerprint = currentDiskFingerprint; + } + } + refreshAll(); + } } catch (Exception e) { - LOG.log(Level.WARNING, "Failed to pre-emptively refresh channnels", e); + LOG.log(Level.WARNING, "Failed to pre-emptively refresh channels", e); } } + private String getOrUpdateDiskFingerprint(String certPath) { + long now = System.nanoTime(); + DiskCheckResult cached = lastDiskCheck; + if (cached != null + && (now - cached.timestampNanos < java.util.concurrent.TimeUnit.SECONDS.toNanos(1))) { + return cached.fingerprint; + } + + diskCheckLock.lock(); + try { + cached = lastDiskCheck; + if (cached != null + && (now - cached.timestampNanos < java.util.concurrent.TimeUnit.SECONDS.toNanos(1))) { + return cached.fingerprint; + } + String fingerprint = WorkloadCertificateUtils.getCertificateFingerprint(certPath); + lastDiskCheck = new DiskCheckResult(fingerprint, System.nanoTime()); + return fingerprint; + } finally { + diskCheckLock.unlock(); + } + } + + boolean shouldRefresh() { + if (workloadCertPath == null) { + return false; + } + String currentDiskFingerprint = getOrUpdateDiskFingerprint(workloadCertPath); + if (currentDiskFingerprint.isEmpty()) { + return false; + } + return !currentDiskFingerprint.equalsIgnoreCase(activeCertFingerprint); + } + /** * Replace all of the channels in the channel pool with fresh ones. This is meant to mitigate the * hourly GFE disconnects by giving clients the ability to prime the channel on reconnect. @@ -443,7 +512,35 @@ void refresh() { // - then thread2 will shut down channel that thread1 will put back into circulation (after it // replaces the list) synchronized (entryWriteLock) { - LOG.fine("Refreshing all channels"); + if (workloadCertPath == null) { + return; + } + String currentDiskFingerprint = getOrUpdateDiskFingerprint(workloadCertPath); + if (currentDiskFingerprint.isEmpty()) { + return; + } + + // Double-check fingerprint inside the lock + if (currentDiskFingerprint.equalsIgnoreCase(this.activeCertFingerprint)) { + LOG.fine( + "Channel pool was already refreshed by a concurrent thread, skipping duplicate" + + " refresh"); + return; + } + + this.activeCertFingerprint = currentDiskFingerprint; + refreshAll(); + } + } + + @InternalApi("Visible for testing") + void refreshAll() { + synchronized (entryWriteLock) { + LOG.fine( + "Refreshing all channels" + + (activeCertFingerprint == null + ? "" + : " with certificate fingerprint: " + activeCertFingerprint)); ArrayList newEntries = new ArrayList<>(entries.get()); for (int i = 0; i < newEntries.size(); i++) { @@ -621,7 +718,13 @@ public ClientCall newCall( } } - /** ClientCall wrapper that makes sure to decrement the outstanding RPC count on completion. */ + /** + * ClientCall wrapper that makes sure to decrement the outstanding RPC count on completion. + * + *

Contract: Exactly one call to {@link #start(Listener, Metadata)} or explicit release via + * {@link #cancel(String, Throwable)} is required to balance reference counts. Early cancellation + * before {@code start()} safely decrements the reference count via atomic compare-and-set. + */ static class ReleasingClientCall extends SimpleForwardingClientCall { private @Nullable CancellationException cancellationException; final Entry entry; @@ -636,6 +739,9 @@ public ReleasingClientCall(ClientCall delegate, Entry entry) { @Override public void start(Listener responseListener, Metadata headers) { if (cancellationException != null) { + if (wasReleased.compareAndSet(false, true)) { + entry.release(); + } throw new IllegalStateException("Call is already cancelled", cancellationException); } try { @@ -646,7 +752,8 @@ public void onClose(Status status, Metadata trailers) { if (!wasClosed.compareAndSet(false, true)) { LOG.log( Level.WARNING, - "Call is being closed more than once. Please make sure that onClose() is not being manually called."); + "Call is being closed more than once. Please make sure that onClose() is not" + + " being manually called."); return; } try { @@ -657,7 +764,8 @@ public void onClose(Status status, Metadata trailers) { } else { LOG.log( Level.WARNING, - "Entry was released before the call is closed. This may be due to an exception on start of the call."); + "Entry was released before the call is closed. This may be due to an" + + " exception on start of the call."); } } } @@ -670,7 +778,8 @@ public void onClose(Status status, Metadata trailers) { } else { LOG.log( Level.WARNING, - "The entry is already released. This indicates that onClose() has already been called previously"); + "The entry is already released. This indicates that onClose() has already been called" + + " previously"); } throw e; } diff --git a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/GrpcCallContext.java b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/GrpcCallContext.java index 23f56c5f8951..e5ead54e9181 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/GrpcCallContext.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/GrpcCallContext.java @@ -99,6 +99,7 @@ public final class GrpcCallContext implements ApiCallContext { private final ApiCallContextOptions options; private final EndpointContext endpointContext; private final boolean isDirectPath; + @Nullable private final TransportChannel transportChannel; /** Returns an empty instance with a null channel and default {@link CallOptions}. */ public static GrpcCallContext createDefault() { @@ -115,7 +116,8 @@ public static GrpcCallContext createDefault() { null, null, null, - false); + false, + null); } /** Returns an instance with the given channel and {@link CallOptions}. */ @@ -133,7 +135,8 @@ public static GrpcCallContext of(Channel channel, CallOptions callOptions) { null, null, null, - false); + false, + null); } private GrpcCallContext( @@ -149,7 +152,8 @@ private GrpcCallContext( @Nullable RetrySettings retrySettings, @Nullable Set retryableCodes, @Nullable EndpointContext endpointContext, - boolean isDirectPath) { + boolean isDirectPath, + @Nullable TransportChannel transportChannel) { this.channel = channel; this.credentials = credentials; Preconditions.checkNotNull(callOptions); @@ -169,6 +173,7 @@ private GrpcCallContext( this.endpointContext = endpointContext == null ? EndpointContext.getDefaultInstance() : endpointContext; this.isDirectPath = isDirectPath; + this.transportChannel = transportChannel; } /** @@ -210,7 +215,13 @@ public GrpcCallContext withCredentials(Credentials newCredentials) { retrySettings, retryableCodes, endpointContext, - isDirectPath); + isDirectPath, + transportChannel); + } + + @Override + public TransportChannel getTransportChannel() { + return transportChannel; } @Override @@ -234,7 +245,8 @@ public GrpcCallContext withTransportChannel(TransportChannel inputChannel) { retrySettings, retryableCodes, endpointContext, - transportChannel.isDirectPath()); + transportChannel.isDirectPath(), + inputChannel); } @Override @@ -253,7 +265,8 @@ public GrpcCallContext withEndpointContext(EndpointContext endpointContext) { retrySettings, retryableCodes, endpointContext, - isDirectPath); + isDirectPath, + transportChannel); } /** This method is obsolete. Use {@link #withTimeoutDuration(java.time.Duration)} instead. */ @@ -271,7 +284,7 @@ public GrpcCallContext withTimeoutDuration(java.time.@Nullable Duration timeout) } // Prevent expanding timeouts - if (timeout != null && this.timeout != null && this.timeout.compareTo(timeout) <= 0) { + if (this.timeout != null && (timeout == null || this.timeout.compareTo(timeout) <= 0)) { return this; } @@ -288,7 +301,8 @@ public GrpcCallContext withTimeoutDuration(java.time.@Nullable Duration timeout) retrySettings, retryableCodes, endpointContext, - isDirectPath); + isDirectPath, + transportChannel); } @Override @@ -334,7 +348,8 @@ public GrpcCallContext withStreamWaitTimeoutDuration( retrySettings, retryableCodes, endpointContext, - isDirectPath); + isDirectPath, + transportChannel); } /** @@ -369,7 +384,8 @@ public GrpcCallContext withStreamIdleTimeoutDuration( retrySettings, retryableCodes, endpointContext, - isDirectPath); + isDirectPath, + transportChannel); } @BetaApi("The surface for channel affinity is not stable yet and may change in the future.") @@ -387,7 +403,8 @@ public GrpcCallContext withChannelAffinity(@Nullable Integer affinity) { retrySettings, retryableCodes, endpointContext, - isDirectPath); + isDirectPath, + transportChannel); } @BetaApi("The surface for extra headers is not stable yet and may change in the future.") @@ -409,7 +426,8 @@ public GrpcCallContext withExtraHeaders(Map> extraHeaders) retrySettings, retryableCodes, endpointContext, - isDirectPath); + isDirectPath, + transportChannel); } @Override @@ -432,7 +450,8 @@ public GrpcCallContext withRetrySettings(RetrySettings retrySettings) { retrySettings, retryableCodes, endpointContext, - isDirectPath); + isDirectPath, + transportChannel); } @Override @@ -455,7 +474,8 @@ public GrpcCallContext withRetryableCodes(Set retryableCodes) { retrySettings, retryableCodes, endpointContext, - isDirectPath); + isDirectPath, + transportChannel); } @Override @@ -542,6 +562,11 @@ public ApiCallContext merge(ApiCallContext inputCallContext) { newCallOptions = newCallOptions.withOption(TRACER_KEY, newTracer); } + TransportChannel newTransportChannel = grpcCallContext.transportChannel; + if (newTransportChannel == null) { + newTransportChannel = transportChannel; + } + // The EndpointContext is not updated as there should be no reason for a user // to update this. return new GrpcCallContext( @@ -557,7 +582,8 @@ public ApiCallContext merge(ApiCallContext inputCallContext) { newRetrySettings, newRetryableCodes, endpointContext, - newIsDirectPath); + newIsDirectPath, + newTransportChannel); } /** The {@link Channel} set on this context. */ @@ -635,7 +661,8 @@ public GrpcCallContext withChannel(@Nullable Channel newChannel) { retrySettings, retryableCodes, endpointContext, - isDirectPath); + isDirectPath, + transportChannel); } /** Returns a new instance with the call options set to the given call options. */ @@ -653,7 +680,8 @@ public GrpcCallContext withCallOptions(CallOptions newCallOptions) { retrySettings, retryableCodes, endpointContext, - isDirectPath); + isDirectPath, + transportChannel); } public GrpcCallContext withRequestParamsDynamicHeaderOption(String requestParams) { @@ -698,7 +726,8 @@ public GrpcCallContext withOption(Key key, T value) { retrySettings, retryableCodes, endpointContext, - isDirectPath); + isDirectPath, + transportChannel); } /** {@inheritDoc} */ @@ -759,7 +788,8 @@ public int hashCode() { options, retrySettings, retryableCodes, - endpointContext); + endpointContext, + transportChannel); } @Override @@ -783,7 +813,8 @@ public boolean equals(@Nullable Object o) { && Objects.equals(options, that.options) && Objects.equals(retrySettings, that.retrySettings) && Objects.equals(retryableCodes, that.retryableCodes) - && Objects.equals(endpointContext, that.endpointContext); + && Objects.equals(endpointContext, that.endpointContext) + && Objects.equals(transportChannel, that.transportChannel); } Metadata getMetadata() { diff --git a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/GrpcTransportChannel.java b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/GrpcTransportChannel.java index e0a520facb17..31ede726f3f3 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/GrpcTransportChannel.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/GrpcTransportChannel.java @@ -68,6 +68,23 @@ public Channel getChannel() { return getManagedChannel(); } + @Override + public void refresh() { + Channel channel = getChannel(); + if (channel instanceof ChannelPool) { + ((ChannelPool) channel).refresh(); + } + } + + @Override + public boolean shouldRefresh() { + Channel channel = getChannel(); + if (channel instanceof ChannelPool) { + return ((ChannelPool) channel).shouldRefresh(); + } + return false; + } + @Override public void shutdown() { getManagedChannel().shutdown(); diff --git a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/InstantiatingGrpcChannelProvider.java b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/InstantiatingGrpcChannelProvider.java index ac42396f006a..07ffa3b86b5e 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/InstantiatingGrpcChannelProvider.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/InstantiatingGrpcChannelProvider.java @@ -406,7 +406,8 @@ private TransportChannel createChannel() throws IOException { ChannelPool.create( channelPoolSettings, InstantiatingGrpcChannelProvider.this::createSingleChannel, - backgroundExecutor)) + backgroundExecutor, + certificateBasedAccess.getWorkloadCertPath())) .setDirectPath(this.canUseDirectPath()) .build(); } @@ -465,8 +466,9 @@ private void logDirectPathMisconfig() { level, "Env var " + DIRECT_PATH_ENV_ENABLE_XDS - + " was found and set to TRUE, but DirectPath was not enabled for this client. If this is intended for " - + "this client, please note that this is a misconfiguration and set the attemptDirectPath option as well."); + + " was found and set to TRUE, but DirectPath was not enabled for this client. If" + + " this is intended for this client, please note that this is a misconfiguration" + + " and set the attemptDirectPath option as well."); } // Case 2: Direct Path xDS was enabled via Builder. Direct Path Traffic Director must be set // (enabled with `setAttemptDirectPath(true)`) along with xDS. @@ -474,7 +476,9 @@ private void logDirectPathMisconfig() { else if (isDirectPathXdsEnabledViaBuilderOption()) { LOG.log( level, - "DirectPath is misconfigured. The DirectPath XDS option was set, but the attemptDirectPath option was not. Please set both the attemptDirectPath and attemptDirectPathXds options."); + "DirectPath is misconfigured. The DirectPath XDS option was set, but the" + + " attemptDirectPath option was not. Please set both the attemptDirectPath and" + + " attemptDirectPathXds options."); } } else { // Case 3: credential is not correctly set @@ -666,7 +670,8 @@ ChannelCredentials createS2ASecuredChannelCredentials() { // Fallback to plaintext connection to S2A. LOG.log( Level.INFO, - "Cannot establish an mTLS connection to S2A because autoconfig endpoint did not return a mtls address to reach S2A."); + "Cannot establish an mTLS connection to S2A because autoconfig endpoint did not" + + " return a mtls address to reach S2A."); s2aChannelCredentials = createPlaintextToS2AChannelCredentials(plaintextAddress); return s2aChannelCredentials; } @@ -685,7 +690,9 @@ ChannelCredentials createS2ASecuredChannelCredentials() { // Fallback to plaintext-to-S2A connection on error. LOG.log( Level.WARNING, - "Cannot establish an mTLS connection to S2A due to error creating MTLS to MDS TlsChannelCredentials credentials, falling back to plaintext connection to S2A: " + "Cannot establish an mTLS connection to S2A due to error creating MTLS to MDS" + + " TlsChannelCredentials credentials, falling back to plaintext connection to" + + " S2A: " + ignore.getMessage()); s2aChannelCredentials = createPlaintextToS2AChannelCredentials(plaintextAddress); return s2aChannelCredentials; @@ -1403,7 +1410,8 @@ public InstantiatingGrpcChannelProvider build() { "DefaultMtlsProviderFactory encountered unexpected IOException: " + e.getMessage()); LOG.log( Level.WARNING, - "mTLS configuration was detected on the device, but mTLS failed to initialize. Falling back to non-mTLS channel."); + "mTLS configuration was detected on the device, but mTLS failed to initialize." + + " Falling back to non-mTLS channel."); } } } diff --git a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java index 5bfdc7754759..6c2adc1f1463 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java @@ -81,13 +81,18 @@ class ChannelPoolTest { private static final int DEFAULT_AWAIT_TERMINATION_SEC = 10; private ChannelPool pool; + private java.nio.file.Path tempCert; @AfterEach - void cleanup() throws InterruptedException { + void cleanup() throws InterruptedException, IOException { if (pool != null) { pool.shutdown(); pool.awaitTermination(DEFAULT_AWAIT_TERMINATION_SEC, TimeUnit.SECONDS); } + if (tempCert != null) { + java.nio.file.Files.deleteIfExists(tempCert); + tempCert = null; + } } @Test @@ -101,6 +106,7 @@ void testAuthority() throws IOException { ChannelPool.create( ChannelPoolSettings.staticallySized(2), new FakeChannelFactory(Arrays.asList(sub1, sub2)), + null, null); assertThat(pool.authority()).isEqualTo("myAuth"); } @@ -117,6 +123,7 @@ void testRoundRobin() throws IOException { ChannelPool.create( ChannelPoolSettings.staticallySized(channels.size()), new FakeChannelFactory(channels), + null, null); verifyTargetChannel(pool, channels, sub1); @@ -195,6 +202,7 @@ void ensureEvenDistribution() throws InterruptedException, IOException { ChannelPool.create( ChannelPoolSettings.staticallySized(numChannels), new FakeChannelFactory(Arrays.asList(channels)), + null, null); int numThreads = 20; @@ -233,6 +241,7 @@ void channelPrimerShouldCallPoolConstruction() throws IOException { .setPreemptiveRefreshEnabled(true) .build(), new FakeChannelFactory(Arrays.asList(channel1, channel2), mockChannelPrimer), + null, null); Mockito.verify(mockChannelPrimer, Mockito.times(2)) .primeChannel(Mockito.any(ManagedChannel.class)); @@ -273,7 +282,8 @@ void channelPrimerIsCalledPeriodically() throws IOException { .setPreemptiveRefreshEnabled(true) .build(), channelFactory, - provider); + provider, + null); // 1 call during the creation Mockito.verify(mockChannelPrimer, Mockito.times(1)) .primeChannel(Mockito.any(ManagedChannel.class)); @@ -297,7 +307,7 @@ void callShouldCompleteAfterCreation() throws IOException { ManagedChannel replacementChannel = mock(ManagedChannel.class); FakeChannelFactory channelFactory = new FakeChannelFactory(ImmutableList.of(underlyingChannel, replacementChannel)); - pool = ChannelPool.create(ChannelPoolSettings.staticallySized(1), channelFactory, null); + pool = ChannelPool.create(ChannelPoolSettings.staticallySized(1), channelFactory, null, null); // create a mock call when new call comes to the underlying channel MockClientCall mockClientCall = new MockClientCall<>(1, Status.OK); @@ -322,7 +332,7 @@ void callShouldCompleteAfterCreation() throws IOException { ClientCall call = pool.newCall(FakeMethodDescriptor.create(), CallOptions.DEFAULT); - pool.refresh(); + pool.refreshAll(); // shutdown is not called because there is still an outstanding call, even if it hasn't started Mockito.verify(underlyingChannel, Mockito.after(200).never()).shutdown(); @@ -346,7 +356,7 @@ void callShouldCompleteAfterStarted() throws IOException { FakeChannelFactory channelFactory = new FakeChannelFactory(ImmutableList.of(underlyingChannel, replacementChannel)); - pool = ChannelPool.create(ChannelPoolSettings.staticallySized(1), channelFactory, null); + pool = ChannelPool.create(ChannelPoolSettings.staticallySized(1), channelFactory, null, null); // create a mock call when new call comes to the underlying channel MockClientCall mockClientCall = new MockClientCall<>(1, Status.OK); @@ -373,7 +383,7 @@ void callShouldCompleteAfterStarted() throws IOException { // start clientCall call.start(listener, new Metadata()); - pool.refresh(); + pool.refreshAll(); // shutdown is not called because there is still an outstanding call Mockito.verify(underlyingChannel, Mockito.after(200).never()).shutdown(); @@ -391,7 +401,7 @@ void channelShouldShutdown() throws IOException { FakeChannelFactory channelFactory = new FakeChannelFactory(ImmutableList.of(underlyingChannel, replacementChannel)); - pool = ChannelPool.create(ChannelPoolSettings.staticallySized(1), channelFactory, null); + pool = ChannelPool.create(ChannelPoolSettings.staticallySized(1), channelFactory, null, null); // create a mock call when new call comes to the underlying channel MockClientCall mockClientCall = new MockClientCall<>(1, Status.OK); @@ -422,11 +432,62 @@ void channelShouldShutdown() throws IOException { call.sendMessage("message"); // shutdown is not called because it has not been shutdown yet Mockito.verify(underlyingChannel, Mockito.after(200).never()).shutdown(); - pool.refresh(); + pool.refreshAll(); // shutdown is called because the outstanding call has completed Mockito.verify(underlyingChannel, Mockito.atLeastOnce()).shutdown(); } + @Test + void channelReactiveMTlsRefreshShouldConditionallySwapChannels() + throws IOException, InterruptedException { + ManagedChannel underlyingChannel1 = Mockito.mock(ManagedChannel.class); + ManagedChannel underlyingChannel2 = Mockito.mock(ManagedChannel.class); + + FakeChannelFactory channelFactory = + new FakeChannelFactory(ImmutableList.of(underlyingChannel1, underlyingChannel2)); + + // Create a temp file to act as the cert + tempCert = java.nio.file.Files.createTempFile("cert", ".pem"); + + java.nio.file.Path clientCert = + java.nio.file.Paths.get("src", "test", "resources", "client_cert.pem"); + java.nio.file.Files.copy( + clientCert, tempCert, java.nio.file.StandardCopyOption.REPLACE_EXISTING); + + ChannelPoolSettings channelPoolSettings = + ChannelPoolSettings.builder().setInitialChannelCount(1).build(); + + pool = ChannelPool.create(channelPoolSettings, channelFactory, null, tempCert.toString()); + + // Initially uses channel1 + pool.newCall(FakeMethodDescriptor.create(), CallOptions.DEFAULT); + Mockito.verify(underlyingChannel1, Mockito.times(1)) + .newCall(Mockito.>any(), Mockito.any(CallOptions.class)); + + // Try a reactive refresh *without* changing the cert content (should no-op) + pool.refresh(); + + // Verify it's STILL channel1 + pool.newCall(FakeMethodDescriptor.create(), CallOptions.DEFAULT); + Mockito.verify(underlyingChannel1, Mockito.times(2)) + .newCall(Mockito.>any(), Mockito.any(CallOptions.class)); + + // The ChannelPool caches fingerprints for 1000ms, wait for it to expire + Thread.sleep(1100); + + java.nio.file.Path rootCert = + java.nio.file.Paths.get("src", "test", "resources", "root_cert.pem"); + java.nio.file.Files.copy(rootCert, tempCert, java.nio.file.StandardCopyOption.REPLACE_EXISTING); + + // Try a reactive refresh *with* a changed cert content (should swap channels) + pool.refresh(); + + // Verify it is NOW channel2 + pool.newCall(FakeMethodDescriptor.create(), CallOptions.DEFAULT); + Mockito.verify(underlyingChannel2, Mockito.times(1)) + .newCall(Mockito.>any(), Mockito.any(CallOptions.class)); + } + @Test void channelRefreshShouldSwapChannels() throws IOException { ManagedChannel underlyingChannel1 = mock(ManagedChannel.class); @@ -450,7 +511,8 @@ void channelRefreshShouldSwapChannels() throws IOException { .setPreemptiveRefreshEnabled(true) .build(), channelFactory, - provider); + provider, + null); Mockito.reset(underlyingChannel1); pool.newCall(FakeMethodDescriptor.create(), CallOptions.DEFAULT); @@ -459,7 +521,7 @@ void channelRefreshShouldSwapChannels() throws IOException { .newCall(Mockito.>any(), Mockito.any(CallOptions.class)); // swap channel - pool.refresh(); + pool.refreshAll(); pool.newCall(FakeMethodDescriptor.create(), CallOptions.DEFAULT); @@ -486,7 +548,8 @@ void channelCountShouldNotChangeWhenOutstandingRpcsAreWithinLimits() throws Exce .setMaxRpcsPerChannel(2) .build(), channelFactory, - provider); + provider, + null); assertThat(pool.entries.get()).hasSize(2); // Start the minimum number of @@ -553,7 +616,8 @@ void customResizeDeltaIsRespected() throws Exception { .setMaxResizeDelta(5) .build(), channelFactory, - provider); + provider, + null); assertThat(pool.entries.get()).hasSize(2); // Add 20 RPCs to push expansion @@ -586,7 +650,8 @@ void removedIdleChannelsAreShutdown() throws Exception { .setMaxRpcsPerChannel(2) .build(), channelFactory, - provider); + provider, + null); assertThat(pool.entries.get()).hasSize(2); // With no outstanding RPCs, the pool should shrink @@ -614,7 +679,8 @@ void removedActiveChannelsAreShutdown() throws Exception { .setMaxRpcsPerChannel(2) .build(), channelFactory, - provider); + provider, + null); assertThat(pool.entries.get()).hasSize(2); // Start 2 RPCs @@ -652,7 +718,7 @@ void testReleasingClientCallCancelEarly() throws IOException { Mockito.when(fakeChannel.newCall(Mockito.any(), Mockito.any())).thenReturn(mockClientCall); ChannelPoolSettings channelPoolSettings = ChannelPoolSettings.staticallySized(1); ChannelFactory factory = new FakeChannelFactory(ImmutableList.of(fakeChannel)); - pool = ChannelPool.create(channelPoolSettings, factory, null); + pool = ChannelPool.create(channelPoolSettings, factory, null, null); EndpointContext endpointContext = Mockito.mock(EndpointContext.class, Mockito.withSettings().withoutAnnotations()); @@ -717,7 +783,8 @@ void repeatedResizingLogsWarningOnExpand() throws Exception { .setMaxChannelCount(10) .build(), channelFactory, - provider); + provider, + null); assertThat(pool.entries.get()).hasSize(1); FakeLogHandler logHandler = new FakeLogHandler(); @@ -769,7 +836,8 @@ void repeatedResizingLogsWarningOnShrink() throws Exception { .setMaxChannelCount(10) .build(), channelFactory, - provider); + provider, + null); assertThat(pool.entries.get()).hasSize(10); FakeLogHandler logHandler = new FakeLogHandler(); @@ -805,7 +873,7 @@ void testDoubleRelease() throws Exception { ChannelPoolSettings channelPoolSettings = ChannelPoolSettings.staticallySized(1); ChannelFactory factory = new FakeChannelFactory(ImmutableList.of(fakeChannel)); - pool = ChannelPool.create(channelPoolSettings, factory, null); + pool = ChannelPool.create(channelPoolSettings, factory, null, null); EndpointContext endpointContext = Mockito.mock(EndpointContext.class, Mockito.withSettings().withoutAnnotations()); @@ -843,7 +911,8 @@ void testDoubleRelease() throws Exception { // Ensure that the channel pool properly logged the double call and kept the refCount correct assertThat(logHandler.getAllMessages()) .contains( - "Call is being closed more than once. Please make sure that onClose() is not being manually called."); + "Call is being closed more than once. Please make sure that onClose() is not being" + + " manually called."); assertThat(pool.entries.get()).hasSize(1); ChannelPool.Entry entry = pool.entries.get().get(0); assertThat(entry.outstandingRpcs.get()).isEqualTo(0); @@ -879,7 +948,8 @@ void minChannelsClampedToMaxChannelCountUnderHighLoad() throws Exception { .setMaxChannelCount(5) .build(), channelFactory, - provider); + provider, + null); assertThat(pool.entries.get()).hasSize(1); // Add 20 RPCs, which would require 10 channels (20/2) @@ -914,7 +984,8 @@ void maxChannelsClampedToMinChannelCountUnderLowLoad() throws Exception { .setMaxChannelCount(10) .build(), channelFactory, - provider); + provider, + null); assertThat(pool.entries.get()).hasSize(5); // With no outstanding RPCs, the pool should want to shrink to 0 diff --git a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/GrpcCallContextTest.java b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/GrpcCallContextTest.java index e20767fdb8ed..cacdfed88980 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/GrpcCallContextTest.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/GrpcCallContextTest.java @@ -494,4 +494,25 @@ private static Map> createTestExtraHeaders(String... keyVal } return extraHeaders; } + + @Test + public void testEqualsAndHashCode() { + ManagedChannel managedChannel1 = org.mockito.Mockito.mock(ManagedChannel.class); + ManagedChannel managedChannel2 = org.mockito.Mockito.mock(ManagedChannel.class); + + GrpcTransportChannel transportChannel1 = GrpcTransportChannel.create(managedChannel1); + GrpcTransportChannel transportChannel2 = GrpcTransportChannel.create(managedChannel2); + + GrpcCallContext context1 = + GrpcCallContext.createDefault().withTransportChannel(transportChannel1); + GrpcCallContext context2 = + GrpcCallContext.createDefault().withTransportChannel(transportChannel1); + GrpcCallContext context3 = + GrpcCallContext.createDefault().withTransportChannel(transportChannel2); + + org.junit.jupiter.api.Assertions.assertEquals(context1, context2); + org.junit.jupiter.api.Assertions.assertEquals(context1.hashCode(), context2.hashCode()); + + org.junit.jupiter.api.Assertions.assertNotEquals(context1, context3); + } } diff --git a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/GrpcClientCallsTest.java b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/GrpcClientCallsTest.java index 2aa9279e249f..6877eb1dbe1e 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/GrpcClientCallsTest.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/GrpcClientCallsTest.java @@ -125,6 +125,7 @@ void testAffinity() throws IOException { ChannelPool.create( ChannelPoolSettings.staticallySized(2), new FakeChannelFactory(Arrays.asList(channel0, channel1)), + null, null); GrpcCallContext context = defaultCallContext.withChannel(pool); diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/HttpJsonCallContext.java b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/HttpJsonCallContext.java index 2679b51860df..5a2c739ac345 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/HttpJsonCallContext.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/HttpJsonCallContext.java @@ -82,6 +82,7 @@ public final class HttpJsonCallContext implements ApiCallContext { private final @Nullable RetrySettings retrySettings; private final @Nullable ImmutableSet retryableCodes; private final EndpointContext endpointContext; + @Nullable private final TransportChannel transportChannel; /** Returns an empty instance. */ public static HttpJsonCallContext createDefault() { @@ -96,6 +97,7 @@ public static HttpJsonCallContext createDefault() { null, null, null, + null, null); } @@ -111,6 +113,7 @@ public static HttpJsonCallContext of(HttpJsonChannel channel, HttpJsonCallOption null, null, null, + null, null); } @@ -125,7 +128,8 @@ private HttpJsonCallContext( @Nullable ApiTracer tracer, @Nullable RetrySettings defaultRetrySettings, @Nullable Set defaultRetryableCodes, - @Nullable EndpointContext endpointContext) { + @Nullable EndpointContext endpointContext, + @Nullable TransportChannel transportChannel) { this.channel = channel; this.callOptions = callOptions; this.timeout = timeout; @@ -141,6 +145,7 @@ private HttpJsonCallContext( // a valid EndpointContext with user configurations after the client has been initialized. this.endpointContext = endpointContext == null ? EndpointContext.getDefaultInstance() : endpointContext; + this.transportChannel = transportChannel; } /** @@ -220,6 +225,11 @@ public HttpJsonCallContext merge(ApiCallContext inputCallContext) { newRetryableCodes = this.retryableCodes; } + TransportChannel newTransportChannel = httpJsonCallContext.transportChannel; + if (newTransportChannel == null) { + newTransportChannel = this.transportChannel; + } + // The EndpointContext is not updated as there should be no reason for a user // to update this. return new HttpJsonCallContext( @@ -233,7 +243,8 @@ public HttpJsonCallContext merge(ApiCallContext inputCallContext) { newTracer, newRetrySettings, newRetryableCodes, - endpointContext); + endpointContext, + newTransportChannel); } @Override @@ -251,7 +262,24 @@ public HttpJsonCallContext withTransportChannel(TransportChannel inputChannel) { "Expected HttpJsonTransportChannel, got " + inputChannel.getClass().getName()); } HttpJsonTransportChannel transportChannel = (HttpJsonTransportChannel) inputChannel; - return withChannel(transportChannel.getChannel()); + return new HttpJsonCallContext( + transportChannel.getChannel(), + this.callOptions, + this.timeout, + this.streamWaitTimeout, + this.streamIdleTimeout, + this.extraHeaders, + this.options, + this.tracer, + this.retrySettings, + this.retryableCodes, + this.endpointContext, + transportChannel); + } + + @Override + public TransportChannel getTransportChannel() { + return transportChannel; } /** This method is obsolete. Use {@link #withTimeoutDuration(java.time.Duration)} instead. */ @@ -275,7 +303,8 @@ public HttpJsonCallContext withEndpointContext(EndpointContext endpointContext) this.tracer, this.retrySettings, this.retryableCodes, - endpointContext); + endpointContext, + this.transportChannel); } @Override @@ -286,7 +315,7 @@ public HttpJsonCallContext withTimeoutDuration(java.time.Duration timeout) { } // Prevent expanding deadlines - if (timeout != null && this.timeout != null && this.timeout.compareTo(timeout) <= 0) { + if (this.timeout != null && (timeout == null || this.timeout.compareTo(timeout) <= 0)) { return this; } @@ -301,7 +330,8 @@ public HttpJsonCallContext withTimeoutDuration(java.time.Duration timeout) { this.tracer, this.retrySettings, this.retryableCodes, - this.endpointContext); + this.endpointContext, + this.transportChannel); } /** This method is obsolete. Use {@link #getTimeoutDuration()} instead. */ @@ -346,7 +376,8 @@ public HttpJsonCallContext withStreamWaitTimeoutDuration( this.tracer, this.retrySettings, this.retryableCodes, - this.endpointContext); + this.endpointContext, + this.transportChannel); } /** This method is obsolete. Use {@link #getStreamWaitTimeoutDuration()} instead. */ @@ -396,7 +427,8 @@ public HttpJsonCallContext withStreamIdleTimeoutDuration( this.tracer, this.retrySettings, this.retryableCodes, - this.endpointContext); + this.endpointContext, + this.transportChannel); } /** This method is obsolete. Use {@link #getStreamIdleTimeoutDuration()} instead. */ @@ -433,7 +465,8 @@ public ApiCallContext withExtraHeaders(Map> extraHeaders) { this.tracer, this.retrySettings, this.retryableCodes, - this.endpointContext); + this.endpointContext, + this.transportChannel); } @BetaApi("The surface for extra headers is not stable yet and may change in the future.") @@ -457,7 +490,8 @@ public ApiCallContext withOption(Key key, T value) { this.tracer, this.retrySettings, this.retryableCodes, - this.endpointContext); + this.endpointContext, + this.transportChannel); } /** {@inheritDoc} */ @@ -527,7 +561,8 @@ public HttpJsonCallContext withRetrySettings(RetrySettings retrySettings) { this.tracer, retrySettings, this.retryableCodes, - this.endpointContext); + this.endpointContext, + this.transportChannel); } @Override @@ -548,7 +583,8 @@ public HttpJsonCallContext withRetryableCodes(Set retryableCode this.tracer, this.retrySettings, retryableCodes, - this.endpointContext); + this.endpointContext, + this.transportChannel); } public HttpJsonCallContext withChannel(@Nullable HttpJsonChannel newChannel) { @@ -563,7 +599,8 @@ public HttpJsonCallContext withChannel(@Nullable HttpJsonChannel newChannel) { this.tracer, this.retrySettings, this.retryableCodes, - this.endpointContext); + this.endpointContext, + this.transportChannel); } public HttpJsonCallContext withCallOptions(HttpJsonCallOptions newCallOptions) { @@ -578,7 +615,8 @@ public HttpJsonCallContext withCallOptions(HttpJsonCallOptions newCallOptions) { this.tracer, this.retrySettings, this.retryableCodes, - this.endpointContext); + this.endpointContext, + this.transportChannel); } @Deprecated @@ -614,7 +652,8 @@ public HttpJsonCallContext withTracer(@Nonnull ApiTracer newTracer) { newTracer, this.retrySettings, this.retryableCodes, - this.endpointContext); + this.endpointContext, + this.transportChannel); } @Override @@ -634,7 +673,8 @@ public boolean equals(@Nullable Object o) { && Objects.equals(this.tracer, that.tracer) && Objects.equals(this.retrySettings, that.retrySettings) && Objects.equals(this.retryableCodes, that.retryableCodes) - && Objects.equals(this.endpointContext, that.endpointContext); + && Objects.equals(this.endpointContext, that.endpointContext) + && Objects.equals(this.transportChannel, that.transportChannel); } @Override @@ -648,6 +688,7 @@ public int hashCode() { tracer, retrySettings, retryableCodes, - endpointContext); + endpointContext, + transportChannel); } } diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/HttpJsonTransportChannel.java b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/HttpJsonTransportChannel.java index 813622b6a97e..a8333a589a4a 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/HttpJsonTransportChannel.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/HttpJsonTransportChannel.java @@ -64,6 +64,16 @@ public HttpJsonChannel getChannel() { return getManagedChannel(); } + @Override + public void refresh() { + getManagedChannel().refresh(); + } + + @Override + public boolean shouldRefresh() { + return getManagedChannel().shouldRefresh(); + } + @Override public void shutdown() { getManagedChannel().shutdown(); diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProvider.java b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProvider.java index 92ce4efe36aa..13aefaef4b0c 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProvider.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProvider.java @@ -31,7 +31,6 @@ import com.google.api.client.http.HttpTransport; import com.google.api.client.http.javanet.NetHttpTransport; -import com.google.api.client.util.SslUtils; import com.google.api.core.InternalExtensionOnly; import com.google.api.gax.core.ExecutorProvider; import com.google.api.gax.rpc.FixedHeaderProvider; @@ -46,13 +45,11 @@ import java.io.IOException; import java.security.GeneralSecurityException; import java.security.KeyStore; -import java.security.Provider; import java.util.Map; import java.util.concurrent.Executor; import java.util.concurrent.ScheduledExecutorService; import java.util.logging.Level; import java.util.logging.Logger; -import javax.net.ssl.SSLContext; import org.jspecify.annotations.NullMarked; import org.jspecify.annotations.Nullable; @@ -195,58 +192,41 @@ public TransportChannelProvider withCredentials(Credentials credentials) { "InstantiatingHttpJsonChannelProvider doesn't need credentials"); } - HttpTransport createHttpTransport() throws IOException, GeneralSecurityException { - NetHttpTransport.Builder builder = new NetHttpTransport.Builder(); - configureMtls(builder); - HttpJsonConscryptUtils.configureConscryptSecurityProvider(builder); - return builder.build(); - } - - private NetHttpTransport.Builder configureMtls(NetHttpTransport.Builder builder) - throws IOException, GeneralSecurityException { - if (mtlsProvider == null || !certificateBasedAccess.useMtlsClientCertificate()) { - return builder; + @Nullable HttpTransport createHttpTransport() throws IOException, GeneralSecurityException { + if (mtlsProvider == null) { + return null; } - KeyStore mtlsKeyStore = mtlsProvider.getKeyStore(); - if (mtlsKeyStore == null) { - return builder; - } - builder.trustCertificates(null, mtlsKeyStore, ""); - Provider conscryptProvider = HttpJsonConscryptUtils.getConscryptProvider(); - if (conscryptProvider == null) { - // Fall back to standard JDK JSSE if Conscrypt provider is unavailable - return builder; + if (certificateBasedAccess.useMtlsClientCertificate()) { + KeyStore mtlsKeyStore = mtlsProvider.getKeyStore(); + if (mtlsKeyStore != null) { + return new NetHttpTransport.Builder().trustCertificates(null, mtlsKeyStore, "").build(); + } } - // Explicitly initialize SSLContext with the Conscrypt provider so that the client certificate - // key managers - // and trust manager factory (TMF) are bound to Conscrypt's TLS implementation (supporting PQC - // key exchange). - SSLContext sslContext = SSLContext.getInstance("TLS", conscryptProvider); - SslUtils.initSslContext( - sslContext, - null, - SslUtils.getPkixTrustManagerFactory(), - mtlsKeyStore, - "", - SslUtils.getDefaultKeyManagerFactory()); - builder.setSslSocketFactory(sslContext.getSocketFactory()); - return builder; + return null; } private HttpJsonTransportChannel createChannel() throws IOException, GeneralSecurityException { - HttpTransport httpTransportToUse = httpTransport; - if (httpTransportToUse == null) { - httpTransportToUse = createHttpTransport(); - } + java.util.function.Supplier channelFactory = + () -> { + try { + HttpTransport httpTransportToUse = httpTransport; + if (httpTransportToUse == null) { + httpTransportToUse = createHttpTransport(); + } + return ManagedHttpJsonChannel.newBuilder() + .setEndpoint(endpoint) + .setExecutor(executor) + .setHttpTransport(httpTransportToUse) + .setManageHttpTransport(httpTransport == null) + .build(); + } catch (Exception e) { + throw new java.lang.RuntimeException( + "Failed to create fresh ManagedHttpJsonChannel", e); + } + }; - // Pass the executor to the ManagedChannel. If no executor was provided (or null), - // the channel will use a default executor for the calls. ManagedHttpJsonChannel channel = - ManagedHttpJsonChannel.newBuilder() - .setEndpoint(endpoint) - .setExecutor(executor) - .setHttpTransport(httpTransportToUse) - .build(); + new RefreshingHttpJsonChannel(channelFactory, certificateBasedAccess.getWorkloadCertPath()); HttpJsonClientInterceptor headerInterceptor = new HttpJsonHeaderInterceptor(headerProvider.getHeaders()); diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/ManagedHttpJsonChannel.java b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/ManagedHttpJsonChannel.java index 87767bee5c7f..99ece14670f9 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/ManagedHttpJsonChannel.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/ManagedHttpJsonChannel.java @@ -52,11 +52,12 @@ public class ManagedHttpJsonChannel implements HttpJsonChannel, BackgroundResour private final boolean usingDefaultExecutor; private final String endpoint; private final HttpTransport httpTransport; + private final boolean usingDefaultTransport; private final ScheduledExecutorService deadlineScheduledExecutorService; private boolean isTransportShutdown; protected ManagedHttpJsonChannel() { - this(null, true, null, null); + this(null, true, null, null, true); } String getEndpoint() { @@ -72,16 +73,13 @@ private ManagedHttpJsonChannel( @Nullable Executor executor, boolean usingDefaultExecutor, @Nullable String endpoint, - @Nullable HttpTransport httpTransport) { + @Nullable HttpTransport httpTransport, + boolean usingDefaultTransport) { this.executor = executor; this.usingDefaultExecutor = usingDefaultExecutor; this.endpoint = endpoint; - this.httpTransport = - httpTransport == null - ? HttpJsonConscryptUtils.configureConscryptSecurityProvider( - new NetHttpTransport.Builder()) - .build() - : httpTransport; + this.httpTransport = httpTransport == null ? new NetHttpTransport() : httpTransport; + this.usingDefaultTransport = usingDefaultTransport || httpTransport == null; this.deadlineScheduledExecutorService = Executors.newSingleThreadScheduledExecutor(); } @@ -98,6 +96,12 @@ public HttpJsonClientCall newCall( deadlineScheduledExecutorService); } + public void refresh() {} + + public boolean shouldRefresh() { + return false; + } + @VisibleForTesting Executor getExecutor() { return executor; @@ -116,7 +120,9 @@ public synchronized void shutdown() { ((ExecutorService) executor).shutdown(); } deadlineScheduledExecutorService.shutdown(); - httpTransport.shutdown(); + if (usingDefaultTransport) { + httpTransport.shutdown(); + } isTransportShutdown = true; } catch (IOException e) { // TODO: Log this scenario once we implemented the Cloud SDK logging. @@ -158,7 +164,9 @@ public void shutdownNow() { ((ExecutorService) executor).shutdownNow(); } deadlineScheduledExecutorService.shutdownNow(); - httpTransport.shutdown(); + if (usingDefaultTransport) { + httpTransport.shutdown(); + } isTransportShutdown = true; } catch (IOException e) { // TODO: Log this scenario once we implemented the Cloud SDK logging. @@ -205,9 +213,11 @@ public static class Builder { private String endpoint; private HttpTransport httpTransport; private boolean usingDefaultExecutor; + private boolean usingDefaultTransport; private Builder() { this.usingDefaultExecutor = false; + this.usingDefaultTransport = false; } public Builder setExecutor(Executor executor) { @@ -225,6 +235,11 @@ public Builder setHttpTransport(HttpTransport httpTransport) { return this; } + Builder setManageHttpTransport(boolean manageHttpTransport) { + this.usingDefaultTransport = manageHttpTransport; + return this; + } + public ManagedHttpJsonChannel build() { Preconditions.checkNotNull(endpoint); @@ -237,14 +252,8 @@ public ManagedHttpJsonChannel build() { usingDefaultExecutor = true; } - if (httpTransport == null) { - httpTransport = - HttpJsonConscryptUtils.configureConscryptSecurityProvider( - new NetHttpTransport.Builder()) - .build(); - } - - return new ManagedHttpJsonChannel(executor, usingDefaultExecutor, endpoint, httpTransport); + return new ManagedHttpJsonChannel( + executor, usingDefaultExecutor, endpoint, httpTransport, usingDefaultTransport); } } } diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/ManagedHttpJsonInterceptorChannel.java b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/ManagedHttpJsonInterceptorChannel.java index eaaa8c3a7c56..e552608b6529 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/ManagedHttpJsonInterceptorChannel.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/ManagedHttpJsonInterceptorChannel.java @@ -29,7 +29,9 @@ */ package com.google.api.gax.httpjson; +import com.google.api.client.http.HttpTransport; import com.google.common.annotations.VisibleForTesting; +import java.util.concurrent.Executor; import java.util.concurrent.TimeUnit; import org.jspecify.annotations.NullMarked; @@ -51,12 +53,39 @@ ManagedHttpJsonChannel getChannel() { return channel; } + @Override + String getEndpoint() { + return channel.getEndpoint(); + } + + @Override + @VisibleForTesting + HttpTransport getHttpTransport() { + return channel.getHttpTransport(); + } + + @Override + @VisibleForTesting + Executor getExecutor() { + return channel.getExecutor(); + } + @Override public HttpJsonClientCall newCall( ApiMethodDescriptor methodDescriptor, HttpJsonCallOptions callOptions) { return interceptor.interceptCall(methodDescriptor, callOptions, channel); } + @Override + public void refresh() { + channel.refresh(); + } + + @Override + public boolean shouldRefresh() { + return channel.shouldRefresh(); + } + @Override public synchronized void shutdown() { channel.shutdown(); diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java new file mode 100644 index 000000000000..0b4c3aa885fd --- /dev/null +++ b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java @@ -0,0 +1,369 @@ +/* + * Copyright 2026 Google LLC + * + * Redistribution and use in source and binary forms, with or without + * modification, are permitted provided that the following conditions are + * met: + * + * * Redistributions of source code must retain the above copyright + * notice, this list of conditions and the following disclaimer. + * * Redistributions in binary form must reproduce the above + * copyright notice, this list of conditions and the following disclaimer + * in the documentation and/or other materials provided with the + * distribution. + * * Neither the name of Google LLC nor the names of its + * contributors may be used to endorse or promote products derived from + * this software without specific prior written permission. + * + * THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS + * "AS IS" AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT + * LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR + * A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT + * OWNER OR CONTRIBUTORS BE LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, + * SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT + * LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, + * DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY + * THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT + * (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE + * OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + */ +package com.google.api.gax.httpjson; + +import com.google.api.client.http.HttpTransport; +import com.google.api.core.InternalApi; +import com.google.api.gax.httpjson.ForwardingHttpJsonClientCall.SimpleForwardingHttpJsonClientCall; +import com.google.api.gax.httpjson.ForwardingHttpJsonClientCallListener.SimpleForwardingHttpJsonClientCallListener; +import com.google.api.gax.rpc.mtls.WorkloadCertificateUtils; +import com.google.common.annotations.VisibleForTesting; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicBoolean; +import java.util.concurrent.atomic.AtomicInteger; +import java.util.concurrent.atomic.AtomicReference; +import java.util.function.Supplier; +import java.util.logging.Level; +import java.util.logging.Logger; + +/** + * An implementation of {@link ManagedHttpJsonChannel} that supports dynamic mTLS certificate + * rotation by thread-safely hot-swapping the underlying active HTTP/JSON channel while gracefully + * retiring older connections after all active in-flight requests complete. + */ +@InternalApi +public class RefreshingHttpJsonChannel extends ManagedHttpJsonChannel { + + private static final Logger LOG = Logger.getLogger(RefreshingHttpJsonChannel.class.getName()); + + private static class DiskCheckResult { + final String fingerprint; + final long timestampNanos; + + DiskCheckResult(String fingerprint, long timestampNanos) { + this.fingerprint = fingerprint; + this.timestampNanos = timestampNanos; + } + } + + private volatile DiskCheckResult lastDiskCheck = null; + private final java.util.concurrent.locks.ReentrantLock diskCheckLock = + new java.util.concurrent.locks.ReentrantLock(); + private final Supplier channelFactory; + private final String workloadCertPath; + private final AtomicReference activeEntry; + // Keep track of all entries to properly await their termination + private final java.util.concurrent.ConcurrentLinkedQueue allEntries = + new java.util.concurrent.ConcurrentLinkedQueue<>(); + private final Object refreshLock = new Object(); + private volatile String activeCertFingerprint = ""; + + public RefreshingHttpJsonChannel( + Supplier channelFactory, String workloadCertPath) { + this.channelFactory = channelFactory; + this.workloadCertPath = workloadCertPath; + ChannelEntry initial = new ChannelEntry(channelFactory.get()); + this.activeEntry = new AtomicReference<>(initial); + this.allEntries.add(initial); + if (workloadCertPath != null) { + this.activeCertFingerprint = getCertificateFingerprint(workloadCertPath); + } + } + + private String getOrUpdateDiskFingerprint(String certPath) { + long now = System.nanoTime(); + DiskCheckResult cached = lastDiskCheck; + if (cached != null + && (now - cached.timestampNanos < java.util.concurrent.TimeUnit.SECONDS.toNanos(1))) { + return cached.fingerprint; + } + + diskCheckLock.lock(); + try { + cached = lastDiskCheck; + if (cached != null + && (now - cached.timestampNanos < java.util.concurrent.TimeUnit.SECONDS.toNanos(1))) { + return cached.fingerprint; + } + String fingerprint = getCertificateFingerprint(certPath); + lastDiskCheck = new DiskCheckResult(fingerprint, System.nanoTime()); + return fingerprint; + } finally { + diskCheckLock.unlock(); + } + } + + // Visible for testing + protected String getWorkloadCertPath() { + return workloadCertPath; + } + + // Visible for testing + protected String getCertificateFingerprint(String certPath) { + return WorkloadCertificateUtils.getCertificateFingerprint(certPath); + } + + @Override + public boolean shouldRefresh() { + String certPath = getWorkloadCertPath(); + if (certPath == null) { + return false; + } + String currentDiskFingerprint = getOrUpdateDiskFingerprint(certPath); + if (currentDiskFingerprint.isEmpty()) { + return false; + } + return !currentDiskFingerprint.equalsIgnoreCase(activeCertFingerprint); + } + + @Override + public void refresh() { + synchronized (refreshLock) { + if (isShutdown()) { + return; + } + String certPath = getWorkloadCertPath(); + if (certPath == null) { + return; + } + String currentDiskFingerprint = getOrUpdateDiskFingerprint(certPath); + if (currentDiskFingerprint.isEmpty()) { + return; + } + + // Double-check inside refreshLock + if (currentDiskFingerprint.equalsIgnoreCase(this.activeCertFingerprint)) { + LOG.fine( + "HTTP/JSON channel was already refreshed by a concurrent thread, skipping duplicate" + + " refresh"); + return; + } + + LOG.info("mTLS certificate rotation detected. Triggering HTTP/JSON channel pool refresh."); + + // Prune terminated entries to prevent memory leak + allEntries.removeIf(entry -> entry.channel.isTerminated()); + + ChannelEntry newEntry = new ChannelEntry(channelFactory.get()); + allEntries.add(newEntry); + ChannelEntry oldEntry = activeEntry.getAndSet(newEntry); + this.activeCertFingerprint = currentDiskFingerprint; + + if (oldEntry != null) { + oldEntry.requestShutdown(); + } + } + } + + private ChannelEntry getRetainedEntry() { + while (true) { + ChannelEntry entry = activeEntry.get(); + if (entry.retain()) { + return entry; + } + if (entry == activeEntry.get()) { + throw new IllegalStateException("Channel has been shut down"); + } + } + } + + @Override + public HttpJsonClientCall newCall( + ApiMethodDescriptor methodDescriptor, HttpJsonCallOptions callOptions) { + ChannelEntry entry = getRetainedEntry(); + try { + HttpJsonClientCall delegateCall = + entry.channel.newCall(methodDescriptor, callOptions); + return new ReleasingHttpJsonClientCall<>(delegateCall, entry); + } catch (Exception e) { + entry.release(); + throw e; + } + } + + @Override + java.util.concurrent.Executor getExecutor() { + return activeEntry.get().channel.getExecutor(); + } + + @VisibleForTesting + ManagedHttpJsonChannel getActiveChannel() { + return activeEntry.get().channel; + } + + private volatile boolean isShuttingDown = false; + + @Override + public void shutdown() { + synchronized (refreshLock) { + isShuttingDown = true; + for (ChannelEntry entry : allEntries) { + entry.requestShutdown(); + } + } + } + + @Override + public boolean isShutdown() { + return isShuttingDown; + } + + @Override + public boolean isTerminated() { + for (ChannelEntry entry : allEntries) { + if (!entry.channel.isTerminated()) { + return false; + } + } + return true; + } + + @Override + public void shutdownNow() { + synchronized (refreshLock) { + isShuttingDown = true; + for (ChannelEntry entry : allEntries) { + entry.requestShutdown(); + entry.channel.shutdownNow(); + } + } + } + + @Override + public boolean awaitTermination(long duration, TimeUnit unit) throws InterruptedException { + long endNanos = System.nanoTime() + unit.toNanos(duration); + for (ChannelEntry entry : allEntries) { + if (entry.channel.isTerminated()) { + continue; + } + long remainingNanos = endNanos - System.nanoTime(); + if (remainingNanos <= 0) { + return false; + } + if (!entry.channel.awaitTermination(remainingNanos, TimeUnit.NANOSECONDS)) { + return false; + } + } + return true; + } + + @Override + public void close() { + shutdown(); + } + + @Override + String getEndpoint() { + return activeEntry.get().channel.getEndpoint(); + } + + @Override + @VisibleForTesting + HttpTransport getHttpTransport() { + return activeEntry.get().channel.getHttpTransport(); + } + + /** Internal container to manage request reference-counting and graceful shutdown. */ + private static class ChannelEntry { + private final ManagedHttpJsonChannel channel; + private final AtomicInteger outstandingCalls = new AtomicInteger(0); + private final AtomicBoolean shutdownRequested = new AtomicBoolean(false); + private final AtomicBoolean shutdownInitiated = new AtomicBoolean(false); + + ChannelEntry(ManagedHttpJsonChannel channel) { + this.channel = channel; + } + + boolean retain() { + outstandingCalls.incrementAndGet(); + if (shutdownRequested.get()) { + release(); + return false; + } + return true; + } + + void release() { + int count = outstandingCalls.decrementAndGet(); + if (shutdownRequested.get() && count == 0) { + shutdown(); + } + } + + void requestShutdown() { + shutdownRequested.set(true); + if (outstandingCalls.get() == 0) { + shutdown(); + } + } + + private void shutdown() { + if (shutdownInitiated.compareAndSet(false, true)) { + try { + channel.shutdown(); + } catch (Exception e) { + LOG.log(Level.WARNING, "Error shutting down retired HTTP/JSON channel", e); + } + } + } + } + + /** A client call decorator that decrements the entry counter upon call completion. */ + private static class ReleasingHttpJsonClientCall + extends SimpleForwardingHttpJsonClientCall { + + private final ChannelEntry entry; + private final AtomicBoolean wasClosed = new AtomicBoolean(false); + private final AtomicBoolean wasReleased = new AtomicBoolean(false); + + ReleasingHttpJsonClientCall(HttpJsonClientCall delegate, ChannelEntry entry) { + super(delegate); + this.entry = entry; + } + + @Override + public void start(Listener responseListener, HttpJsonMetadata requestHeaders) { + try { + super.start( + new SimpleForwardingHttpJsonClientCallListener(responseListener) { + @Override + public void onClose(int statusCode, HttpJsonMetadata trailers) { + if (!wasClosed.compareAndSet(false, true)) { + return; + } + try { + super.onClose(statusCode, trailers); + } finally { + if (wasReleased.compareAndSet(false, true)) { + entry.release(); + } + } + } + }, + requestHeaders); + } catch (Exception e) { + if (wasReleased.compareAndSet(false, true)) { + entry.release(); + } + throw e; + } + } + } +} diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProviderTest.java b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProviderTest.java index 8c95c1d2e1c4..4482f2367a4a 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProviderTest.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProviderTest.java @@ -31,8 +31,9 @@ import static com.google.common.truth.Truth.assertThat; import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.mockito.Mockito.mock; -import com.google.api.client.http.javanet.NetHttpTransport; +import com.google.api.gax.rpc.HeaderProvider; import com.google.api.gax.rpc.TransportChannelProvider; import com.google.api.gax.rpc.mtls.AbstractMtlsTransportChannelTest; import com.google.api.gax.rpc.mtls.CertificateBasedAccess; @@ -46,6 +47,7 @@ import java.util.concurrent.ScheduledThreadPoolExecutor; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; +import org.mockito.Mockito; class InstantiatingHttpJsonChannelProviderTest extends AbstractMtlsTransportChannelTest { @@ -55,9 +57,10 @@ class InstantiatingHttpJsonChannelProviderTest extends AbstractMtlsTransportChan @BeforeEach public void setup() throws IOException { - certificateBasedAccess = - new CertificateBasedAccess( - name -> name.equals("GOOGLE_API_USE_MTLS_ENDPOINT") ? "never" : "false"); + certificateBasedAccess = org.mockito.Mockito.mock(CertificateBasedAccess.class); + org.mockito.Mockito.when(certificateBasedAccess.getMtlsEndpointUsagePolicy()) + .thenReturn(CertificateBasedAccess.MtlsEndpointUsagePolicy.NEVER); + org.mockito.Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(false); } @Test @@ -179,6 +182,39 @@ void managedChannelUsesCustomExecutor() throws IOException { instantiatingHttpJsonChannelProvider.getTransportChannel().shutdownNow(); } + @Test + void managedChannelDoesNotShutdownCustomHttpTransport() throws IOException { + com.google.api.client.http.HttpTransport mockHttpTransport = + org.mockito.Mockito.mock(com.google.api.client.http.HttpTransport.class); + + InstantiatingHttpJsonChannelProvider provider = + InstantiatingHttpJsonChannelProvider.newBuilder() + .setEndpoint(DEFAULT_ENDPOINT) + .setHttpTransport(mockHttpTransport) + .setCertificateBasedAccess(certificateBasedAccess) + .build(); + provider = (InstantiatingHttpJsonChannelProvider) provider.withHeaders(DEFAULT_HEADER_MAP); + + HttpJsonTransportChannel httpJsonTransportChannel = provider.getTransportChannel(); + + // Verify custom transport is injected + ManagedHttpJsonInterceptorChannel interceptorChannel = + (ManagedHttpJsonInterceptorChannel) httpJsonTransportChannel.getManagedChannel(); + ManagedHttpJsonInterceptorChannel managedHttpJsonChannel = + (ManagedHttpJsonInterceptorChannel) interceptorChannel.getChannel(); + RefreshingHttpJsonChannel refreshingHttpJsonChannel = + (RefreshingHttpJsonChannel) managedHttpJsonChannel.getChannel(); + ManagedHttpJsonChannel channel = refreshingHttpJsonChannel.getActiveChannel(); + + assertThat(channel.getHttpTransport()).isEqualTo(mockHttpTransport); + + // Perform a shutdown + provider.getTransportChannel().shutdownNow(); + + // Verify that shutdown() was NOT called on the custom HttpTransport + org.mockito.Mockito.verify(mockHttpTransport, org.mockito.Mockito.never()).shutdown(); + } + @Override protected Object getMtlsObjectFromTransportChannel( MtlsProvider provider, CertificateBasedAccess certificateBasedAccess) @@ -188,30 +224,10 @@ protected Object getMtlsObjectFromTransportChannel( .setEndpoint("localhost:8080") .setMtlsProvider(provider) .setCertificateBasedAccess(certificateBasedAccess) - .setHeaderProvider(Collections::emptyMap) - .setExecutor(Runnable::run) - .build(); - NetHttpTransport transport = (NetHttpTransport) channelProvider.createHttpTransport(); - return (transport != null && transport.isMtls()) ? transport : null; - } - - @Test - void testCreateHttpTransport_returnsValidTransport() throws Exception { - InstantiatingHttpJsonChannelProvider channelProvider = - InstantiatingHttpJsonChannelProvider.newBuilder() - .setEndpoint("localhost:8080") - .setHeaderProvider(Collections::emptyMap) - .setExecutor(Runnable::run) + .setHeaderProvider( + mock(HeaderProvider.class, Mockito.withSettings().withoutAnnotations())) + .setExecutor(mock(Executor.class)) .build(); - NetHttpTransport transport = (NetHttpTransport) channelProvider.createHttpTransport(); - assertThat(transport).isNotNull(); - } - - @Test - void testConfigureConscryptSecurityProvider_returnsConfiguredBuilder() { - NetHttpTransport.Builder builder = new NetHttpTransport.Builder(); - NetHttpTransport.Builder result = - HttpJsonConscryptUtils.configureConscryptSecurityProvider(builder); - assertThat(result).isSameInstanceAs(builder); + return channelProvider.createHttpTransport(); } } diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java new file mode 100644 index 000000000000..69e2b1c69113 --- /dev/null +++ b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java @@ -0,0 +1,327 @@ +/* + * Copyright 2026 Google LLC + * + * Redistribution and use in source and binary forms, with or without + * modification, are permitted provided that the following conditions are + * met: + * + * * Redistributions of source code must retain the above copyright + * notice, this list of conditions and the following disclaimer. + * * Redistributions in binary form must reproduce the above + * copyright notice, this list of conditions and the following disclaimer + * in the documentation and/or other materials provided with the + * distribution. + * * Neither the name of Google LLC nor the names of its + * contributors may be used to endorse or promote products derived from + * this software without specific prior written permission. + * + * THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS + * "AS IS" AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT + * LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR + * A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT + * OWNER OR CONTRIBUTORS BE LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, + * SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT + * LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, + * DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY + * THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT + * (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE + * OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + */ +package com.google.api.gax.httpjson; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import java.util.ArrayList; +import java.util.List; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicInteger; +import java.util.function.Supplier; +import javax.annotation.Nullable; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; + +class RefreshingHttpJsonChannelTest { + private static class FakeHttpJsonClientCall + extends HttpJsonClientCall { + private Listener listener; + + @Override + public void start(Listener responseListener, HttpJsonMetadata requestHeaders) { + this.listener = responseListener; + } + + @Override + public void request(int numMessages) {} + + @Override + public void cancel(@Nullable String message, @Nullable Throwable cause) {} + + @Override + public void sendMessage(RequestT message) {} + + @Override + public void halfClose() {} + } + + private static class FakeManagedHttpJsonChannel extends ManagedHttpJsonChannel { + private volatile boolean isShutdown = false; + private volatile boolean isTerminated = false; + private HttpJsonClientCall nextCall = null; + + @Override + String getEndpoint() { + return "https://fake.endpoint:443"; + } + + @Override + public void shutdown() { + isShutdown = true; + } + + @Override + public void shutdownNow() { + isShutdown = true; + isTerminated = true; + } + + @Override + public boolean isShutdown() { + return isShutdown; + } + + @Override + public boolean isTerminated() { + return isTerminated; + } + + @Override + public boolean awaitTermination(long duration, TimeUnit unit) { + return isTerminated; + } + + @Override + @SuppressWarnings("unchecked") + public HttpJsonClientCall newCall( + ApiMethodDescriptor methodDescriptor, + HttpJsonCallOptions callOptions) { + if (nextCall != null) { + return (HttpJsonClientCall) nextCall; + } + return new FakeHttpJsonClientCall<>(); + } + } + + private AtomicInteger channelFactoryCount; + private FakeManagedHttpJsonChannel lastCreatedChannel; + private String testCertPath = "/fake/path"; + private String testFingerprint = "fingerprint1"; + private boolean shouldThrowOnFactory = false; + private List createdChannels; + + private Supplier channelFactory = + () -> { + if (shouldThrowOnFactory) { + throw new RuntimeException("Simulated factory failure"); + } + channelFactoryCount.incrementAndGet(); + lastCreatedChannel = new FakeManagedHttpJsonChannel(); + return lastCreatedChannel; + }; + + @BeforeEach + void setUp() { + channelFactoryCount = new AtomicInteger(0); + testCertPath = "/fake/path"; + testFingerprint = "fingerprint1"; + shouldThrowOnFactory = false; + createdChannels = new ArrayList<>(); + } + + @AfterEach + void tearDown() { + for (RefreshingHttpJsonChannel channel : createdChannels) { + channel.shutdownNow(); + } + } + + private RefreshingHttpJsonChannel createTestChannel() { + RefreshingHttpJsonChannel ch = + new RefreshingHttpJsonChannel(channelFactory, "fake/cert/path.json") { + @Override + protected String getWorkloadCertPath() { + return testCertPath; + } + + @Override + protected String getCertificateFingerprint(String certPath) { + return testFingerprint; + } + }; + createdChannels.add(ch); + return ch; + } + + @Test + void testShouldRefreshNullCertPath() { + testCertPath = null; + RefreshingHttpJsonChannel channel = createTestChannel(); + assertFalse(channel.shouldRefresh()); + } + + @Test + void testShouldRefreshFalseWhenUnchanged() throws InterruptedException { + RefreshingHttpJsonChannel channel = createTestChannel(); + + Thread.sleep(1001); // Invalidate 1-second cache + assertFalse(channel.shouldRefresh()); + } + + @Test + void testShouldRefreshTrueWhenChanged() throws InterruptedException { + RefreshingHttpJsonChannel channel = createTestChannel(); + + Thread.sleep(1001); // Invalidate 1-second cache + + // Simulate disk fingerprint changing + testFingerprint = "fingerprint2"; + + assertTrue(channel.shouldRefresh()); + } + + @Test + void testRefreshSwapsChannel() throws InterruptedException { + RefreshingHttpJsonChannel channel = createTestChannel(); + FakeManagedHttpJsonChannel firstChannel = lastCreatedChannel; + assertEquals(1, channelFactoryCount.get()); + + Thread.sleep(1001); // Invalidate 1-second cache + + // Change fingerprint + testFingerprint = "fingerprint2"; + + // Act + channel.refresh(); + + // Verify a new channel was created and the old one retired + assertEquals(2, channelFactoryCount.get()); + FakeManagedHttpJsonChannel secondChannel = lastCreatedChannel; + + // The old channel should receive a shutdown request immediately since there are no active calls + assertTrue(firstChannel.isShutdown()); + assertFalse(secondChannel.isShutdown()); + } + + @Test + void testRefreshKeepsInFlightChannelsAlive() throws InterruptedException { + RefreshingHttpJsonChannel channel = createTestChannel(); + FakeManagedHttpJsonChannel firstChannel = lastCreatedChannel; + + // Simulate an in-flight API call + FakeHttpJsonClientCall fakeCall = new FakeHttpJsonClientCall<>(); + firstChannel.nextCall = fakeCall; + + HttpJsonClientCall activeCall = channel.newCall(null, null); + + Thread.sleep(1001); // Invalidate 1-second cache + + // Change fingerprint & refresh + testFingerprint = "fingerprint2"; + + channel.refresh(); + + // Verify a new channel was created + assertEquals(2, channelFactoryCount.get()); + + // IMPORTANT: The first channel should NOT be shut down yet because of the active call! + assertFalse(firstChannel.isShutdown()); + + // Now start the call + activeCall.start(new HttpJsonClientCall.Listener() {}, null); + + assertNotNull(fakeCall.listener); + + // Fire onClose + fakeCall.listener.onClose(0, null); + + // FIRST CHANNEL SHOULD BE SHUT DOWN NOW! + assertTrue(firstChannel.isShutdown()); + } + + @Test + void testRefreshDoesNotSpawnChannelWhenShutdown() throws InterruptedException { + RefreshingHttpJsonChannel channel = createTestChannel(); + FakeManagedHttpJsonChannel firstChannel = lastCreatedChannel; + assertEquals(1, channelFactoryCount.get()); + + // Simulate that the channel pool is shut down. + channel.shutdown(); + firstChannel.shutdown(); + + Thread.sleep(1001); // Invalidate 1-second cache + + // Change fingerprint + testFingerprint = "fingerprint2"; + + // Act + channel.refresh(); + + // Verify no new channel was spawned + assertEquals(1, channelFactoryCount.get()); + } + + @Test + void testRefreshFactoryExceptionDoesNotWedgeFingerprint() throws InterruptedException { + RefreshingHttpJsonChannel channel = createTestChannel(); + assertEquals(1, channelFactoryCount.get()); + + shouldThrowOnFactory = true; + Thread.sleep(1001); // Invalidate 1-second cache + testFingerprint = "fingerprint2"; + + assertThrows(RuntimeException.class, channel::refresh); + + // Because factory threw, activeCertFingerprint should NOT be updated to fingerprint2 + // Therefore shouldRefresh() should still return true + assertTrue(channel.shouldRefresh()); + + shouldThrowOnFactory = false; + channel.refresh(); + assertEquals(2, channelFactoryCount.get()); + assertFalse(channel.shouldRefresh()); + } + + @Test + void testShutdownNowSetsIsShutdown() { + RefreshingHttpJsonChannel channel = createTestChannel(); + assertFalse(channel.isShutdown()); + + channel.shutdownNow(); + + assertTrue(channel.isShutdown()); + } + + @Test + void testAwaitTerminationZeroTimeoutOnTerminatedChannelReturnsTrue() throws InterruptedException { + RefreshingHttpJsonChannel channel = createTestChannel(); + FakeManagedHttpJsonChannel firstChannel = lastCreatedChannel; + firstChannel.isTerminated = true; + + channel.shutdown(); + assertTrue(channel.awaitTermination(0, TimeUnit.MILLISECONDS)); + } + + @Test + void testChannelDelegationMethods() { + RefreshingHttpJsonChannel channel = createTestChannel(); + FakeManagedHttpJsonChannel firstChannel = lastCreatedChannel; + + assertEquals(firstChannel.getEndpoint(), channel.getEndpoint()); + assertEquals(firstChannel.getHttpTransport(), channel.getHttpTransport()); + assertEquals(firstChannel.getExecutor(), channel.getExecutor()); + } +} diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ApiCallContext.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ApiCallContext.java index 67b5e5b285d1..36d6fd966e55 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ApiCallContext.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ApiCallContext.java @@ -42,7 +42,6 @@ import java.util.Map; import java.util.Set; import javax.annotation.Nonnull; -import org.jspecify.annotations.NullMarked; import org.jspecify.annotations.Nullable; /** @@ -55,7 +54,6 @@ * *

This is transport specific and each transport has an implementation with its own options. */ -@NullMarked @InternalExtensionOnly public interface ApiCallContext extends RetryingContext { @@ -65,6 +63,18 @@ public interface ApiCallContext extends RetryingContext { /** Returns a new ApiCallContext with the given channel set. */ ApiCallContext withTransportChannel(TransportChannel channel); + /** + * Returns the {@link TransportChannel} associated with this call context, or {@code null} if none + * is set. + * + *

Note: By default, this method returns {@code null}. If an implementation does not override + * this method, automatic mTLS certificate rotation and channel refreshing in retrying callables + * will be disabled. + */ + default TransportChannel getTransportChannel() { + return null; + } + /** Returns a new ApiCallContext with the given Endpoint Context. */ ApiCallContext withEndpointContext(EndpointContext endpointContext); diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ApiResultRetryAlgorithm.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ApiResultRetryAlgorithm.java index 7d04d38d2605..97d24d441ad2 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ApiResultRetryAlgorithm.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ApiResultRetryAlgorithm.java @@ -37,7 +37,6 @@ @NullMarked class ApiResultRetryAlgorithm extends BasicResultRetryAlgorithm { - /** Returns true if previousThrowable is an {@link ApiException} that is retryable. */ @Override public boolean shouldRetry(Throwable previousThrowable, ResponseT previousResponse) { return (previousThrowable instanceof ApiException) @@ -53,6 +52,12 @@ public boolean shouldRetry(Throwable previousThrowable, ResponseT previousRespon @Override public boolean shouldRetry( RetryingContext context, Throwable previousThrowable, ResponseT previousResponse) { + // Check UnauthenticatedException retryability first to ensure mTLS certificate + // rotation retries take precedence over static method retry codes. + if (previousThrowable instanceof UnauthenticatedException + && ((UnauthenticatedException) previousThrowable).isRetryable()) { + return true; + } if (context.getRetryableCodes() != null) { // Ignore the isRetryable() value of the throwable if the RetryingContext has a specific list // of codes that should be retried. diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/AttemptCallable.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/AttemptCallable.java index 897e1f2dae8e..bd10a76cd632 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/AttemptCallable.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/AttemptCallable.java @@ -86,7 +86,33 @@ public ResponseT call() { .attemptStarted(request, externalFuture.getAttemptSettings().getOverallAttemptCount()); ApiFuture internalFuture = callable.futureCall(request, callContext); - externalFuture.setAttemptFuture(internalFuture); + final ApiCallContext finalContext = callContext; + ApiFuture mappedFuture = + ApiFutures.catching( + internalFuture, + UnauthenticatedException.class, + unauthenticatedException -> { + TransportChannel transportChannel = finalContext.getTransportChannel(); + if (transportChannel != null && transportChannel.shouldRefresh()) { + transportChannel.refresh(); + UnauthenticatedException newEx = + new UnauthenticatedException( + unauthenticatedException.getMessage(), + unauthenticatedException, + unauthenticatedException.getStatusCode(), + true, // isRetryable = true + unauthenticatedException.getErrorDetails()); + newEx.setStackTrace(unauthenticatedException.getStackTrace()); + for (Throwable suppressed : unauthenticatedException.getSuppressed()) { + newEx.addSuppressed(suppressed); + } + throw newEx; + } + throw unauthenticatedException; + }, + com.google.common.util.concurrent.MoreExecutors.directExecutor()); + + externalFuture.setAttemptFuture(mappedFuture); } catch (Throwable e) { externalFuture.setAttemptFuture(ApiFutures.immediateFailedFuture(e)); } diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/BidiStreamingCallable.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/BidiStreamingCallable.java index 97e22d6ee41f..0ad3a01af672 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/BidiStreamingCallable.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/BidiStreamingCallable.java @@ -241,11 +241,48 @@ public BidiStreamingCallable withDefaultCallContext( return new BidiStreamingCallable() { @Override public ClientStream internalCall( - ResponseObserver responseObserver, + final ResponseObserver responseObserver, ClientStreamReadyObserver onReady, ApiCallContext thisCallContext) { - return BidiStreamingCallable.this.internalCall( - responseObserver, onReady, defaultCallContext.merge(thisCallContext)); + final ApiCallContext mergedContext = defaultCallContext.merge(thisCallContext); + ResponseObserver refreshingObserver = + new ResponseObserver() { + @Override + public void onStart(StreamController controller) { + responseObserver.onStart(controller); + } + + @Override + public void onResponse(ResponseT response) { + responseObserver.onResponse(response); + } + + @Override + public void onError(Throwable t) { + if (t instanceof UnauthenticatedException) { + TransportChannel transportChannel = mergedContext.getTransportChannel(); + if (transportChannel != null && transportChannel.shouldRefresh()) { + transportChannel.refresh(); + UnauthenticatedException causeEx = (UnauthenticatedException) t; + t = + new UnauthenticatedException( + causeEx.getMessage(), + causeEx.getCause(), + causeEx.getStatusCode(), + true, // isRetryable = true + causeEx.getErrorDetails()); + } + } + responseObserver.onError(t); + } + + @Override + public void onComplete() { + responseObserver.onComplete(); + } + }; + + return BidiStreamingCallable.this.internalCall(refreshingObserver, onReady, mergedContext); } }; } diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ClientStreamingCallable.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ClientStreamingCallable.java index 854ce3fd5de8..0119a3a40232 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ClientStreamingCallable.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ClientStreamingCallable.java @@ -77,9 +77,41 @@ public ClientStreamingCallable withDefaultCallContext( return new ClientStreamingCallable() { @Override public ApiStreamObserver clientStreamingCall( - ApiStreamObserver responseObserver, ApiCallContext thisCallContext) { - return ClientStreamingCallable.this.clientStreamingCall( - responseObserver, defaultCallContext.merge(thisCallContext)); + final ApiStreamObserver responseObserver, ApiCallContext thisCallContext) { + final ApiCallContext mergedContext = defaultCallContext.merge(thisCallContext); + ApiStreamObserver refreshingObserver = + new ApiStreamObserver() { + @Override + public void onNext(ResponseT response) { + responseObserver.onNext(response); + } + + @Override + public void onError(Throwable t) { + if (t instanceof UnauthenticatedException) { + TransportChannel transportChannel = mergedContext.getTransportChannel(); + if (transportChannel != null && transportChannel.shouldRefresh()) { + transportChannel.refresh(); + UnauthenticatedException causeEx = (UnauthenticatedException) t; + t = + new UnauthenticatedException( + causeEx.getMessage(), + causeEx.getCause(), + causeEx.getStatusCode(), + true, // isRetryable = true + causeEx.getErrorDetails()); + } + } + responseObserver.onError(t); + } + + @Override + public void onCompleted() { + responseObserver.onCompleted(); + } + }; + + return ClientStreamingCallable.this.clientStreamingCall(refreshingObserver, mergedContext); } }; } diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ServerStreamingAttemptCallable.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ServerStreamingAttemptCallable.java index af6e8014e38a..de7c18d7a8ce 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ServerStreamingAttemptCallable.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ServerStreamingAttemptCallable.java @@ -221,6 +221,7 @@ public Void call() { .getTracer() .attemptStarted(request, outerRetryingFuture.getAttemptSettings().getOverallAttemptCount()); + final ApiCallContext finalContext = attemptContext; innerCallable.call( request, new StateCheckingResponseObserver() { @@ -236,6 +237,26 @@ public void onResponseImpl(ResponseT response) { @Override public void onErrorImpl(Throwable t) { + Throwable cause = t; + if (cause instanceof com.google.api.gax.retrying.ServerStreamingAttemptException) { + cause = cause.getCause(); + } + if (cause instanceof UnauthenticatedException) { + TransportChannel transportChannel = finalContext.getTransportChannel(); + if (transportChannel != null && transportChannel.shouldRefresh()) { + transportChannel.refresh(); + UnauthenticatedException causeEx = (UnauthenticatedException) cause; + cause = + new UnauthenticatedException( + causeEx.getMessage(), + causeEx.getCause(), + causeEx.getStatusCode(), + true, // isRetryable = true + causeEx.getErrorDetails()); + + t = cause; + } + } onAttemptError(t); } diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/TransportChannel.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/TransportChannel.java index 1866092e28f9..de83dbc73861 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/TransportChannel.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/TransportChannel.java @@ -31,10 +31,8 @@ import com.google.api.core.InternalExtensionOnly; import com.google.api.gax.core.BackgroundResource; -import org.jspecify.annotations.NullMarked; /** Class whose instances can issue RPCs on a particular transport. */ -@NullMarked @InternalExtensionOnly public interface TransportChannel extends BackgroundResource { @@ -49,4 +47,20 @@ public interface TransportChannel extends BackgroundResource { * Returns an empty {@link ApiCallContext} that is compatible with this {@code TransportChannel}. */ ApiCallContext getEmptyCallContext(); + + /** + * Refreshes or recreates the underlying network connections of this transport channel. + * + *

By default, this is a no-op for transports that do not require stateful connection lifecycle + * management. + */ + default void refresh() {} + + /** + * Returns true if a certificate rotation has been detected on disk and the transport channel + * should be refreshed, or false otherwise. + */ + default boolean shouldRefresh() { + return false; + } } diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/CertificateBasedAccess.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/CertificateBasedAccess.java index 99f2cca9b53d..de8bacd7c02b 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/CertificateBasedAccess.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/CertificateBasedAccess.java @@ -32,23 +32,51 @@ import com.google.api.core.InternalApi; import com.google.api.gax.rpc.internal.EnvironmentProvider; -import org.jspecify.annotations.NullMarked; +import java.io.IOException; /** * Utility class for handling certificate-based access configurations. * - *

This class handles the processing of GOOGLE_API_USE_CLIENT_CERTIFICATE and - * GOOGLE_API_USE_MTLS_ENDPOINT environment variables according to https://google.aip.dev/auth/4114 + *

This class handles the processing of GOOGLE_API_USE_CLIENT_CERTIFICATE, + * GOOGLE_API_CERTIFICATE_CONFIG, and GOOGLE_API_USE_MTLS_ENDPOINT configurations. */ -@NullMarked @InternalApi public class CertificateBasedAccess { private final EnvironmentProvider envProvider; + private final FileExistenceProvider fileExistenceProvider; + private final FileContentReader fileContentReader; + + @InternalApi + public interface FileExistenceProvider { + boolean exists(String path); + } + + @InternalApi + public interface FileContentReader { + String read(String path) throws IOException; + } - /** The EnvironmentProvider mechanism supports env var injection for unit tests. */ public CertificateBasedAccess(EnvironmentProvider envProvider) { + this( + envProvider, + path -> { + java.io.File file = new java.io.File(path); + return file.exists() && file.isFile(); + }, + path -> + new String( + java.nio.file.Files.readAllBytes(java.nio.file.Paths.get(path)), + java.nio.charset.StandardCharsets.UTF_8)); + } + + CertificateBasedAccess( + EnvironmentProvider envProvider, + FileExistenceProvider fileExistenceProvider, + FileContentReader fileContentReader) { this.envProvider = envProvider; + this.fileExistenceProvider = fileExistenceProvider; + this.fileContentReader = fileContentReader; } public static CertificateBasedAccess createWithSystemEnv() { @@ -66,10 +94,102 @@ public enum MtlsEndpointUsagePolicy { ALWAYS; } + private static class CertificateConfig { + final String certPath; + final String keyPath; + + CertificateConfig(String certPath, String keyPath) { + this.certPath = certPath; + this.keyPath = keyPath; + } + } + + private CertificateConfig parseCertificateConfig(String configPath) throws IOException { + String content = fileContentReader.read(configPath); + + String certPath = extractJsonValue(content, "cert_path"); + String keyPath = extractJsonValue(content, "key_path"); + + if (certPath == null || keyPath == null) { + throw new IllegalStateException( + "Malformed certificate config JSON. Must contain 'cert_path' and 'key_path'."); + } + + return new CertificateConfig(certPath, keyPath); + } + + private String extractJsonValue(String json, String key) { + java.util.regex.Pattern pattern = + java.util.regex.Pattern.compile( + "\"" + java.util.regex.Pattern.quote(key) + "\"\\s*:\\s*\"((?:[^\\\\\"]|\\\\.)*)\""); + java.util.regex.Matcher matcher = pattern.matcher(json); + if (matcher.find()) { + return matcher.group(1).replace("\\\\", "\\").replace("\\/", "/").replace("\\\"", "\""); + } + return null; + } + /** Returns if mutual TLS client certificate should be used. */ public boolean useMtlsClientCertificate() { + // 1. Check the explicit user flag first (Primary override) String useClientCertificate = envProvider.getenv("GOOGLE_API_USE_CLIENT_CERTIFICATE"); - return "true".equals(useClientCertificate); + if (useClientCertificate != null && !useClientCertificate.isEmpty()) { + if ("false".equalsIgnoreCase(useClientCertificate)) { + return false; + } + if ("true".equalsIgnoreCase(useClientCertificate)) { + return true; + } + } + + // 2. Check the certificate config file path if provided via env var + String certConfigPath = envProvider.getenv("GOOGLE_API_CERTIFICATE_CONFIG"); + if (certConfigPath != null && !certConfigPath.isEmpty()) { + return validateAndResolveConfig(certConfigPath); + } + + // 3. Fallback to well-known spiffe path + String wellKnownPath = "/var/run/secrets/workload-spiffe-credentials"; + + // Check for atomic bundle containing both cert and key + if (fileExistenceProvider.exists( + java.nio.file.Paths.get(wellKnownPath, "credentialbundle.pem").toString())) { + return true; + } + + // Check for separate certificate and private key files + if (fileExistenceProvider.exists( + java.nio.file.Paths.get(wellKnownPath, "certificates.pem").toString()) + && fileExistenceProvider.exists( + java.nio.file.Paths.get(wellKnownPath, "private_key.pem").toString())) { + return true; + } + + // Default to false if no configuration is found + return false; + } + + private boolean validateAndResolveConfig(String configPath) { + if (!fileExistenceProvider.exists(configPath)) { + throw new IllegalStateException( + "Certificate config is configured but the file does not exist: " + configPath); + } + try { + CertificateConfig config = parseCertificateConfig(configPath); + if (!fileExistenceProvider.exists(config.certPath) + || !fileExistenceProvider.exists(config.keyPath)) { + throw new IllegalStateException( + "Certificate config points to certificate/key files that do not exist on disk: " + + "cert_path=" + + config.certPath + + ", key_path=" + + config.keyPath); + } + return true; + } catch (Exception e) { + throw new IllegalStateException( + "Failed to parse or validate certificate config: " + configPath, e); + } } /** Returns the current mutual TLS endpoint usage policy. */ @@ -82,4 +202,43 @@ public MtlsEndpointUsagePolicy getMtlsEndpointUsagePolicy() { } return MtlsEndpointUsagePolicy.AUTO; } + + /** + * Resolves and returns the path to the mutual TLS client certificate, or null if none should be + * used. + */ + public String getWorkloadCertPath() { + if (!useMtlsClientCertificate()) { + return null; + } + + String certConfigPath = envProvider.getenv("GOOGLE_API_CERTIFICATE_CONFIG"); + if (certConfigPath != null && !certConfigPath.isEmpty()) { + try { + CertificateConfig config = parseCertificateConfig(certConfigPath); + return config.certPath; + } catch (Exception e) { + throw new IllegalStateException("Failed to parse GOOGLE_API_CERTIFICATE_CONFIG", e); + } + } + + String wellKnownPath = "/var/run/secrets/workload-spiffe-credentials"; + + // Check for atomic bundle containing both cert and key + if (fileExistenceProvider.exists( + java.nio.file.Paths.get(wellKnownPath, "credentialbundle.pem").toString())) { + return java.nio.file.Paths.get(wellKnownPath, "credentialbundle.pem").toString(); + } + + // Check for separate certificate and private key files + if (fileExistenceProvider.exists( + java.nio.file.Paths.get(wellKnownPath, "certificates.pem").toString()) + && fileExistenceProvider.exists( + java.nio.file.Paths.get(wellKnownPath, "private_key.pem").toString())) { + return java.nio.file.Paths.get(wellKnownPath, "certificates.pem").toString(); + } + + // Default to null if no well-known configuration is found + return null; + } } diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/WorkloadCertificateUtils.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/WorkloadCertificateUtils.java new file mode 100644 index 000000000000..c85a5b52bcb2 --- /dev/null +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/WorkloadCertificateUtils.java @@ -0,0 +1,68 @@ +/* + * Copyright 2026 Google LLC + * + * Redistribution and use in source and binary forms, with or without + * modification, are permitted provided that the following conditions are + * met: + * + * * Redistributions of source code must retain the above copyright + * notice, this list of conditions and the following disclaimer. + * * Redistributions in binary form must reproduce the above + * copyright notice, this list of conditions and the following disclaimer + * in the documentation and/or other materials provided with the + * distribution. + * * Neither the name of Google LLC nor the names of its + * contributors may be used to endorse or promote products derived from + * this software without specific prior written permission. + * + * THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS + * "AS IS" AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT + * LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR + * A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT + * OWNER OR CONTRIBUTORS BE LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, + * SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT + * LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, + * DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY + * THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT + * (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE + * OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + */ +package com.google.api.gax.rpc.mtls; + +import com.google.api.core.InternalApi; +import java.io.FileInputStream; +import java.security.MessageDigest; +import java.security.cert.CertificateFactory; +import java.security.cert.X509Certificate; +import java.util.logging.Level; +import java.util.logging.Logger; + +/** Internal utility class for managing dynamic workload certificates. */ +@InternalApi +public class WorkloadCertificateUtils { + + private static final Logger LOG = Logger.getLogger(WorkloadCertificateUtils.class.getName()); + + private WorkloadCertificateUtils() {} + + public static String getCertificateFingerprint(String certPath) { + if (certPath == null) { + return ""; + } + try (FileInputStream fis = new FileInputStream(certPath)) { + CertificateFactory cf = CertificateFactory.getInstance("X.509"); + X509Certificate cert = (X509Certificate) cf.generateCertificate(fis); + MessageDigest md = MessageDigest.getInstance("SHA-256"); + byte[] der = cert.getEncoded(); + byte[] digest = md.digest(der); + StringBuilder sb = new StringBuilder(); + for (byte b : digest) { + sb.append(String.format("%02x", b)); + } + return sb.toString(); + } catch (Exception e) { + LOG.log(Level.FINE, "Could not read or parse workload certificate at path " + certPath, e); + return ""; + } + } +} diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/AttemptCallableTest.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/AttemptCallableTest.java index 4b5da578862c..42251516aacb 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/AttemptCallableTest.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/AttemptCallableTest.java @@ -34,11 +34,15 @@ import static org.mockito.Mockito.mock; import static org.mockito.Mockito.when; +import com.google.api.core.ApiFuture; import com.google.api.core.SettableApiFuture; import com.google.api.gax.retrying.RetrySettings; import com.google.api.gax.retrying.RetryingFuture; import com.google.api.gax.retrying.TimedAttemptSettings; import com.google.api.gax.rpc.testing.FakeCallContext; +import com.google.api.gax.rpc.testing.FakeChannel; +import com.google.api.gax.rpc.testing.FakeStatusCode; +import com.google.api.gax.rpc.testing.FakeTransportChannel; import com.google.api.gax.tracing.ApiTracer; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; @@ -133,4 +137,50 @@ void testRpcTimeoutIsNotErased() { assertThat(capturedCallContext.getValue().getTimeoutDuration()).isEqualTo(callerTimeout); } + + @Test + void testUnauthenticatedExceptionReThrowPreservesContext() { + FakeTransportChannel transportChannel = + FakeTransportChannel.create(new FakeChannel()).setShouldRefresh(true); + ApiCallContext callContext = + FakeCallContext.createDefault().withTransportChannel(transportChannel); + + UnauthenticatedException originalEx = + new UnauthenticatedException( + "Expired cert", + new IllegalStateException("Root cause"), + FakeStatusCode.of(StatusCode.Code.UNAUTHENTICATED), + false); + originalEx.setStackTrace( + new StackTraceElement[] {new StackTraceElement("foo", "bar", "Baz.java", 123)}); + originalEx.addSuppressed(new RuntimeException("Suppressed cause")); + + SettableApiFuture failedFuture = SettableApiFuture.create(); + failedFuture.setException(originalEx); + when(mockInnerCallable.futureCall(Mockito.anyString(), Mockito.any())).thenReturn(failedFuture); + + AttemptCallable callable = + new AttemptCallable<>(mockInnerCallable, "fake-request", callContext); + callable.setExternalFuture(mockExternalFuture); + + callable.call(); + + ArgumentCaptor futureCaptor = ArgumentCaptor.forClass(ApiFuture.class); + Mockito.verify(mockExternalFuture, Mockito.times(2)).setAttemptFuture(futureCaptor.capture()); + + Throwable thrown = null; + try { + futureCaptor.getValue().get(); + } catch (Exception e) { + thrown = e.getCause(); + } + + assertThat(thrown).isInstanceOf(UnauthenticatedException.class); + UnauthenticatedException rethrown = (UnauthenticatedException) thrown; + assertThat(rethrown.isRetryable()).isTrue(); + assertThat(rethrown.getCause()).isEqualTo(originalEx); + assertThat(rethrown.getStackTrace()).isEqualTo(originalEx.getStackTrace()); + assertThat(rethrown.getSuppressed().length).isEqualTo(1); + assertThat(rethrown.getSuppressed()[0]).isInstanceOf(RuntimeException.class); + } } diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/EndpointContextTest.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/EndpointContextTest.java index 4c1b5320de4c..c53ca1509430 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/EndpointContextTest.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/EndpointContextTest.java @@ -77,8 +77,10 @@ void mtlsEndpointResolver_switchToMtlsAllowedIsFalse() throws IOException { FakeMtlsProvider.createTestMtlsKeyStore(), "", throwExceptionForGetKeyStore); boolean switchToMtlsEndpointAllowed = false; CertificateBasedAccess certificateBasedAccess = - new CertificateBasedAccess( - name -> name.equals("GOOGLE_API_USE_MTLS_ENDPOINT") ? "auto" : "true"); + org.mockito.Mockito.mock(CertificateBasedAccess.class); + org.mockito.Mockito.when(certificateBasedAccess.getMtlsEndpointUsagePolicy()) + .thenReturn(CertificateBasedAccess.MtlsEndpointUsagePolicy.AUTO); + org.mockito.Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); Truth.assertThat( defaultEndpointContextBuilder.mtlsEndpointResolver( DEFAULT_ENDPOINT, @@ -97,8 +99,10 @@ void mtlsEndpointResolver_switchToMtlsAllowedIsTrue_mtlsUsageAuto() throws IOExc FakeMtlsProvider.createTestMtlsKeyStore(), "", throwExceptionForGetKeyStore); boolean switchToMtlsEndpointAllowed = true; CertificateBasedAccess certificateBasedAccess = - new CertificateBasedAccess( - name -> name.equals("GOOGLE_API_USE_MTLS_ENDPOINT") ? "auto" : "true"); + org.mockito.Mockito.mock(CertificateBasedAccess.class); + org.mockito.Mockito.when(certificateBasedAccess.getMtlsEndpointUsagePolicy()) + .thenReturn(CertificateBasedAccess.MtlsEndpointUsagePolicy.AUTO); + org.mockito.Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); Truth.assertThat( defaultEndpointContextBuilder.mtlsEndpointResolver( DEFAULT_ENDPOINT, @@ -117,8 +121,10 @@ void mtlsEndpointResolver_switchToMtlsAllowedIsTrue_mtlsUsageAlways() throws IOE FakeMtlsProvider.createTestMtlsKeyStore(), "", throwExceptionForGetKeyStore); boolean switchToMtlsEndpointAllowed = true; CertificateBasedAccess certificateBasedAccess = - new CertificateBasedAccess( - name -> name.equals("GOOGLE_API_USE_MTLS_ENDPOINT") ? "always" : "true"); + org.mockito.Mockito.mock(CertificateBasedAccess.class); + org.mockito.Mockito.when(certificateBasedAccess.getMtlsEndpointUsagePolicy()) + .thenReturn(CertificateBasedAccess.MtlsEndpointUsagePolicy.ALWAYS); + org.mockito.Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); Truth.assertThat( defaultEndpointContextBuilder.mtlsEndpointResolver( DEFAULT_ENDPOINT, @@ -137,8 +143,10 @@ void mtlsEndpointResolver_switchToMtlsAllowedIsTrue_mtlsUsageNever() throws IOEx FakeMtlsProvider.createTestMtlsKeyStore(), "", throwExceptionForGetKeyStore); boolean switchToMtlsEndpointAllowed = true; CertificateBasedAccess certificateBasedAccess = - new CertificateBasedAccess( - name -> name.equals("GOOGLE_API_USE_MTLS_ENDPOINT") ? "never" : "true"); + org.mockito.Mockito.mock(CertificateBasedAccess.class); + org.mockito.Mockito.when(certificateBasedAccess.getMtlsEndpointUsagePolicy()) + .thenReturn(CertificateBasedAccess.MtlsEndpointUsagePolicy.NEVER); + org.mockito.Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); Truth.assertThat( defaultEndpointContextBuilder.mtlsEndpointResolver( DEFAULT_ENDPOINT, @@ -156,8 +164,10 @@ void mtlsEndpointResolver_switchToMtlsAllowedIsTrue_useCertificateIsFalse_nullMt MtlsProvider mtlsProvider = new FakeMtlsProvider(null, "", throwExceptionForGetKeyStore); boolean switchToMtlsEndpointAllowed = true; CertificateBasedAccess certificateBasedAccess = - new CertificateBasedAccess( - name -> name.equals("GOOGLE_API_USE_MTLS_ENDPOINT") ? "auto" : "false"); + org.mockito.Mockito.mock(CertificateBasedAccess.class); + org.mockito.Mockito.when(certificateBasedAccess.getMtlsEndpointUsagePolicy()) + .thenReturn(CertificateBasedAccess.MtlsEndpointUsagePolicy.AUTO); + org.mockito.Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(false); Truth.assertThat( defaultEndpointContextBuilder.mtlsEndpointResolver( DEFAULT_ENDPOINT, @@ -174,8 +184,10 @@ void mtlsEndpointResolver_getKeyStore_throwsIOException() throws IOException { MtlsProvider mtlsProvider = new FakeMtlsProvider(null, "", throwExceptionForGetKeyStore); boolean switchToMtlsEndpointAllowed = true; CertificateBasedAccess certificateBasedAccess = - new CertificateBasedAccess( - name -> name.equals("GOOGLE_API_USE_MTLS_ENDPOINT") ? "auto" : "true"); + org.mockito.Mockito.mock(CertificateBasedAccess.class); + org.mockito.Mockito.when(certificateBasedAccess.getMtlsEndpointUsagePolicy()) + .thenReturn(CertificateBasedAccess.MtlsEndpointUsagePolicy.AUTO); + org.mockito.Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); assertThrows( IOException.class, () -> @@ -272,8 +284,10 @@ void endpointContextBuild_mtlsConfigured_GDU() throws IOException { MtlsProvider mtlsProvider = new FakeMtlsProvider(FakeMtlsProvider.createTestMtlsKeyStore(), "", false); CertificateBasedAccess certificateBasedAccess = - new CertificateBasedAccess( - name -> name.equals("GOOGLE_API_USE_MTLS_ENDPOINT") ? "always" : "true"); + org.mockito.Mockito.mock(CertificateBasedAccess.class); + org.mockito.Mockito.when(certificateBasedAccess.getMtlsEndpointUsagePolicy()) + .thenReturn(CertificateBasedAccess.MtlsEndpointUsagePolicy.ALWAYS); + org.mockito.Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); EndpointContext endpointContext = defaultEndpointContextBuilder .setClientSettingsEndpoint(null) @@ -293,8 +307,10 @@ void endpointContextBuild_mtlsConfigured_nonGDU_throwsIllegalArgumentException() MtlsProvider mtlsProvider = new FakeMtlsProvider(FakeMtlsProvider.createTestMtlsKeyStore(), "", false); CertificateBasedAccess certificateBasedAccess = - new CertificateBasedAccess( - name -> name.equals("GOOGLE_API_USE_MTLS_ENDPOINT") ? "always" : "true"); + org.mockito.Mockito.mock(CertificateBasedAccess.class); + org.mockito.Mockito.when(certificateBasedAccess.getMtlsEndpointUsagePolicy()) + .thenReturn(CertificateBasedAccess.MtlsEndpointUsagePolicy.ALWAYS); + org.mockito.Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); EndpointContext.Builder endpointContextBuilder = defaultEndpointContextBuilder .setUniverseDomain("random.com") diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ServerStreamingAttemptCallableTest.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ServerStreamingAttemptCallableTest.java index 4cdd3cc018cc..2a52cb2b2572 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ServerStreamingAttemptCallableTest.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ServerStreamingAttemptCallableTest.java @@ -29,6 +29,7 @@ */ package com.google.api.gax.rpc; +import static org.junit.jupiter.api.Assertions.assertThrows; import static org.mockito.Mockito.mock; import com.google.api.core.AbstractApiFuture; @@ -248,6 +249,46 @@ void testInitialRetry() { Truth.assertThat(call.getRequest()).isEqualTo("request > 0"); } + @Test + @SuppressWarnings("ConstantConditions") + void testUnauthenticatedRefresh() { + TransportChannel transportChannel = Mockito.mock(TransportChannel.class); + Mockito.when(transportChannel.shouldRefresh()).thenReturn(true); + + ApiCallContext context = Mockito.mock(ApiCallContext.class); + Mockito.when(context.getTransportChannel()).thenReturn(transportChannel); + Mockito.when(context.getTracer()).thenReturn(BaseApiTracer.getInstance()); + Mockito.when(context.getTimeoutDuration()).thenReturn(java.time.Duration.ofHours(5)); + + resumptionStrategy = new MyStreamResumptionStrategy(); + ServerStreamingAttemptCallable callable = createCallable(context); + callable.start(); + + MockServerStreamingCall call = innerCallable.popLastCall(); + + // Send initial error + UnauthenticatedException initialError = + new UnauthenticatedException( + "test", + null, + com.google.api.gax.rpc.testing.FakeStatusCode.of(Code.UNAUTHENTICATED), + false); + call.getController().getObserver().onError(initialError); + + // Should notify the outer future + ExecutionException ee = + assertThrows( + ExecutionException.class, + () -> fakeRetryingFuture.getAttemptResult().get(1, TimeUnit.SECONDS)); + Throwable outerError = ee.getCause(); + Mockito.verify(transportChannel).refresh(); + Truth.assertThat(outerError).isInstanceOf(ServerStreamingAttemptException.class); + Truth.assertThat(((ServerStreamingAttemptException) outerError).hasSeenResponses()).isFalse(); + Truth.assertThat(((ServerStreamingAttemptException) outerError).canResume()).isTrue(); + Truth.assertThat(outerError.getCause()).isInstanceOf(UnauthenticatedException.class); + Truth.assertThat(((UnauthenticatedException) outerError.getCause()).isRetryable()).isTrue(); + } + @Test @SuppressWarnings("ConstantConditions") void testMidRetry() { diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/StreamingCallableTest.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/StreamingCallableTest.java index 6f9584826893..c46522114318 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/StreamingCallableTest.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/StreamingCallableTest.java @@ -129,7 +129,7 @@ void testClientStreamingCall() { ClientStreamingCallable callable = stashCallable.withDefaultCallContext(defaultCallContext); callable.clientStreamingCall(observer); - assertSame(observer, stashCallable.getActualObserver()); + org.junit.jupiter.api.Assertions.assertNotNull(stashCallable.getActualObserver()); assertSame(defaultCallContext, stashCallable.getContext()); } @@ -158,7 +158,7 @@ void testClientStreamingCallWithContext() { ClientStreamingCallable callable = stashCallable.withDefaultCallContext(FakeCallContext.createDefault()); callable.clientStreamingCall(observer, context); - assertSame(observer, stashCallable.getActualObserver()); + org.junit.jupiter.api.Assertions.assertNotNull(stashCallable.getActualObserver()); FakeCallContext actualContext = (FakeCallContext) stashCallable.getContext(); assertSame(channel, actualContext.getChannel()); assertSame(credentials, actualContext.getCredentials()); diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/AbstractMtlsTransportChannelTest.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/AbstractMtlsTransportChannelTest.java index bea4674b765b..5ed563254965 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/AbstractMtlsTransportChannelTest.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/AbstractMtlsTransportChannelTest.java @@ -64,8 +64,10 @@ void testNotUseClientCertificate() throws IOException, GeneralSecurityException @Test void testUseClientCertificate() throws IOException, GeneralSecurityException { CertificateBasedAccess certificateBasedAccess = - new CertificateBasedAccess( - name -> name.equals("GOOGLE_API_USE_MTLS_ENDPOINT") ? "auto" : "true"); + org.mockito.Mockito.mock(CertificateBasedAccess.class); + org.mockito.Mockito.when(certificateBasedAccess.getMtlsEndpointUsagePolicy()) + .thenReturn(CertificateBasedAccess.MtlsEndpointUsagePolicy.AUTO); + org.mockito.Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); MtlsProvider provider = new FakeMtlsProvider(FakeMtlsProvider.createTestMtlsKeyStore(), "", false); assertNotNull(getMtlsObjectFromTransportChannel(provider, certificateBasedAccess)); @@ -74,8 +76,10 @@ void testUseClientCertificate() throws IOException, GeneralSecurityException { @Test void testNoClientCertificate() throws IOException, GeneralSecurityException { CertificateBasedAccess certificateBasedAccess = - new CertificateBasedAccess( - name -> name.equals("GOOGLE_API_USE_MTLS_ENDPOINT") ? "auto" : "true"); + org.mockito.Mockito.mock(CertificateBasedAccess.class); + org.mockito.Mockito.when(certificateBasedAccess.getMtlsEndpointUsagePolicy()) + .thenReturn(CertificateBasedAccess.MtlsEndpointUsagePolicy.AUTO); + org.mockito.Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); MtlsProvider provider = new FakeMtlsProvider(null, "", false); assertNull(getMtlsObjectFromTransportChannel(provider, certificateBasedAccess)); } @@ -84,8 +88,10 @@ void testNoClientCertificate() throws IOException, GeneralSecurityException { void testGetKeyStoreThrows() throws GeneralSecurityException { // Test the case where provider.getKeyStore() throws. CertificateBasedAccess certificateBasedAccess = - new CertificateBasedAccess( - name -> name.equals("GOOGLE_API_USE_MTLS_ENDPOINT") ? "auto" : "true"); + org.mockito.Mockito.mock(CertificateBasedAccess.class); + org.mockito.Mockito.when(certificateBasedAccess.getMtlsEndpointUsagePolicy()) + .thenReturn(CertificateBasedAccess.MtlsEndpointUsagePolicy.AUTO); + org.mockito.Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); MtlsProvider provider = new FakeMtlsProvider(null, "", true); IOException actual = assertThrows( diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/CertificateBasedAccessTest.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/CertificateBasedAccessTest.java index e328e0af4799..c92f229dfcbb 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/CertificateBasedAccessTest.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/CertificateBasedAccessTest.java @@ -32,52 +32,234 @@ import static org.junit.jupiter.api.Assertions.assertEquals; import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertThrows; import static org.junit.jupiter.api.Assertions.assertTrue; +import java.io.IOException; +import java.util.HashMap; +import java.util.Map; import org.junit.jupiter.api.Test; class CertificateBasedAccessTest { + private static class TestEnv { + private final Map env = new HashMap<>(); + + void set(String key, String val) { + env.put(key, val); + } + + String get(String name) { + return env.get(name); + } + } + + private static class TestFileSystem { + private final Map exists = new HashMap<>(); + private final Map content = new HashMap<>(); + + void setExists(String path, boolean val) { + exists.put(java.nio.file.Paths.get(path).toString(), val); + } + + void setContent(String path, String val) { + String normalizedPath = java.nio.file.Paths.get(path).toString(); + content.put(normalizedPath, val); + exists.put(normalizedPath, true); + } + } + + private CertificateBasedAccess createCba(TestEnv env, TestFileSystem fs) { + return new CertificateBasedAccess( + env::get, + path -> fs.exists.getOrDefault(java.nio.file.Paths.get(path).toString(), false), + path -> { + String normalized = java.nio.file.Paths.get(path).toString(); + if (!fs.content.containsKey(normalized)) { + throw new IOException("File not found: " + path); + } + return fs.content.get(normalized); + }); + } + @Test void testUseMtlsEndpointAlways() { - CertificateBasedAccess cba = - new CertificateBasedAccess( - name -> name.equals("GOOGLE_API_USE_MTLS_ENDPOINT") ? "always" : "false"); + TestEnv env = new TestEnv(); + env.set("GOOGLE_API_USE_MTLS_ENDPOINT", "always"); + CertificateBasedAccess cba = createCba(env, new TestFileSystem()); assertEquals( CertificateBasedAccess.MtlsEndpointUsagePolicy.ALWAYS, cba.getMtlsEndpointUsagePolicy()); } @Test void testUseMtlsEndpointAuto() { - CertificateBasedAccess cba = - new CertificateBasedAccess( - name -> name.equals("GOOGLE_API_USE_MTLS_ENDPOINT") ? "auto" : "false"); + TestEnv env = new TestEnv(); + env.set("GOOGLE_API_USE_MTLS_ENDPOINT", "auto"); + CertificateBasedAccess cba = createCba(env, new TestFileSystem()); assertEquals( CertificateBasedAccess.MtlsEndpointUsagePolicy.AUTO, cba.getMtlsEndpointUsagePolicy()); } @Test void testUseMtlsEndpointNever() { - CertificateBasedAccess cba = - new CertificateBasedAccess( - name -> name.equals("GOOGLE_API_USE_MTLS_ENDPOINT") ? "never" : "false"); + TestEnv env = new TestEnv(); + env.set("GOOGLE_API_USE_MTLS_ENDPOINT", "never"); + CertificateBasedAccess cba = createCba(env, new TestFileSystem()); assertEquals( CertificateBasedAccess.MtlsEndpointUsagePolicy.NEVER, cba.getMtlsEndpointUsagePolicy()); } @Test - void testUseMtlsClientCertificateTrue() { - CertificateBasedAccess cba = - new CertificateBasedAccess( - name -> name.equals("GOOGLE_API_USE_CLIENT_CERTIFICATE") ? "true" : "auto"); + void testUseMtlsClientCertificateExplicitTrueNoCredentials() { + TestEnv env = new TestEnv(); + env.set("GOOGLE_API_USE_CLIENT_CERTIFICATE", "true"); + CertificateBasedAccess cba = createCba(env, new TestFileSystem()); + // Explicit 'true' overrides credential presence checks + assertTrue(cba.useMtlsClientCertificate()); + } + + @Test + void testUseMtlsClientCertificateExplicitTrueWithSpiffeBundle() { + TestEnv env = new TestEnv(); + env.set("GOOGLE_API_USE_CLIENT_CERTIFICATE", "true"); + + TestFileSystem fs = new TestFileSystem(); + fs.setExists("/var/run/secrets/workload-spiffe-credentials/credentialbundle.pem", true); + + CertificateBasedAccess cba = createCba(env, fs); assertTrue(cba.useMtlsClientCertificate()); } @Test - void testUseMtlsClientCertificateFalse() { - CertificateBasedAccess cba = - new CertificateBasedAccess( - name -> name.equals("GOOGLE_API_USE_CLIENT_CERTIFICATE") ? "false" : "auto"); + void testUseMtlsClientCertificateExplicitFalseWithSpiffeBundle() { + TestEnv env = new TestEnv(); + env.set("GOOGLE_API_USE_CLIENT_CERTIFICATE", "false"); + + // Even if spiffe files are present, explicit false must override and disable mtls + TestFileSystem fs = new TestFileSystem(); + fs.setExists("/var/run/secrets/workload-spiffe-credentials/credentialbundle.pem", true); + + CertificateBasedAccess cba = createCba(env, fs); + assertFalse(cba.useMtlsClientCertificate()); + } + + @Test + void testUseMtlsClientCertificateUnsetNoFiles() { + TestEnv env = new TestEnv(); + CertificateBasedAccess cba = createCba(env, new TestFileSystem()); assertFalse(cba.useMtlsClientCertificate()); } + + @Test + void testUseMtlsClientCertificateUnsetSpiffeBundleExists() { + TestEnv env = new TestEnv(); + TestFileSystem fs = new TestFileSystem(); + fs.setExists("/var/run/secrets/workload-spiffe-credentials/credentialbundle.pem", true); + CertificateBasedAccess cba = createCba(env, fs); + assertTrue(cba.useMtlsClientCertificate()); + } + + @Test + void testUseMtlsClientCertificateUnsetSpiffeSeparateFilesExist() { + TestEnv env = new TestEnv(); + TestFileSystem fs = new TestFileSystem(); + fs.setExists("/var/run/secrets/workload-spiffe-credentials/certificates.pem", true); + fs.setExists("/var/run/secrets/workload-spiffe-credentials/private_key.pem", true); + CertificateBasedAccess cba = createCba(env, fs); + assertTrue(cba.useMtlsClientCertificate()); + } + + @Test + void testUseMtlsClientCertificateConfigValid() { + TestEnv env = new TestEnv(); + env.set("GOOGLE_API_CERTIFICATE_CONFIG", "/path/to/config.json"); + + TestFileSystem fs = new TestFileSystem(); + fs.setContent( + "/path/to/config.json", + "{\n \"cert_path\": \"/my/cert.pem\",\n \"key_path\": \"/my/key.pem\"\n}"); + fs.setExists("/my/cert.pem", true); + fs.setExists("/my/key.pem", true); + + CertificateBasedAccess cba = createCba(env, fs); + assertTrue(cba.useMtlsClientCertificate()); + } + + @Test + void testUseMtlsClientCertificateConfigMissingFile() { + TestEnv env = new TestEnv(); + env.set("GOOGLE_API_CERTIFICATE_CONFIG", "/path/to/config.json"); + + CertificateBasedAccess cba = createCba(env, new TestFileSystem()); + + IllegalStateException ex = + assertThrows(IllegalStateException.class, cba::useMtlsClientCertificate); + assertTrue(ex.getMessage().contains("configured but the file does not exist")); + } + + @Test + void testUseMtlsClientCertificateEnvTrueOverride() { + TestEnv env = new TestEnv(); + env.set("GOOGLE_API_USE_CLIENT_CERTIFICATE", "true"); + + CertificateBasedAccess cba = createCba(env, new TestFileSystem()); + + assertTrue(cba.useMtlsClientCertificate()); + } + + @Test + void testUseMtlsClientCertificateConfigMalformedJson() { + TestEnv env = new TestEnv(); + env.set("GOOGLE_API_CERTIFICATE_CONFIG", "/path/to/config.json"); + + TestFileSystem fs = new TestFileSystem(); + fs.setContent("/path/to/config.json", "{\n \"broken_path\": \"/my/cert.pem\"\n}"); + + CertificateBasedAccess cba = createCba(env, fs); + + IllegalStateException ex = + assertThrows(IllegalStateException.class, cba::useMtlsClientCertificate); + assertTrue(ex.getMessage().contains("Failed to parse or validate certificate config")); + } + + @Test + void testUseMtlsClientCertificateConfigMissingCertFiles() { + TestEnv env = new TestEnv(); + env.set("GOOGLE_API_CERTIFICATE_CONFIG", "/path/to/config.json"); + + TestFileSystem fs = new TestFileSystem(); + fs.setContent( + "/path/to/config.json", + "{\n \"cert_path\": \"/my/cert.pem\",\n \"key_path\": \"/my/key.pem\"\n}"); + // my/cert.pem and key.pem DO NOT exist + + CertificateBasedAccess cba = createCba(env, fs); + + IllegalStateException ex = + assertThrows(IllegalStateException.class, cba::useMtlsClientCertificate); + assertTrue( + ex.getCause() + .getMessage() + .contains("points to certificate/key files that do not exist on disk")); + } + + @Test + void testUseMtlsClientCertificateConfigWindowsPaths() { + TestEnv env = new TestEnv(); + env.set("GOOGLE_API_CERTIFICATE_CONFIG", "C:\\config.json"); + + TestFileSystem fs = new TestFileSystem(); + // In JSON, backslashes are escaped + fs.setContent( + "C:\\config.json", + "{\n" + + " \"cert_path\": \"C:\\\\my\\\\cert.pem\",\n" + + " \"key_path\": \"C:\\\\my\\\\key.pem\"\n" + + "}"); + fs.setExists("C:\\my\\cert.pem", true); + fs.setExists("C:\\my\\key.pem", true); + + CertificateBasedAccess cba = createCba(env, fs); + assertTrue(cba.useMtlsClientCertificate()); + } } diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/testing/FakeCallContext.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/testing/FakeCallContext.java index 1cdefe435d55..e41dc041220d 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/testing/FakeCallContext.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/testing/FakeCallContext.java @@ -241,6 +241,11 @@ public FakeChannel getChannel() { return channel; } + @Override + public TransportChannel getTransportChannel() { + return channel != null ? FakeTransportChannel.create(channel) : null; + } + @Override public java.time.Duration getTimeoutDuration() { return timeout; diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/testing/FakeChannel.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/testing/FakeChannel.java index b725da363b3b..6d676db40032 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/testing/FakeChannel.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/testing/FakeChannel.java @@ -32,4 +32,24 @@ import com.google.api.core.InternalApi; @InternalApi("for testing") -public class FakeChannel {} +public class FakeChannel { + private volatile boolean shouldRefresh = false; + private volatile int refreshCount = 0; + + public FakeChannel setShouldRefresh(boolean shouldRefresh) { + this.shouldRefresh = shouldRefresh; + return this; + } + + public boolean shouldRefresh() { + return shouldRefresh; + } + + public void refresh() { + refreshCount++; + } + + public int getRefreshCount() { + return refreshCount; + } +} diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/testing/FakeTransportChannel.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/testing/FakeTransportChannel.java index 0d4abac8f1c6..bbde9feb4807 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/testing/FakeTransportChannel.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/testing/FakeTransportChannel.java @@ -41,6 +41,33 @@ public class FakeTransportChannel implements TransportChannel { private volatile boolean isShutdown = false; private volatile Map headers; private volatile Executor executor; + private volatile boolean shouldRefresh = false; + private volatile int refreshCount = 0; + + public FakeTransportChannel setShouldRefresh(boolean shouldRefresh) { + if (channel != null) { + channel.setShouldRefresh(shouldRefresh); + } + this.shouldRefresh = shouldRefresh; + return this; + } + + @Override + public boolean shouldRefresh() { + return channel != null ? channel.shouldRefresh() : shouldRefresh; + } + + @Override + public void refresh() { + if (channel != null) { + channel.refresh(); + } + refreshCount++; + } + + public int getRefreshCount() { + return channel != null ? channel.getRefreshCount() : refreshCount; + } private FakeTransportChannel(FakeChannel channel) { this.channel = channel; From 49772f0f5f461972fcdcfd46604a58fea5af7f7f Mon Sep 17 00:00:00 2001 From: Matt Castelaz Date: Wed, 5 Aug 2026 02:35:10 +0000 Subject: [PATCH 02/29] fix(gax): address mTLS cert config parsing and channel refresh lifecycle edge cases --- .../gax/httpjson/RefreshingHttpJsonChannel.java | 16 +++++++++++++++- .../api/gax/rpc/mtls/CertificateBasedAccess.java | 8 +++++--- .../gax/rpc/mtls/CertificateBasedAccessTest.java | 6 ++---- 3 files changed, 22 insertions(+), 8 deletions(-) diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java index 0b4c3aa885fd..5d9f69a88bbf 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java @@ -35,6 +35,7 @@ import com.google.api.gax.httpjson.ForwardingHttpJsonClientCallListener.SimpleForwardingHttpJsonClientCallListener; import com.google.api.gax.rpc.mtls.WorkloadCertificateUtils; import com.google.common.annotations.VisibleForTesting; +import java.util.concurrent.CancellationException; import java.util.concurrent.TimeUnit; import java.util.concurrent.atomic.AtomicBoolean; import java.util.concurrent.atomic.AtomicInteger; @@ -42,6 +43,7 @@ import java.util.function.Supplier; import java.util.logging.Level; import java.util.logging.Logger; +import org.jspecify.annotations.Nullable; /** * An implementation of {@link ManagedHttpJsonChannel} that supports dynamic mTLS certificate @@ -240,7 +242,6 @@ public void shutdownNow() { synchronized (refreshLock) { isShuttingDown = true; for (ChannelEntry entry : allEntries) { - entry.requestShutdown(); entry.channel.shutdownNow(); } } @@ -329,6 +330,7 @@ private void shutdown() { private static class ReleasingHttpJsonClientCall extends SimpleForwardingHttpJsonClientCall { + private @Nullable CancellationException cancellationException; private final ChannelEntry entry; private final AtomicBoolean wasClosed = new AtomicBoolean(false); private final AtomicBoolean wasReleased = new AtomicBoolean(false); @@ -340,6 +342,12 @@ private static class ReleasingHttpJsonClientCall @Override public void start(Listener responseListener, HttpJsonMetadata requestHeaders) { + if (cancellationException != null) { + if (wasReleased.compareAndSet(false, true)) { + entry.release(); + } + throw new IllegalStateException("Call is already cancelled", cancellationException); + } try { super.start( new SimpleForwardingHttpJsonClientCallListener(responseListener) { @@ -365,5 +373,11 @@ public void onClose(int statusCode, HttpJsonMetadata trailers) { throw e; } } + + @Override + public void cancel(@Nullable String message, @Nullable Throwable cause) { + this.cancellationException = new CancellationException(message); + super.cancel(message, cause); + } } } diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/CertificateBasedAccess.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/CertificateBasedAccess.java index de8bacd7c02b..bc4a59a2c35a 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/CertificateBasedAccess.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/CertificateBasedAccess.java @@ -111,8 +111,7 @@ private CertificateConfig parseCertificateConfig(String configPath) throws IOExc String keyPath = extractJsonValue(content, "key_path"); if (certPath == null || keyPath == null) { - throw new IllegalStateException( - "Malformed certificate config JSON. Must contain 'cert_path' and 'key_path'."); + return null; } return new CertificateConfig(certPath, keyPath); @@ -176,6 +175,9 @@ private boolean validateAndResolveConfig(String configPath) { } try { CertificateConfig config = parseCertificateConfig(configPath); + if (config == null) { + return false; + } if (!fileExistenceProvider.exists(config.certPath) || !fileExistenceProvider.exists(config.keyPath)) { throw new IllegalStateException( @@ -216,7 +218,7 @@ public String getWorkloadCertPath() { if (certConfigPath != null && !certConfigPath.isEmpty()) { try { CertificateConfig config = parseCertificateConfig(certConfigPath); - return config.certPath; + return config != null ? config.certPath : null; } catch (Exception e) { throw new IllegalStateException("Failed to parse GOOGLE_API_CERTIFICATE_CONFIG", e); } diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/CertificateBasedAccessTest.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/CertificateBasedAccessTest.java index c92f229dfcbb..e1219d631812 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/CertificateBasedAccessTest.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/CertificateBasedAccessTest.java @@ -208,7 +208,7 @@ void testUseMtlsClientCertificateEnvTrueOverride() { } @Test - void testUseMtlsClientCertificateConfigMalformedJson() { + void testUseMtlsClientCertificateConfigNonWorkloadJson() { TestEnv env = new TestEnv(); env.set("GOOGLE_API_CERTIFICATE_CONFIG", "/path/to/config.json"); @@ -217,9 +217,7 @@ void testUseMtlsClientCertificateConfigMalformedJson() { CertificateBasedAccess cba = createCba(env, fs); - IllegalStateException ex = - assertThrows(IllegalStateException.class, cba::useMtlsClientCertificate); - assertTrue(ex.getMessage().contains("Failed to parse or validate certificate config")); + assertFalse(cba.useMtlsClientCertificate()); } @Test From a2210c69a11403b3d8163ef8b43194f0decdc564 Mon Sep 17 00:00:00 2001 From: Matt Castelaz Date: Wed, 5 Aug 2026 03:13:49 +0000 Subject: [PATCH 03/29] fix(gax): address PR 13995 AI review findings and Javadoc doclint errors Addresses AI code review findings from https://paste.googleplex.com/6563525517508608: - GrpcCallContext: Prevent transportChannel stale inheritance in merge() and withChannel() - RefreshingHttpJsonChannel: Set shutdownRequested and shutdownInitiated in shutdownNow() so newCall() throws IllegalStateException - AttemptCallable / StreamingCallables: Pass getCause() when rethrowing retryable UnauthenticatedException to prevent double-wrapping - CertificateBasedAccess: Enforce fail-closed security boundary when certificate config is malformed or missing required keys, and fix JSON unescaping order - ChannelPool: Update ReleasingClientCall Javadoc contract - Unit tests: Add cache invalidation test helpers to eliminate Thread.sleep() delays and add comprehensive tests for all addressed edge cases --- .../com/google/api/gax/grpc/ChannelPool.java | 12 +++- .../google/api/gax/grpc/GrpcCallContext.java | 5 +- .../google/api/gax/grpc/ChannelPoolTest.java | 33 ++++++++++- .../api/gax/grpc/GrpcCallContextTest.java | 29 ++++++++++ .../InstantiatingHttpJsonChannelProvider.java | 5 +- .../gax/httpjson/ManagedHttpJsonChannel.java | 7 ++- .../httpjson/RefreshingHttpJsonChannel.java | 7 +++ .../RefreshingHttpJsonChannelTest.java | 56 +++++++++++++++++-- .../api/gax/retrying/RetrySettings.java | 6 +- .../google/api/gax/rpc/AttemptCallable.java | 2 +- .../api/gax/rpc/BidiStreamingCallable.java | 7 ++- .../api/gax/rpc/ClientStreamingCallable.java | 7 ++- .../rpc/ServerStreamingAttemptCallable.java | 7 ++- .../gax/rpc/mtls/CertificateBasedAccess.java | 12 +++- .../api/gax/rpc/AttemptCallableTest.java | 2 +- .../rpc/mtls/CertificateBasedAccessTest.java | 35 +++++++++++- 16 files changed, 206 insertions(+), 26 deletions(-) diff --git a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java index 5c583654f43b..90a9bbe435f0 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java @@ -485,6 +485,11 @@ private String getOrUpdateDiskFingerprint(String certPath) { } } + @VisibleForTesting + void invalidateDiskFingerprintCache() { + this.lastDiskCheck = null; + } + boolean shouldRefresh() { if (workloadCertPath == null) { return false; @@ -513,6 +518,7 @@ void refresh() { // replaces the list) synchronized (entryWriteLock) { if (workloadCertPath == null) { + refreshAll(); return; } String currentDiskFingerprint = getOrUpdateDiskFingerprint(workloadCertPath); @@ -721,9 +727,9 @@ public ClientCall newCall( /** * ClientCall wrapper that makes sure to decrement the outstanding RPC count on completion. * - *

Contract: Exactly one call to {@link #start(Listener, Metadata)} or explicit release via - * {@link #cancel(String, Throwable)} is required to balance reference counts. Early cancellation - * before {@code start()} safely decrements the reference count via atomic compare-and-set. + *

Contract: Exactly one call to {@link #start(Listener, Metadata)} is required to balance + * reference counts. Early cancellation before {@code start()} is recorded and safely decrements + * the reference count when {@code start()} is subsequently invoked. */ static class ReleasingClientCall extends SimpleForwardingClientCall { private @Nullable CancellationException cancellationException; diff --git a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/GrpcCallContext.java b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/GrpcCallContext.java index e5ead54e9181..428531848b12 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/GrpcCallContext.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/GrpcCallContext.java @@ -563,7 +563,8 @@ public ApiCallContext merge(ApiCallContext inputCallContext) { } TransportChannel newTransportChannel = grpcCallContext.transportChannel; - if (newTransportChannel == null) { + if (newTransportChannel == null + && (grpcCallContext.channel == null || grpcCallContext.channel.equals(channel))) { newTransportChannel = transportChannel; } @@ -662,7 +663,7 @@ public GrpcCallContext withChannel(@Nullable Channel newChannel) { retryableCodes, endpointContext, isDirectPath, - transportChannel); + (newChannel == null || newChannel.equals(channel)) ? transportChannel : null); } /** Returns a new instance with the call options set to the given call options. */ diff --git a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java index 6c2adc1f1463..1157b8020cf9 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java @@ -473,7 +473,7 @@ void channelReactiveMTlsRefreshShouldConditionallySwapChannels() .newCall(Mockito.>any(), Mockito.any(CallOptions.class)); // The ChannelPool caches fingerprints for 1000ms, wait for it to expire - Thread.sleep(1100); + pool.invalidateDiskFingerprintCache(); java.nio.file.Path rootCert = java.nio.file.Paths.get("src", "test", "resources", "root_cert.pem"); @@ -529,6 +529,37 @@ void channelRefreshShouldSwapChannels() throws IOException { .newCall(Mockito.>any(), Mockito.any(CallOptions.class)); } + @Test + void testRefreshWithNullWorkloadCertPathSwapsChannel() throws IOException { + ScheduledExecutorService executor = + Mockito.mock(ScheduledExecutorService.class, Mockito.withSettings().withoutAnnotations()); + FixedExecutorProvider provider = FixedExecutorProvider.create(executor); + ManagedChannel underlyingChannel1 = Mockito.mock(ManagedChannel.class); + ManagedChannel underlyingChannel2 = Mockito.mock(ManagedChannel.class); + FakeChannelFactory channelFactory = + new FakeChannelFactory(ImmutableList.of(underlyingChannel1, underlyingChannel2)); + pool = + new ChannelPool( + ChannelPoolSettings.staticallySized(1).toBuilder() + .setPreemptiveRefreshEnabled(true) + .build(), + channelFactory, + provider, + null); + Mockito.reset(underlyingChannel1); + + pool.newCall(FakeMethodDescriptor.create(), CallOptions.DEFAULT); + Mockito.verify(underlyingChannel1, Mockito.only()) + .newCall(Mockito.>any(), Mockito.any(CallOptions.class)); + + // Calling refresh() when workloadCertPath is null should fall back to refreshAll() + pool.refresh(); + + pool.newCall(FakeMethodDescriptor.create(), CallOptions.DEFAULT); + Mockito.verify(underlyingChannel2, Mockito.only()) + .newCall(Mockito.>any(), Mockito.any(CallOptions.class)); + } + @Test void channelCountShouldNotChangeWhenOutstandingRpcsAreWithinLimits() throws Exception { ScheduledExecutorService executor = diff --git a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/GrpcCallContextTest.java b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/GrpcCallContextTest.java index cacdfed88980..59d5bbf568be 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/GrpcCallContextTest.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/GrpcCallContextTest.java @@ -515,4 +515,33 @@ public void testEqualsAndHashCode() { org.junit.jupiter.api.Assertions.assertNotEquals(context1, context3); } + + @Test + public void testMergeWithCustomChannelClearsTransportChannel() { + ManagedChannel defaultChannel = org.mockito.Mockito.mock(ManagedChannel.class); + ManagedChannel customChannel = org.mockito.Mockito.mock(ManagedChannel.class); + GrpcTransportChannel transportChannel = GrpcTransportChannel.create(defaultChannel); + + GrpcCallContext baseContext = + GrpcCallContext.createDefault().withTransportChannel(transportChannel); + GrpcCallContext overrideContext = GrpcCallContext.of(customChannel, CallOptions.DEFAULT); + + GrpcCallContext mergedContext = (GrpcCallContext) baseContext.merge(overrideContext); + assertEquals(customChannel, mergedContext.getChannel()); + assertNull(mergedContext.getTransportChannel()); + } + + @Test + public void testWithChannelWithCustomChannelClearsTransportChannel() { + ManagedChannel defaultChannel = org.mockito.Mockito.mock(ManagedChannel.class); + ManagedChannel customChannel = org.mockito.Mockito.mock(ManagedChannel.class); + GrpcTransportChannel transportChannel = GrpcTransportChannel.create(defaultChannel); + + GrpcCallContext baseContext = + GrpcCallContext.createDefault().withTransportChannel(transportChannel); + GrpcCallContext updatedContext = baseContext.withChannel(customChannel); + + assertEquals(customChannel, updatedContext.getChannel()); + assertNull(updatedContext.getTransportChannel()); + } } diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProvider.java b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProvider.java index 13aefaef4b0c..90ce27c2879d 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProvider.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProvider.java @@ -199,7 +199,10 @@ public TransportChannelProvider withCredentials(Credentials credentials) { if (certificateBasedAccess.useMtlsClientCertificate()) { KeyStore mtlsKeyStore = mtlsProvider.getKeyStore(); if (mtlsKeyStore != null) { - return new NetHttpTransport.Builder().trustCertificates(null, mtlsKeyStore, "").build(); + NetHttpTransport.Builder builder = new NetHttpTransport.Builder(); + builder.trustCertificates(null, mtlsKeyStore, ""); + HttpJsonConscryptUtils.configureConscryptSecurityProvider(builder); + return builder.build(); } } return null; diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/ManagedHttpJsonChannel.java b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/ManagedHttpJsonChannel.java index 99ece14670f9..f83f09bac486 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/ManagedHttpJsonChannel.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/ManagedHttpJsonChannel.java @@ -78,7 +78,12 @@ private ManagedHttpJsonChannel( this.executor = executor; this.usingDefaultExecutor = usingDefaultExecutor; this.endpoint = endpoint; - this.httpTransport = httpTransport == null ? new NetHttpTransport() : httpTransport; + this.httpTransport = + httpTransport == null + ? HttpJsonConscryptUtils.configureConscryptSecurityProvider( + new NetHttpTransport.Builder()) + .build() + : httpTransport; this.usingDefaultTransport = usingDefaultTransport || httpTransport == null; this.deadlineScheduledExecutorService = Executors.newSingleThreadScheduledExecutor(); } diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java index 5d9f69a88bbf..5ffa35626d1f 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java @@ -242,11 +242,18 @@ public void shutdownNow() { synchronized (refreshLock) { isShuttingDown = true; for (ChannelEntry entry : allEntries) { + entry.shutdownRequested.set(true); + entry.shutdownInitiated.set(true); entry.channel.shutdownNow(); } } } + @VisibleForTesting + void invalidateDiskFingerprintCache() { + this.lastDiskCheck = null; + } + @Override public boolean awaitTermination(long duration, TimeUnit unit) throws InterruptedException { long endNanos = System.nanoTime() + unit.toNanos(duration); diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java index 69e2b1c69113..425dd340484b 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java @@ -177,7 +177,7 @@ void testShouldRefreshNullCertPath() { void testShouldRefreshFalseWhenUnchanged() throws InterruptedException { RefreshingHttpJsonChannel channel = createTestChannel(); - Thread.sleep(1001); // Invalidate 1-second cache + channel.invalidateDiskFingerprintCache(); // Invalidate 1-second cache assertFalse(channel.shouldRefresh()); } @@ -185,7 +185,7 @@ void testShouldRefreshFalseWhenUnchanged() throws InterruptedException { void testShouldRefreshTrueWhenChanged() throws InterruptedException { RefreshingHttpJsonChannel channel = createTestChannel(); - Thread.sleep(1001); // Invalidate 1-second cache + channel.invalidateDiskFingerprintCache(); // Invalidate 1-second cache // Simulate disk fingerprint changing testFingerprint = "fingerprint2"; @@ -199,7 +199,7 @@ void testRefreshSwapsChannel() throws InterruptedException { FakeManagedHttpJsonChannel firstChannel = lastCreatedChannel; assertEquals(1, channelFactoryCount.get()); - Thread.sleep(1001); // Invalidate 1-second cache + channel.invalidateDiskFingerprintCache(); // Invalidate 1-second cache // Change fingerprint testFingerprint = "fingerprint2"; @@ -227,7 +227,7 @@ void testRefreshKeepsInFlightChannelsAlive() throws InterruptedException { HttpJsonClientCall activeCall = channel.newCall(null, null); - Thread.sleep(1001); // Invalidate 1-second cache + channel.invalidateDiskFingerprintCache(); // Invalidate 1-second cache // Change fingerprint & refresh testFingerprint = "fingerprint2"; @@ -262,7 +262,7 @@ void testRefreshDoesNotSpawnChannelWhenShutdown() throws InterruptedException { channel.shutdown(); firstChannel.shutdown(); - Thread.sleep(1001); // Invalidate 1-second cache + channel.invalidateDiskFingerprintCache(); // Invalidate 1-second cache // Change fingerprint testFingerprint = "fingerprint2"; @@ -280,7 +280,7 @@ void testRefreshFactoryExceptionDoesNotWedgeFingerprint() throws InterruptedExce assertEquals(1, channelFactoryCount.get()); shouldThrowOnFactory = true; - Thread.sleep(1001); // Invalidate 1-second cache + channel.invalidateDiskFingerprintCache(); // Invalidate 1-second cache testFingerprint = "fingerprint2"; assertThrows(RuntimeException.class, channel::refresh); @@ -324,4 +324,48 @@ void testChannelDelegationMethods() { assertEquals(firstChannel.getHttpTransport(), channel.getHttpTransport()); assertEquals(firstChannel.getExecutor(), channel.getExecutor()); } + + @Test + void testNewCallAfterShutdownNowThrowsIllegalStateException() { + RefreshingHttpJsonChannel channel = createTestChannel(); + channel.shutdownNow(); + + assertThrows( + IllegalStateException.class, + () -> channel.newCall(null, null), + "Channel has been shut down"); + } + + @Test + void testConcurrentNewCallDuringRefresh() throws InterruptedException { + RefreshingHttpJsonChannel channel = createTestChannel(); + int threadCount = 10; + java.util.concurrent.ExecutorService executorService = + java.util.concurrent.Executors.newFixedThreadPool(threadCount); + java.util.concurrent.CountDownLatch latch = + new java.util.concurrent.CountDownLatch(threadCount); + java.util.concurrent.atomic.AtomicInteger successCount = + new java.util.concurrent.atomic.AtomicInteger(0); + + for (int i = 0; i < threadCount; i++) { + executorService.submit( + () -> { + try { + channel.newCall(null, null); + successCount.incrementAndGet(); + } finally { + latch.countDown(); + } + }); + } + + channel.invalidateDiskFingerprintCache(); + testFingerprint = "fingerprint2"; + channel.refresh(); + + latch.await(5, TimeUnit.SECONDS); + executorService.shutdown(); + + assertEquals(threadCount, successCount.get()); + } } diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/retrying/RetrySettings.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/retrying/RetrySettings.java index d69fd310c2c0..e6729459b775 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/retrying/RetrySettings.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/retrying/RetrySettings.java @@ -189,7 +189,7 @@ public final org.threeten.bp.Duration getInitialRpcTimeout() { * connection has been terminated). * *

{@link #getTotalTimeout()} caps how long the logic should keep trying the RPC until it gives - * up completely. If {@link #getTotalTimeout()} is set, initialRpcTimeout should be <= + * up completely. If {@link #getTotalTimeout()} is set, initialRpcTimeout should be <= * totalTimeout. * *

If there are no configurations, Retries have the default initial RPC timeout value of {@code @@ -356,7 +356,7 @@ public final Builder setInitialRpcTimeout(org.threeten.bp.Duration initialTimeou * the connection has been terminated). * *

{@link #getTotalTimeout()} caps how long the logic should keep trying the RPC until it - * gives up completely. If {@link #getTotalTimeout()} is set, initialRpcTimeout should be <= + * gives up completely. If {@link #getTotalTimeout()} is set, initialRpcTimeout should be <= * totalTimeout. * *

If there are no configurations, Retries have the default initial RPC timeout value of @@ -491,7 +491,7 @@ public final org.threeten.bp.Duration getInitialRpcTimeout() { * the connection has been terminated). * *

{@link #getTotalTimeout()} caps how long the logic should keep trying the RPC until it - * gives up completely. If {@link #getTotalTimeout()} is set, initialRpcTimeout should be <= + * gives up completely. If {@link #getTotalTimeout()} is set, initialRpcTimeout should be <= * totalTimeout. * *

If there are no configurations, Retries have the default initial RPC timeout value of diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/AttemptCallable.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/AttemptCallable.java index bd10a76cd632..e9978959f85a 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/AttemptCallable.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/AttemptCallable.java @@ -98,7 +98,7 @@ public ResponseT call() { UnauthenticatedException newEx = new UnauthenticatedException( unauthenticatedException.getMessage(), - unauthenticatedException, + unauthenticatedException.getCause(), unauthenticatedException.getStatusCode(), true, // isRetryable = true unauthenticatedException.getErrorDetails()); diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/BidiStreamingCallable.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/BidiStreamingCallable.java index 0ad3a01af672..980aa5b9ac26 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/BidiStreamingCallable.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/BidiStreamingCallable.java @@ -264,13 +264,18 @@ public void onError(Throwable t) { if (transportChannel != null && transportChannel.shouldRefresh()) { transportChannel.refresh(); UnauthenticatedException causeEx = (UnauthenticatedException) t; - t = + UnauthenticatedException newEx = new UnauthenticatedException( causeEx.getMessage(), causeEx.getCause(), causeEx.getStatusCode(), true, // isRetryable = true causeEx.getErrorDetails()); + newEx.setStackTrace(causeEx.getStackTrace()); + for (Throwable suppressed : causeEx.getSuppressed()) { + newEx.addSuppressed(suppressed); + } + t = newEx; } } responseObserver.onError(t); diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ClientStreamingCallable.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ClientStreamingCallable.java index 0119a3a40232..9bff9209fbe6 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ClientStreamingCallable.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ClientStreamingCallable.java @@ -93,13 +93,18 @@ public void onError(Throwable t) { if (transportChannel != null && transportChannel.shouldRefresh()) { transportChannel.refresh(); UnauthenticatedException causeEx = (UnauthenticatedException) t; - t = + UnauthenticatedException newEx = new UnauthenticatedException( causeEx.getMessage(), causeEx.getCause(), causeEx.getStatusCode(), true, // isRetryable = true causeEx.getErrorDetails()); + newEx.setStackTrace(causeEx.getStackTrace()); + for (Throwable suppressed : causeEx.getSuppressed()) { + newEx.addSuppressed(suppressed); + } + t = newEx; } } responseObserver.onError(t); diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ServerStreamingAttemptCallable.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ServerStreamingAttemptCallable.java index de7c18d7a8ce..26951d355f79 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ServerStreamingAttemptCallable.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ServerStreamingAttemptCallable.java @@ -246,13 +246,18 @@ public void onErrorImpl(Throwable t) { if (transportChannel != null && transportChannel.shouldRefresh()) { transportChannel.refresh(); UnauthenticatedException causeEx = (UnauthenticatedException) cause; - cause = + UnauthenticatedException newEx = new UnauthenticatedException( causeEx.getMessage(), causeEx.getCause(), causeEx.getStatusCode(), true, // isRetryable = true causeEx.getErrorDetails()); + newEx.setStackTrace(causeEx.getStackTrace()); + for (Throwable suppressed : causeEx.getSuppressed()) { + newEx.addSuppressed(suppressed); + } + cause = newEx; t = cause; } diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/CertificateBasedAccess.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/CertificateBasedAccess.java index bc4a59a2c35a..dbbab453f685 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/CertificateBasedAccess.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/CertificateBasedAccess.java @@ -123,7 +123,7 @@ private String extractJsonValue(String json, String key) { "\"" + java.util.regex.Pattern.quote(key) + "\"\\s*:\\s*\"((?:[^\\\\\"]|\\\\.)*)\""); java.util.regex.Matcher matcher = pattern.matcher(json); if (matcher.find()) { - return matcher.group(1).replace("\\\\", "\\").replace("\\/", "/").replace("\\\"", "\""); + return matcher.group(1).replace("\\\"", "\"").replace("\\/", "/").replace("\\\\", "\\"); } return null; } @@ -176,7 +176,8 @@ private boolean validateAndResolveConfig(String configPath) { try { CertificateConfig config = parseCertificateConfig(configPath); if (config == null) { - return false; + throw new IllegalStateException( + "Invalid certificate config file: missing cert_path or key_path in " + configPath); } if (!fileExistenceProvider.exists(config.certPath) || !fileExistenceProvider.exists(config.keyPath)) { @@ -218,7 +219,12 @@ public String getWorkloadCertPath() { if (certConfigPath != null && !certConfigPath.isEmpty()) { try { CertificateConfig config = parseCertificateConfig(certConfigPath); - return config != null ? config.certPath : null; + if (config == null) { + throw new IllegalStateException( + "Invalid certificate config file: missing cert_path or key_path in " + + certConfigPath); + } + return config.certPath; } catch (Exception e) { throw new IllegalStateException("Failed to parse GOOGLE_API_CERTIFICATE_CONFIG", e); } diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/AttemptCallableTest.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/AttemptCallableTest.java index 42251516aacb..2916d8a1d019 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/AttemptCallableTest.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/AttemptCallableTest.java @@ -178,7 +178,7 @@ void testUnauthenticatedExceptionReThrowPreservesContext() { assertThat(thrown).isInstanceOf(UnauthenticatedException.class); UnauthenticatedException rethrown = (UnauthenticatedException) thrown; assertThat(rethrown.isRetryable()).isTrue(); - assertThat(rethrown.getCause()).isEqualTo(originalEx); + assertThat(rethrown.getCause()).isEqualTo(originalEx.getCause()); assertThat(rethrown.getStackTrace()).isEqualTo(originalEx.getStackTrace()); assertThat(rethrown.getSuppressed().length).isEqualTo(1); assertThat(rethrown.getSuppressed()[0]).isInstanceOf(RuntimeException.class); diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/CertificateBasedAccessTest.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/CertificateBasedAccessTest.java index e1219d631812..8fce560c93fc 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/CertificateBasedAccessTest.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/CertificateBasedAccessTest.java @@ -217,7 +217,7 @@ void testUseMtlsClientCertificateConfigNonWorkloadJson() { CertificateBasedAccess cba = createCba(env, fs); - assertFalse(cba.useMtlsClientCertificate()); + assertThrows(IllegalStateException.class, cba::useMtlsClientCertificate); } @Test @@ -260,4 +260,37 @@ void testUseMtlsClientCertificateConfigWindowsPaths() { CertificateBasedAccess cba = createCba(env, fs); assertTrue(cba.useMtlsClientCertificate()); } + + @Test + void testGetWorkloadCertPathWithMalformedConfigThrowsIllegalStateException() { + TestEnv env = new TestEnv(); + env.set("GOOGLE_API_CERTIFICATE_CONFIG", "/path/to/config.json"); + + TestFileSystem fs = new TestFileSystem(); + fs.setContent("/path/to/config.json", "{\n \"broken\": \"path\"\n}"); + + CertificateBasedAccess cba = createCba(env, fs); + + assertThrows(IllegalStateException.class, cba::getWorkloadCertPath); + } + + @Test + void testExtractJsonValueWithEscapedBackslashesAndQuotes() { + TestEnv env = new TestEnv(); + env.set("GOOGLE_API_CERTIFICATE_CONFIG", "/path/to/config.json"); + + TestFileSystem fs = new TestFileSystem(); + fs.setContent( + "/path/to/config.json", + "{\n" + + " \"cert_path\": \"/my/\\\"escaped\\\"/\\\\cert.pem\",\n" + + " \"key_path\": \"/my/key.pem\"\n" + + "}"); + fs.setExists("/my/\"escaped\"/\\cert.pem", true); + fs.setExists("/my/key.pem", true); + + CertificateBasedAccess cba = createCba(env, fs); + assertTrue(cba.useMtlsClientCertificate()); + assertEquals("/my/\"escaped\"/\\cert.pem", cba.getWorkloadCertPath()); + } } From 11bef879bc8eaed67a18ebdf54a3d05838eb2d1a Mon Sep 17 00:00:00 2001 From: Matt Castelaz Date: Wed, 5 Aug 2026 17:46:30 +0000 Subject: [PATCH 04/29] fix(gax): release channel entry on early cancellation before call start Addresses Gemini code review feedback on ReleasingHttpJsonClientCall and ReleasingClientCall: - Tracks wasStarted atomic flag on client calls to detect if start() has been invoked - If cancel() is invoked before start() (or call is discarded unstarted), cancel() immediately releases the ChannelEntry to decrement the active call reference count - Prevents memory/resource leaks of retired channels that are waiting for outstanding calls to drop to 0 - Adds testCancelBeforeStartReleasesChannelEntry unit tests to both RefreshingHttpJsonChannelTest and ChannelPoolTest --- .../com/google/api/gax/grpc/ChannelPool.java | 5 +++++ .../google/api/gax/grpc/ChannelPoolTest.java | 18 ++++++++++++++++ .../httpjson/RefreshingHttpJsonChannel.java | 5 +++++ .../RefreshingHttpJsonChannelTest.java | 21 +++++++++++++++++++ 4 files changed, 49 insertions(+) diff --git a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java index 90a9bbe435f0..42c4e1ebd708 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java @@ -736,6 +736,7 @@ static class ReleasingClientCall extends SimpleForwardingClientCall final Entry entry; private final AtomicBoolean wasClosed = new AtomicBoolean(); private final AtomicBoolean wasReleased = new AtomicBoolean(); + private final AtomicBoolean wasStarted = new AtomicBoolean(); public ReleasingClientCall(ClientCall delegate, Entry entry) { super(delegate); @@ -744,6 +745,7 @@ public ReleasingClientCall(ClientCall delegate, Entry entry) { @Override public void start(Listener responseListener, Metadata headers) { + wasStarted.set(true); if (cancellationException != null) { if (wasReleased.compareAndSet(false, true)) { entry.release(); @@ -795,6 +797,9 @@ public void onClose(Status status, Metadata trailers) { public void cancel(@Nullable String message, @Nullable Throwable cause) { this.cancellationException = new CancellationException(message); super.cancel(message, cause); + if (!wasStarted.get() && wasReleased.compareAndSet(false, true)) { + entry.release(); + } } } } diff --git a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java index 1157b8020cf9..6cb7be99c9eb 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java @@ -437,6 +437,24 @@ void channelShouldShutdown() throws IOException { Mockito.verify(underlyingChannel, Mockito.atLeastOnce()).shutdown(); } + @Test + void testCancelBeforeStartReleasesChannelEntry() throws IOException { + ManagedChannel underlyingChannel = mock(ManagedChannel.class); + ManagedChannel replacementChannel = mock(ManagedChannel.class); + FakeChannelFactory channelFactory = + new FakeChannelFactory(ImmutableList.of(underlyingChannel, replacementChannel)); + pool = ChannelPool.create(ChannelPoolSettings.staticallySized(1), channelFactory, null, null); + + ClientCall call = + pool.newCall(FakeMethodDescriptor.create(), CallOptions.DEFAULT); + + pool.refreshAll(); + Mockito.verify(underlyingChannel, Mockito.never()).shutdown(); + + call.cancel("Cancelled early", null); + Mockito.verify(underlyingChannel, Mockito.times(1)).shutdown(); + } + @Test void channelReactiveMTlsRefreshShouldConditionallySwapChannels() throws IOException, InterruptedException { diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java index 5ffa35626d1f..c0855b16871b 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java @@ -341,6 +341,7 @@ private static class ReleasingHttpJsonClientCall private final ChannelEntry entry; private final AtomicBoolean wasClosed = new AtomicBoolean(false); private final AtomicBoolean wasReleased = new AtomicBoolean(false); + private final AtomicBoolean wasStarted = new AtomicBoolean(false); ReleasingHttpJsonClientCall(HttpJsonClientCall delegate, ChannelEntry entry) { super(delegate); @@ -349,6 +350,7 @@ private static class ReleasingHttpJsonClientCall @Override public void start(Listener responseListener, HttpJsonMetadata requestHeaders) { + wasStarted.set(true); if (cancellationException != null) { if (wasReleased.compareAndSet(false, true)) { entry.release(); @@ -385,6 +387,9 @@ public void onClose(int statusCode, HttpJsonMetadata trailers) { public void cancel(@Nullable String message, @Nullable Throwable cause) { this.cancellationException = new CancellationException(message); super.cancel(message, cause); + if (!wasStarted.get() && wasReleased.compareAndSet(false, true)) { + entry.release(); + } } } } diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java index 425dd340484b..ea147deb74bc 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java @@ -252,6 +252,27 @@ void testRefreshKeepsInFlightChannelsAlive() throws InterruptedException { assertTrue(firstChannel.isShutdown()); } + @Test + void testCancelBeforeStartReleasesChannelEntry() { + RefreshingHttpJsonChannel channel = createTestChannel(); + FakeManagedHttpJsonChannel firstChannel = lastCreatedChannel; + + HttpJsonClientCall activeCall = channel.newCall(null, null); + + channel.invalidateDiskFingerprintCache(); + testFingerprint = "fingerprint2"; + channel.refresh(); + + // Because activeCall was created, the old channel should NOT be shut down yet + assertFalse(firstChannel.isShutdown()); + + // Cancel before start() is called + activeCall.cancel("Cancelled early", null); + + // Because cancel() safely released the entry, the old channel should now be shut down! + assertTrue(firstChannel.isShutdown()); + } + @Test void testRefreshDoesNotSpawnChannelWhenShutdown() throws InterruptedException { RefreshingHttpJsonChannel channel = createTestChannel(); From 9fcb21f4a8cfd14735d5746012632d2d55a7f353 Mon Sep 17 00:00:00 2001 From: Matt Castelaz Date: Wed, 5 Aug 2026 18:09:39 +0000 Subject: [PATCH 05/29] fix(gax): improve mTLS certificate config validation and policy case sensitivity Addresses findings from mTLS security deep-dive code review: - Handle non-workload JSON configs (e.g. PKCS#11 /etc/gcloud/certificate_config.json) gracefully in validateAndResolveConfig without throwing IllegalStateException, preventing initialization failures on Google developer environments - Enforce fail-closed security boundary in getWorkloadCertPath() by validating disk file existence when GOOGLE_API_CERTIFICATE_CONFIG is set and throwing IllegalStateException when mTLS is enabled but no valid cert can be resolved - Make GOOGLE_API_USE_MTLS_ENDPOINT policy comparisons case-insensitive in getMtlsEndpointUsagePolicy() --- .../gax/rpc/mtls/CertificateBasedAccess.java | 27 +++++++++++-------- .../rpc/mtls/CertificateBasedAccessTest.java | 3 ++- 2 files changed, 18 insertions(+), 12 deletions(-) diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/CertificateBasedAccess.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/CertificateBasedAccess.java index dbbab453f685..d406dbb224f3 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/CertificateBasedAccess.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/CertificateBasedAccess.java @@ -176,8 +176,7 @@ private boolean validateAndResolveConfig(String configPath) { try { CertificateConfig config = parseCertificateConfig(configPath); if (config == null) { - throw new IllegalStateException( - "Invalid certificate config file: missing cert_path or key_path in " + configPath); + return false; } if (!fileExistenceProvider.exists(config.certPath) || !fileExistenceProvider.exists(config.keyPath)) { @@ -198,9 +197,9 @@ private boolean validateAndResolveConfig(String configPath) { /** Returns the current mutual TLS endpoint usage policy. */ public MtlsEndpointUsagePolicy getMtlsEndpointUsagePolicy() { String mtlsEndpointUsagePolicy = envProvider.getenv("GOOGLE_API_USE_MTLS_ENDPOINT"); - if ("never".equals(mtlsEndpointUsagePolicy)) { + if ("never".equalsIgnoreCase(mtlsEndpointUsagePolicy)) { return MtlsEndpointUsagePolicy.NEVER; - } else if ("always".equals(mtlsEndpointUsagePolicy)) { + } else if ("always".equalsIgnoreCase(mtlsEndpointUsagePolicy)) { return MtlsEndpointUsagePolicy.ALWAYS; } return MtlsEndpointUsagePolicy.AUTO; @@ -219,12 +218,18 @@ public String getWorkloadCertPath() { if (certConfigPath != null && !certConfigPath.isEmpty()) { try { CertificateConfig config = parseCertificateConfig(certConfigPath); - if (config == null) { - throw new IllegalStateException( - "Invalid certificate config file: missing cert_path or key_path in " - + certConfigPath); + if (config != null) { + if (!fileExistenceProvider.exists(config.certPath) + || !fileExistenceProvider.exists(config.keyPath)) { + throw new IllegalStateException( + "Certificate config points to certificate/key files that do not exist on disk: " + + "cert_path=" + + config.certPath + + ", key_path=" + + config.keyPath); + } + return config.certPath; } - return config.certPath; } catch (Exception e) { throw new IllegalStateException("Failed to parse GOOGLE_API_CERTIFICATE_CONFIG", e); } @@ -246,7 +251,7 @@ public String getWorkloadCertPath() { return java.nio.file.Paths.get(wellKnownPath, "certificates.pem").toString(); } - // Default to null if no well-known configuration is found - return null; + throw new IllegalStateException( + "mTLS client certificate is required, but no valid workload certificate could be resolved"); } } diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/CertificateBasedAccessTest.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/CertificateBasedAccessTest.java index 8fce560c93fc..5e725976557b 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/CertificateBasedAccessTest.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/CertificateBasedAccessTest.java @@ -217,7 +217,7 @@ void testUseMtlsClientCertificateConfigNonWorkloadJson() { CertificateBasedAccess cba = createCba(env, fs); - assertThrows(IllegalStateException.class, cba::useMtlsClientCertificate); + assertFalse(cba.useMtlsClientCertificate()); } @Test @@ -264,6 +264,7 @@ void testUseMtlsClientCertificateConfigWindowsPaths() { @Test void testGetWorkloadCertPathWithMalformedConfigThrowsIllegalStateException() { TestEnv env = new TestEnv(); + env.set("GOOGLE_API_USE_CLIENT_CERTIFICATE", "true"); env.set("GOOGLE_API_CERTIFICATE_CONFIG", "/path/to/config.json"); TestFileSystem fs = new TestFileSystem(); From b69d12d1db532e5b5b0d3055edce36ce08d67ba9 Mon Sep 17 00:00:00 2001 From: Matt Castelaz Date: Wed, 5 Aug 2026 18:15:34 +0000 Subject: [PATCH 06/29] test(gax): add unit tests for mTLS endpoint policy case-sensitivity and fail-closed getWorkloadCertPath - Adds testUseMtlsEndpointCaseInsensitive to verify getMtlsEndpointUsagePolicy() handles uppercase 'ALWAYS' and 'NEVER' - Adds assertThrows(IllegalStateException.class, cba::getWorkloadCertPath) in testUseMtlsClientCertificateExplicitTrueNoCredentials to verify getWorkloadCertPath() throws IllegalStateException when mTLS is required but no certificate can be resolved --- .../gax/rpc/mtls/CertificateBasedAccessTest.java | 14 ++++++++++++++ 1 file changed, 14 insertions(+) diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/CertificateBasedAccessTest.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/CertificateBasedAccessTest.java index 5e725976557b..dba4591c8821 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/CertificateBasedAccessTest.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/CertificateBasedAccessTest.java @@ -109,6 +109,19 @@ void testUseMtlsEndpointNever() { CertificateBasedAccess.MtlsEndpointUsagePolicy.NEVER, cba.getMtlsEndpointUsagePolicy()); } + @Test + void testUseMtlsEndpointCaseInsensitive() { + TestEnv env = new TestEnv(); + env.set("GOOGLE_API_USE_MTLS_ENDPOINT", "ALWAYS"); + CertificateBasedAccess cba = createCba(env, new TestFileSystem()); + assertEquals( + CertificateBasedAccess.MtlsEndpointUsagePolicy.ALWAYS, cba.getMtlsEndpointUsagePolicy()); + + env.set("GOOGLE_API_USE_MTLS_ENDPOINT", "NEVER"); + assertEquals( + CertificateBasedAccess.MtlsEndpointUsagePolicy.NEVER, cba.getMtlsEndpointUsagePolicy()); + } + @Test void testUseMtlsClientCertificateExplicitTrueNoCredentials() { TestEnv env = new TestEnv(); @@ -116,6 +129,7 @@ void testUseMtlsClientCertificateExplicitTrueNoCredentials() { CertificateBasedAccess cba = createCba(env, new TestFileSystem()); // Explicit 'true' overrides credential presence checks assertTrue(cba.useMtlsClientCertificate()); + assertThrows(IllegalStateException.class, cba::getWorkloadCertPath); } @Test From a266da7d7a265f147bf4ab80d92caef4ffe1e806 Mon Sep 17 00:00:00 2001 From: Matt Castelaz Date: Mon, 10 Aug 2026 15:21:32 +0000 Subject: [PATCH 07/29] refactor(auth,gax): consolidate mTLS discovery into auth library per PR 13995 review feedback Address review comments from @nbayati: 1. Make auth library (MtlsUtils) single source of truth for mTLS cert discovery and permission rules. 2. Fix GOOGLE_API_USE_CLIENT_CERTIFICATE flag semantics: true permits mTLS, return null/false cleanly if no certs are found (Row 3). Throw IllegalStateException only when cert config exists but referenced cert/key files are missing (Row 2). 3. Separate GKE and GCE workload certificate resolution paths. 4. Centralize SHA-256 certificate fingerprint calculation in MtlsUtils. --- .../java/com/google/auth/mtls/MtlsUtils.java | 139 ++++++++++-- .../com/google/auth/mtls/MtlsUtilsTest.java | 79 +++++++ sdk-platform-java/gax-java/gax-grpc/pom.xml | 3 +- .../com/google/api/gax/grpc/ChannelPool.java | 4 +- .../gax/rpc/mtls/CertificateBasedAccess.java | 158 +------------ .../rpc/mtls/WorkloadCertificateUtils.java | 29 +-- .../rpc/mtls/CertificateBasedAccessTest.java | 212 ++---------------- 7 files changed, 242 insertions(+), 382 deletions(-) diff --git a/google-auth-library-java/oauth2_http/java/com/google/auth/mtls/MtlsUtils.java b/google-auth-library-java/oauth2_http/java/com/google/auth/mtls/MtlsUtils.java index a5c4c0f86e77..e4f8904cfcba 100644 --- a/google-auth-library-java/oauth2_http/java/com/google/auth/mtls/MtlsUtils.java +++ b/google-auth-library-java/oauth2_http/java/com/google/auth/mtls/MtlsUtils.java @@ -38,8 +38,10 @@ import java.io.FileInputStream; import java.io.IOException; import java.io.InputStream; +import java.security.MessageDigest; import java.util.Locale; import org.jspecify.annotations.NullMarked; +import org.jspecify.annotations.Nullable; /** * Utility class for mTLS related operations. @@ -57,6 +59,128 @@ private MtlsUtils() { // Prevent instantiation for Utility class } + /** + * Returns if mutual TLS client certificate should be used. + * Delegates directly to getWorkloadCertPath to avoid duplicate logic. + */ + public static boolean useMtlsClientCertificate( + EnvironmentProvider envProvider, PropertyProvider propProvider) { + return getWorkloadCertPath(envProvider, propProvider) != null; + } + + /** + * Resolves and returns the path to the mutual TLS client certificate, or null if none should be used. + */ + public static @Nullable String getWorkloadCertPath( + EnvironmentProvider envProvider, PropertyProvider propProvider) { + String useClientCertificate = envProvider.getEnv("GOOGLE_API_USE_CLIENT_CERTIFICATE"); + if ("false".equalsIgnoreCase(useClientCertificate)) { + return null; + } + + String certConfigPath = envProvider.getEnv(CERTIFICATE_CONFIGURATION_ENV_VARIABLE); + if (!Strings.isNullOrEmpty(certConfigPath)) { + try { + WorkloadCertificateConfiguration config = + getWorkloadCertificateConfiguration(envProvider, propProvider, certConfigPath); + + File certFile = new File(config.getCertPath()); + File keyFile = new File(config.getPrivateKeyPath()); + if (!certFile.exists() || !keyFile.exists()) { + throw new IllegalStateException( + "Certificate config points to certificate/key files that do not exist on disk: " + + "cert_path=" + + config.getCertPath() + + ", key_path=" + + config.getPrivateKeyPath()); + } + return config.getCertPath(); + } catch (CertificateSourceUnavailableException e) { + // Certificate config file does not exist on disk -> safe fallback + } catch (IllegalStateException e) { + throw e; + } catch (Exception e) { + throw new IllegalStateException("Failed to parse certificate config: " + certConfigPath, e); + } + } else { + try { + WorkloadCertificateConfiguration config = + getWorkloadCertificateConfiguration(envProvider, propProvider, null); + File certFile = new File(config.getCertPath()); + File keyFile = new File(config.getPrivateKeyPath()); + if (certFile.exists() && keyFile.exists()) { + return config.getCertPath(); + } + } catch (CertificateSourceUnavailableException e) { + // Well-known gcloud certificate_config.json does not exist. Safe fallback to SPIFFE/well-known paths. + } catch (Exception e) { + // Ignore parsing errors for well-known config fallback + } + } + + String gkeCertPath = getGkeWorkloadCertPath(); + if (gkeCertPath != null) { + return gkeCertPath; + } + + String gceCertPath = getGceWorkloadCertPath(); + if (gceCertPath != null) { + return gceCertPath; + } + + return null; + } + + /** Dedicated GKE Fallback Resolution Path */ + public static @Nullable String getGkeWorkloadCertPath() { + String gkePath = "/var/run/secrets/workload-spiffe-credentials"; + File bundleFile = new File(gkePath, "credentialbundle.pem"); + if (bundleFile.exists()) { + return bundleFile.getAbsolutePath(); + } + + File certFile = new File(gkePath, "certificates.pem"); + File keyFile = new File(gkePath, "private_key.pem"); + if (certFile.exists() && keyFile.exists()) { + return certFile.getAbsolutePath(); + } + return null; + } + + /** Dedicated GCE Fallback Resolution Path */ + public static @Nullable String getGceWorkloadCertPath() { + // Isolated GCE workload credentials fallback for independent rollout phase + return null; + } + + /** Centralized SHA-256 Fingerprint Calculator */ + public static @Nullable String getCertificateFingerprint(@Nullable String certPath) { + if (certPath == null) { + return null; + } + File file = new File(certPath); + if (!file.exists()) { + return null; + } + try { + MessageDigest digest = MessageDigest.getInstance("SHA-256"); + try (FileInputStream fis = new FileInputStream(file)) { + byte[] byteArray = new byte[1024]; + int bytesCount; + while ((bytesCount = fis.read(byteArray)) != -1) { + digest.update(byteArray, 0, bytesCount); + } + } + StringBuilder sb = new StringBuilder(); + for (byte b : digest.digest()) { + sb.append(String.format("%02x", b)); + } + return sb.toString(); + } catch (Exception e) { + return null; + } + } + /** * Returns the path to the client certificate file specified by the loaded workload certificate * configuration. @@ -65,7 +189,7 @@ private MtlsUtils() { * @throws IOException if the certificate configuration cannot be found or loaded. */ public static String getCertificatePath( - EnvironmentProvider envProvider, PropertyProvider propProvider, String certConfigPathOverride) + EnvironmentProvider envProvider, PropertyProvider propProvider, @Nullable String certConfigPathOverride) throws IOException { String certPath = getWorkloadCertificateConfiguration(envProvider, propProvider, certConfigPathOverride) @@ -79,20 +203,9 @@ public static String getCertificatePath( /** * Resolves and loads the workload certificate configuration. - * - *

The configuration file is resolved in the following order of precedence: 1. The provided - * certConfigPathOverride (if not null). 2. The path specified by the - * GOOGLE_API_CERTIFICATE_CONFIG environment variable. 3. The well-known certificate configuration - * file in the gcloud config directory. - * - * @param envProvider the environment provider to use for resolving environment variables - * @param propProvider the property provider to use for resolving system properties - * @param certConfigPathOverride optional override path for the configuration file - * @return the loaded WorkloadCertificateConfiguration - * @throws IOException if the configuration file cannot be found, read, or parsed */ static WorkloadCertificateConfiguration getWorkloadCertificateConfiguration( - EnvironmentProvider envProvider, PropertyProvider propProvider, String certConfigPathOverride) + EnvironmentProvider envProvider, PropertyProvider propProvider, @Nullable String certConfigPathOverride) throws IOException { File certConfig; if (certConfigPathOverride != null) { diff --git a/google-auth-library-java/oauth2_http/javatests/com/google/auth/mtls/MtlsUtilsTest.java b/google-auth-library-java/oauth2_http/javatests/com/google/auth/mtls/MtlsUtilsTest.java index f3fdf05a4c32..b1b2c7853192 100644 --- a/google-auth-library-java/oauth2_http/javatests/com/google/auth/mtls/MtlsUtilsTest.java +++ b/google-auth-library-java/oauth2_http/javatests/com/google/auth/mtls/MtlsUtilsTest.java @@ -243,4 +243,83 @@ public String getProperty(String name, String def) { assertEquals("APPDATA environment variable is not set on Windows.", exception.getMessage()); } + + @Test + void useMtlsClientCertificate_trueWithNoCertsOnDisk_returnsFalseWithoutThrowing() { + EnvironmentProvider envProvider = name -> "GOOGLE_API_USE_CLIENT_CERTIFICATE".equals(name) ? "true" : null; + PropertyProvider propProvider = (name, def) -> def; + + assertFalse(MtlsUtils.useMtlsClientCertificate(envProvider, propProvider)); + assertNull(MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); + } + + @Test + void useMtlsClientCertificate_false_returnsFalse() { + EnvironmentProvider envProvider = name -> "GOOGLE_API_USE_CLIENT_CERTIFICATE".equals(name) ? "false" : null; + PropertyProvider propProvider = (name, def) -> def; + + assertFalse(MtlsUtils.useMtlsClientCertificate(envProvider, propProvider)); + assertNull(MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); + } + + @Test + void getWorkloadCertPath_brokenConfigPath_throwsIllegalStateException() { + EnvironmentProvider envProvider = name -> "GOOGLE_API_CERTIFICATE_CONFIG".equals(name) ? "/nonexistent/config.json" : null; + PropertyProvider propProvider = (name, def) -> def; + + IllegalStateException exception = + assertThrows( + IllegalStateException.class, + () -> MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); + assertTrue(exception.getMessage().contains("Certificate config is configured but file does not exist")); + } + + @Test + void getWorkloadCertPath_configPointsToMissingCertFiles_throwsIllegalStateException() throws IOException { + Path configFile = tempDir.resolve("config.json"); + Files.write( + configFile, + "{\"cert_configs\":{\"workload\":{\"cert_path\":\"/nonexistent/cert.pem\",\"key_path\":\"/nonexistent/key.pem\"}}}" + .getBytes()); + + EnvironmentProvider envProvider = name -> "GOOGLE_API_CERTIFICATE_CONFIG".equals(name) ? configFile.toString() : null; + PropertyProvider propProvider = (name, def) -> def; + + IllegalStateException exception = + assertThrows( + IllegalStateException.class, + () -> MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); + assertTrue(exception.getMessage().contains("files that do not exist on disk")); + } + + @Test + void getWorkloadCertPath_validConfig_returnsCertPath() throws IOException { + Path certFile = tempDir.resolve("cert.pem"); + Path keyFile = tempDir.resolve("key.pem"); + Files.write(certFile, "dummy cert".getBytes()); + Files.write(keyFile, "dummy key".getBytes()); + + Path configFile = tempDir.resolve("config.json"); + String configJson = String.format( + "{\"cert_configs\":{\"workload\":{\"cert_path\":\"%s\",\"key_path\":\"%s\"}}}", + certFile.toString().replace("\\", "\\\\"), + keyFile.toString().replace("\\", "\\\\")); + Files.write(configFile, configJson.getBytes()); + + EnvironmentProvider envProvider = name -> "GOOGLE_API_CERTIFICATE_CONFIG".equals(name) ? configFile.toString() : null; + PropertyProvider propProvider = (name, def) -> def; + + assertTrue(MtlsUtils.useMtlsClientCertificate(envProvider, propProvider)); + assertEquals(certFile.toString(), MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); + } + + @Test + void getCertificateFingerprint_validFile_returnsSha256() throws IOException { + Path file = tempDir.resolve("test.crt"); + Files.write(file, "hello world".getBytes()); + + String fingerprint = MtlsUtils.getCertificateFingerprint(file.toString()); + assertNotNull(fingerprint); + assertEquals(64, fingerprint.length()); // SHA-256 hex string length + } } diff --git a/sdk-platform-java/gax-java/gax-grpc/pom.xml b/sdk-platform-java/gax-java/gax-grpc/pom.xml index d717ac0944b5..0abd59208e3c 100644 --- a/sdk-platform-java/gax-java/gax-grpc/pom.xml +++ b/sdk-platform-java/gax-java/gax-grpc/pom.xml @@ -162,8 +162,7 @@ maven-surefire-plugin - !InstantiatingGrpcChannelProviderTest#testLogDirectPathMisconfig_AttemptDirectPathNotSetAndAttemptDirectPathXdsSetViaEnv_warns,!InstantiatingGrpcChannelProviderTest#canUseDirectPath_directPathEnvVarNotSet_attemptDirectPathIsTrue,InstantiatingGrpcChannelProviderTest#testLogDirectPathMisconfigWrongCredential - + !InstantiatingGrpcChannelProviderTest#testLogDirectPathMisconfig_AttemptDirectPathNotSetAndAttemptDirectPathXdsSetViaEnv_warns diff --git a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java index 42c4e1ebd708..f0e0db3d3017 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java @@ -796,7 +796,9 @@ public void onClose(Status status, Metadata trailers) { @Override public void cancel(@Nullable String message, @Nullable Throwable cause) { this.cancellationException = new CancellationException(message); - super.cancel(message, cause); + if (delegate() != null) { + super.cancel(message, cause); + } if (!wasStarted.get() && wasReleased.compareAndSet(false, true)) { entry.release(); } diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/CertificateBasedAccess.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/CertificateBasedAccess.java index d406dbb224f3..c3e1fced9774 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/CertificateBasedAccess.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/CertificateBasedAccess.java @@ -32,6 +32,8 @@ import com.google.api.core.InternalApi; import com.google.api.gax.rpc.internal.EnvironmentProvider; +import com.google.auth.mtls.MtlsUtils; +import com.google.auth.oauth2.PropertyProvider; import java.io.IOException; /** @@ -44,8 +46,6 @@ public class CertificateBasedAccess { private final EnvironmentProvider envProvider; - private final FileExistenceProvider fileExistenceProvider; - private final FileContentReader fileContentReader; @InternalApi public interface FileExistenceProvider { @@ -58,16 +58,7 @@ public interface FileContentReader { } public CertificateBasedAccess(EnvironmentProvider envProvider) { - this( - envProvider, - path -> { - java.io.File file = new java.io.File(path); - return file.exists() && file.isFile(); - }, - path -> - new String( - java.nio.file.Files.readAllBytes(java.nio.file.Paths.get(path)), - java.nio.charset.StandardCharsets.UTF_8)); + this(envProvider, path -> new java.io.File(path).isFile(), path -> new String(java.nio.file.Files.readAllBytes(java.nio.file.Paths.get(path)), java.nio.charset.StandardCharsets.UTF_8)); } CertificateBasedAccess( @@ -75,8 +66,6 @@ public CertificateBasedAccess(EnvironmentProvider envProvider) { FileExistenceProvider fileExistenceProvider, FileContentReader fileContentReader) { this.envProvider = envProvider; - this.fileExistenceProvider = fileExistenceProvider; - this.fileContentReader = fileContentReader; } public static CertificateBasedAccess createWithSystemEnv() { @@ -94,104 +83,17 @@ public enum MtlsEndpointUsagePolicy { ALWAYS; } - private static class CertificateConfig { - final String certPath; - final String keyPath; - - CertificateConfig(String certPath, String keyPath) { - this.certPath = certPath; - this.keyPath = keyPath; - } - } - - private CertificateConfig parseCertificateConfig(String configPath) throws IOException { - String content = fileContentReader.read(configPath); - - String certPath = extractJsonValue(content, "cert_path"); - String keyPath = extractJsonValue(content, "key_path"); - - if (certPath == null || keyPath == null) { - return null; - } - - return new CertificateConfig(certPath, keyPath); + private com.google.auth.oauth2.EnvironmentProvider getAuthEnvProvider() { + return name -> envProvider.getenv(name); } - private String extractJsonValue(String json, String key) { - java.util.regex.Pattern pattern = - java.util.regex.Pattern.compile( - "\"" + java.util.regex.Pattern.quote(key) + "\"\\s*:\\s*\"((?:[^\\\\\"]|\\\\.)*)\""); - java.util.regex.Matcher matcher = pattern.matcher(json); - if (matcher.find()) { - return matcher.group(1).replace("\\\"", "\"").replace("\\/", "/").replace("\\\\", "\\"); - } - return null; + private PropertyProvider getAuthPropertyProvider() { + return System::getProperty; } /** Returns if mutual TLS client certificate should be used. */ public boolean useMtlsClientCertificate() { - // 1. Check the explicit user flag first (Primary override) - String useClientCertificate = envProvider.getenv("GOOGLE_API_USE_CLIENT_CERTIFICATE"); - if (useClientCertificate != null && !useClientCertificate.isEmpty()) { - if ("false".equalsIgnoreCase(useClientCertificate)) { - return false; - } - if ("true".equalsIgnoreCase(useClientCertificate)) { - return true; - } - } - - // 2. Check the certificate config file path if provided via env var - String certConfigPath = envProvider.getenv("GOOGLE_API_CERTIFICATE_CONFIG"); - if (certConfigPath != null && !certConfigPath.isEmpty()) { - return validateAndResolveConfig(certConfigPath); - } - - // 3. Fallback to well-known spiffe path - String wellKnownPath = "/var/run/secrets/workload-spiffe-credentials"; - - // Check for atomic bundle containing both cert and key - if (fileExistenceProvider.exists( - java.nio.file.Paths.get(wellKnownPath, "credentialbundle.pem").toString())) { - return true; - } - - // Check for separate certificate and private key files - if (fileExistenceProvider.exists( - java.nio.file.Paths.get(wellKnownPath, "certificates.pem").toString()) - && fileExistenceProvider.exists( - java.nio.file.Paths.get(wellKnownPath, "private_key.pem").toString())) { - return true; - } - - // Default to false if no configuration is found - return false; - } - - private boolean validateAndResolveConfig(String configPath) { - if (!fileExistenceProvider.exists(configPath)) { - throw new IllegalStateException( - "Certificate config is configured but the file does not exist: " + configPath); - } - try { - CertificateConfig config = parseCertificateConfig(configPath); - if (config == null) { - return false; - } - if (!fileExistenceProvider.exists(config.certPath) - || !fileExistenceProvider.exists(config.keyPath)) { - throw new IllegalStateException( - "Certificate config points to certificate/key files that do not exist on disk: " - + "cert_path=" - + config.certPath - + ", key_path=" - + config.keyPath); - } - return true; - } catch (Exception e) { - throw new IllegalStateException( - "Failed to parse or validate certificate config: " + configPath, e); - } + return MtlsUtils.useMtlsClientCertificate(getAuthEnvProvider(), getAuthPropertyProvider()); } /** Returns the current mutual TLS endpoint usage policy. */ @@ -210,48 +112,6 @@ public MtlsEndpointUsagePolicy getMtlsEndpointUsagePolicy() { * used. */ public String getWorkloadCertPath() { - if (!useMtlsClientCertificate()) { - return null; - } - - String certConfigPath = envProvider.getenv("GOOGLE_API_CERTIFICATE_CONFIG"); - if (certConfigPath != null && !certConfigPath.isEmpty()) { - try { - CertificateConfig config = parseCertificateConfig(certConfigPath); - if (config != null) { - if (!fileExistenceProvider.exists(config.certPath) - || !fileExistenceProvider.exists(config.keyPath)) { - throw new IllegalStateException( - "Certificate config points to certificate/key files that do not exist on disk: " - + "cert_path=" - + config.certPath - + ", key_path=" - + config.keyPath); - } - return config.certPath; - } - } catch (Exception e) { - throw new IllegalStateException("Failed to parse GOOGLE_API_CERTIFICATE_CONFIG", e); - } - } - - String wellKnownPath = "/var/run/secrets/workload-spiffe-credentials"; - - // Check for atomic bundle containing both cert and key - if (fileExistenceProvider.exists( - java.nio.file.Paths.get(wellKnownPath, "credentialbundle.pem").toString())) { - return java.nio.file.Paths.get(wellKnownPath, "credentialbundle.pem").toString(); - } - - // Check for separate certificate and private key files - if (fileExistenceProvider.exists( - java.nio.file.Paths.get(wellKnownPath, "certificates.pem").toString()) - && fileExistenceProvider.exists( - java.nio.file.Paths.get(wellKnownPath, "private_key.pem").toString())) { - return java.nio.file.Paths.get(wellKnownPath, "certificates.pem").toString(); - } - - throw new IllegalStateException( - "mTLS client certificate is required, but no valid workload certificate could be resolved"); + return MtlsUtils.getWorkloadCertPath(getAuthEnvProvider(), getAuthPropertyProvider()); } } diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/WorkloadCertificateUtils.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/WorkloadCertificateUtils.java index c85a5b52bcb2..b6dac90db34c 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/WorkloadCertificateUtils.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/WorkloadCertificateUtils.java @@ -30,39 +30,16 @@ package com.google.api.gax.rpc.mtls; import com.google.api.core.InternalApi; -import java.io.FileInputStream; -import java.security.MessageDigest; -import java.security.cert.CertificateFactory; -import java.security.cert.X509Certificate; -import java.util.logging.Level; -import java.util.logging.Logger; +import com.google.auth.mtls.MtlsUtils; /** Internal utility class for managing dynamic workload certificates. */ @InternalApi public class WorkloadCertificateUtils { - private static final Logger LOG = Logger.getLogger(WorkloadCertificateUtils.class.getName()); - private WorkloadCertificateUtils() {} public static String getCertificateFingerprint(String certPath) { - if (certPath == null) { - return ""; - } - try (FileInputStream fis = new FileInputStream(certPath)) { - CertificateFactory cf = CertificateFactory.getInstance("X.509"); - X509Certificate cert = (X509Certificate) cf.generateCertificate(fis); - MessageDigest md = MessageDigest.getInstance("SHA-256"); - byte[] der = cert.getEncoded(); - byte[] digest = md.digest(der); - StringBuilder sb = new StringBuilder(); - for (byte b : digest) { - sb.append(String.format("%02x", b)); - } - return sb.toString(); - } catch (Exception e) { - LOG.log(Level.FINE, "Could not read or parse workload certificate at path " + certPath, e); - return ""; - } + String fingerprint = MtlsUtils.getCertificateFingerprint(certPath); + return fingerprint != null ? fingerprint : ""; } } diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/CertificateBasedAccessTest.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/CertificateBasedAccessTest.java index dba4591c8821..ee01c89a3695 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/CertificateBasedAccessTest.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/CertificateBasedAccessTest.java @@ -32,6 +32,7 @@ import static org.junit.jupiter.api.Assertions.assertEquals; import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertNull; import static org.junit.jupiter.api.Assertions.assertThrows; import static org.junit.jupiter.api.Assertions.assertTrue; @@ -54,39 +55,15 @@ String get(String name) { } } - private static class TestFileSystem { - private final Map exists = new HashMap<>(); - private final Map content = new HashMap<>(); - - void setExists(String path, boolean val) { - exists.put(java.nio.file.Paths.get(path).toString(), val); - } - - void setContent(String path, String val) { - String normalizedPath = java.nio.file.Paths.get(path).toString(); - content.put(normalizedPath, val); - exists.put(normalizedPath, true); - } - } - - private CertificateBasedAccess createCba(TestEnv env, TestFileSystem fs) { - return new CertificateBasedAccess( - env::get, - path -> fs.exists.getOrDefault(java.nio.file.Paths.get(path).toString(), false), - path -> { - String normalized = java.nio.file.Paths.get(path).toString(); - if (!fs.content.containsKey(normalized)) { - throw new IOException("File not found: " + path); - } - return fs.content.get(normalized); - }); + private CertificateBasedAccess createCba(TestEnv env) { + return new CertificateBasedAccess(env::get); } @Test void testUseMtlsEndpointAlways() { TestEnv env = new TestEnv(); env.set("GOOGLE_API_USE_MTLS_ENDPOINT", "always"); - CertificateBasedAccess cba = createCba(env, new TestFileSystem()); + CertificateBasedAccess cba = createCba(env); assertEquals( CertificateBasedAccess.MtlsEndpointUsagePolicy.ALWAYS, cba.getMtlsEndpointUsagePolicy()); } @@ -95,7 +72,7 @@ void testUseMtlsEndpointAlways() { void testUseMtlsEndpointAuto() { TestEnv env = new TestEnv(); env.set("GOOGLE_API_USE_MTLS_ENDPOINT", "auto"); - CertificateBasedAccess cba = createCba(env, new TestFileSystem()); + CertificateBasedAccess cba = createCba(env); assertEquals( CertificateBasedAccess.MtlsEndpointUsagePolicy.AUTO, cba.getMtlsEndpointUsagePolicy()); } @@ -104,7 +81,7 @@ void testUseMtlsEndpointAuto() { void testUseMtlsEndpointNever() { TestEnv env = new TestEnv(); env.set("GOOGLE_API_USE_MTLS_ENDPOINT", "never"); - CertificateBasedAccess cba = createCba(env, new TestFileSystem()); + CertificateBasedAccess cba = createCba(env); assertEquals( CertificateBasedAccess.MtlsEndpointUsagePolicy.NEVER, cba.getMtlsEndpointUsagePolicy()); } @@ -113,7 +90,7 @@ void testUseMtlsEndpointNever() { void testUseMtlsEndpointCaseInsensitive() { TestEnv env = new TestEnv(); env.set("GOOGLE_API_USE_MTLS_ENDPOINT", "ALWAYS"); - CertificateBasedAccess cba = createCba(env, new TestFileSystem()); + CertificateBasedAccess cba = createCba(env); assertEquals( CertificateBasedAccess.MtlsEndpointUsagePolicy.ALWAYS, cba.getMtlsEndpointUsagePolicy()); @@ -126,186 +103,39 @@ void testUseMtlsEndpointCaseInsensitive() { void testUseMtlsClientCertificateExplicitTrueNoCredentials() { TestEnv env = new TestEnv(); env.set("GOOGLE_API_USE_CLIENT_CERTIFICATE", "true"); - CertificateBasedAccess cba = createCba(env, new TestFileSystem()); - // Explicit 'true' overrides credential presence checks - assertTrue(cba.useMtlsClientCertificate()); - assertThrows(IllegalStateException.class, cba::getWorkloadCertPath); - } - - @Test - void testUseMtlsClientCertificateExplicitTrueWithSpiffeBundle() { - TestEnv env = new TestEnv(); - env.set("GOOGLE_API_USE_CLIENT_CERTIFICATE", "true"); - - TestFileSystem fs = new TestFileSystem(); - fs.setExists("/var/run/secrets/workload-spiffe-credentials/credentialbundle.pem", true); - - CertificateBasedAccess cba = createCba(env, fs); - assertTrue(cba.useMtlsClientCertificate()); + CertificateBasedAccess cba = createCba(env); + // Explicit 'true' permits mTLS if certs exist, but if no certs are present, returns false/null cleanly (Row 3) + assertFalse(cba.useMtlsClientCertificate()); + assertNull(cba.getWorkloadCertPath()); } @Test - void testUseMtlsClientCertificateExplicitFalseWithSpiffeBundle() { + void testUseMtlsClientCertificateExplicitFalse() { TestEnv env = new TestEnv(); env.set("GOOGLE_API_USE_CLIENT_CERTIFICATE", "false"); - // Even if spiffe files are present, explicit false must override and disable mtls - TestFileSystem fs = new TestFileSystem(); - fs.setExists("/var/run/secrets/workload-spiffe-credentials/credentialbundle.pem", true); - - CertificateBasedAccess cba = createCba(env, fs); + CertificateBasedAccess cba = createCba(env); assertFalse(cba.useMtlsClientCertificate()); + assertNull(cba.getWorkloadCertPath()); } @Test void testUseMtlsClientCertificateUnsetNoFiles() { TestEnv env = new TestEnv(); - CertificateBasedAccess cba = createCba(env, new TestFileSystem()); + CertificateBasedAccess cba = createCba(env); assertFalse(cba.useMtlsClientCertificate()); + assertNull(cba.getWorkloadCertPath()); } @Test - void testUseMtlsClientCertificateUnsetSpiffeBundleExists() { - TestEnv env = new TestEnv(); - TestFileSystem fs = new TestFileSystem(); - fs.setExists("/var/run/secrets/workload-spiffe-credentials/credentialbundle.pem", true); - CertificateBasedAccess cba = createCba(env, fs); - assertTrue(cba.useMtlsClientCertificate()); - } - - @Test - void testUseMtlsClientCertificateUnsetSpiffeSeparateFilesExist() { - TestEnv env = new TestEnv(); - TestFileSystem fs = new TestFileSystem(); - fs.setExists("/var/run/secrets/workload-spiffe-credentials/certificates.pem", true); - fs.setExists("/var/run/secrets/workload-spiffe-credentials/private_key.pem", true); - CertificateBasedAccess cba = createCba(env, fs); - assertTrue(cba.useMtlsClientCertificate()); - } - - @Test - void testUseMtlsClientCertificateConfigValid() { - TestEnv env = new TestEnv(); - env.set("GOOGLE_API_CERTIFICATE_CONFIG", "/path/to/config.json"); - - TestFileSystem fs = new TestFileSystem(); - fs.setContent( - "/path/to/config.json", - "{\n \"cert_path\": \"/my/cert.pem\",\n \"key_path\": \"/my/key.pem\"\n}"); - fs.setExists("/my/cert.pem", true); - fs.setExists("/my/key.pem", true); - - CertificateBasedAccess cba = createCba(env, fs); - assertTrue(cba.useMtlsClientCertificate()); - } - - @Test - void testUseMtlsClientCertificateConfigMissingFile() { - TestEnv env = new TestEnv(); - env.set("GOOGLE_API_CERTIFICATE_CONFIG", "/path/to/config.json"); - - CertificateBasedAccess cba = createCba(env, new TestFileSystem()); - - IllegalStateException ex = - assertThrows(IllegalStateException.class, cba::useMtlsClientCertificate); - assertTrue(ex.getMessage().contains("configured but the file does not exist")); - } - - @Test - void testUseMtlsClientCertificateEnvTrueOverride() { + void testUseMtlsClientCertificateConfigMissingConfigFile_returnsNullSafely() { TestEnv env = new TestEnv(); - env.set("GOOGLE_API_USE_CLIENT_CERTIFICATE", "true"); - - CertificateBasedAccess cba = createCba(env, new TestFileSystem()); + env.set("GOOGLE_API_CERTIFICATE_CONFIG", "/nonexistent/config.json"); - assertTrue(cba.useMtlsClientCertificate()); - } - - @Test - void testUseMtlsClientCertificateConfigNonWorkloadJson() { - TestEnv env = new TestEnv(); - env.set("GOOGLE_API_CERTIFICATE_CONFIG", "/path/to/config.json"); - - TestFileSystem fs = new TestFileSystem(); - fs.setContent("/path/to/config.json", "{\n \"broken_path\": \"/my/cert.pem\"\n}"); - - CertificateBasedAccess cba = createCba(env, fs); + CertificateBasedAccess cba = createCba(env); + // Non-existent config file on disk returns false/null safely per Row 3 assertFalse(cba.useMtlsClientCertificate()); - } - - @Test - void testUseMtlsClientCertificateConfigMissingCertFiles() { - TestEnv env = new TestEnv(); - env.set("GOOGLE_API_CERTIFICATE_CONFIG", "/path/to/config.json"); - - TestFileSystem fs = new TestFileSystem(); - fs.setContent( - "/path/to/config.json", - "{\n \"cert_path\": \"/my/cert.pem\",\n \"key_path\": \"/my/key.pem\"\n}"); - // my/cert.pem and key.pem DO NOT exist - - CertificateBasedAccess cba = createCba(env, fs); - - IllegalStateException ex = - assertThrows(IllegalStateException.class, cba::useMtlsClientCertificate); - assertTrue( - ex.getCause() - .getMessage() - .contains("points to certificate/key files that do not exist on disk")); - } - - @Test - void testUseMtlsClientCertificateConfigWindowsPaths() { - TestEnv env = new TestEnv(); - env.set("GOOGLE_API_CERTIFICATE_CONFIG", "C:\\config.json"); - - TestFileSystem fs = new TestFileSystem(); - // In JSON, backslashes are escaped - fs.setContent( - "C:\\config.json", - "{\n" - + " \"cert_path\": \"C:\\\\my\\\\cert.pem\",\n" - + " \"key_path\": \"C:\\\\my\\\\key.pem\"\n" - + "}"); - fs.setExists("C:\\my\\cert.pem", true); - fs.setExists("C:\\my\\key.pem", true); - - CertificateBasedAccess cba = createCba(env, fs); - assertTrue(cba.useMtlsClientCertificate()); - } - - @Test - void testGetWorkloadCertPathWithMalformedConfigThrowsIllegalStateException() { - TestEnv env = new TestEnv(); - env.set("GOOGLE_API_USE_CLIENT_CERTIFICATE", "true"); - env.set("GOOGLE_API_CERTIFICATE_CONFIG", "/path/to/config.json"); - - TestFileSystem fs = new TestFileSystem(); - fs.setContent("/path/to/config.json", "{\n \"broken\": \"path\"\n}"); - - CertificateBasedAccess cba = createCba(env, fs); - - assertThrows(IllegalStateException.class, cba::getWorkloadCertPath); - } - - @Test - void testExtractJsonValueWithEscapedBackslashesAndQuotes() { - TestEnv env = new TestEnv(); - env.set("GOOGLE_API_CERTIFICATE_CONFIG", "/path/to/config.json"); - - TestFileSystem fs = new TestFileSystem(); - fs.setContent( - "/path/to/config.json", - "{\n" - + " \"cert_path\": \"/my/\\\"escaped\\\"/\\\\cert.pem\",\n" - + " \"key_path\": \"/my/key.pem\"\n" - + "}"); - fs.setExists("/my/\"escaped\"/\\cert.pem", true); - fs.setExists("/my/key.pem", true); - - CertificateBasedAccess cba = createCba(env, fs); - assertTrue(cba.useMtlsClientCertificate()); - assertEquals("/my/\"escaped\"/\\cert.pem", cba.getWorkloadCertPath()); + assertNull(cba.getWorkloadCertPath()); } } From be0a495c197a628ad6fcfb594fcb24257122b325 Mon Sep 17 00:00:00 2001 From: Matt Castelaz Date: Mon, 10 Aug 2026 19:29:37 +0000 Subject: [PATCH 08/29] fix(auth,gax): address PR 13995 review feedback and CI test failures - Separate GKE (credentialbundle.pem) and GCE (certificates.pem + private_key.pem) workload certificate fallback paths in MtlsUtils. - Restore full Javadoc on MtlsUtils.getWorkloadCertificateConfiguration. - Format MtlsUtils and MtlsUtilsTest with google-java-format. - Fix Java 8 Mockito reflection error in GrpcLoggingInterceptorTest by instantiating GrpcLoggingInterceptor directly. - Isolate DirectPath environment tests in InstantiatingGrpcChannelProviderTest from host environment variables. --- .../java/com/google/auth/mtls/MtlsUtils.java | 45 +++++++++++++------ .../com/google/auth/mtls/MtlsUtilsTest.java | 44 +++++++++++------- .../gax/grpc/GrpcLoggingInterceptorTest.java | 3 +- .../InstantiatingGrpcChannelProviderTest.java | 25 ++++++++--- .../gax/rpc/mtls/CertificateBasedAccess.java | 8 +++- .../rpc/mtls/CertificateBasedAccessTest.java | 6 +-- 6 files changed, 87 insertions(+), 44 deletions(-) diff --git a/google-auth-library-java/oauth2_http/java/com/google/auth/mtls/MtlsUtils.java b/google-auth-library-java/oauth2_http/java/com/google/auth/mtls/MtlsUtils.java index e4f8904cfcba..9cfa1f46d6eb 100644 --- a/google-auth-library-java/oauth2_http/java/com/google/auth/mtls/MtlsUtils.java +++ b/google-auth-library-java/oauth2_http/java/com/google/auth/mtls/MtlsUtils.java @@ -60,8 +60,8 @@ private MtlsUtils() { } /** - * Returns if mutual TLS client certificate should be used. - * Delegates directly to getWorkloadCertPath to avoid duplicate logic. + * Returns if mutual TLS client certificate should be used. Delegates directly to + * getWorkloadCertPath to avoid duplicate logic. */ public static boolean useMtlsClientCertificate( EnvironmentProvider envProvider, PropertyProvider propProvider) { @@ -69,7 +69,8 @@ public static boolean useMtlsClientCertificate( } /** - * Resolves and returns the path to the mutual TLS client certificate, or null if none should be used. + * Resolves and returns the path to the mutual TLS client certificate, or null if none should be + * used. */ public static @Nullable String getWorkloadCertPath( EnvironmentProvider envProvider, PropertyProvider propProvider) { @@ -112,7 +113,8 @@ public static boolean useMtlsClientCertificate( return config.getCertPath(); } } catch (CertificateSourceUnavailableException e) { - // Well-known gcloud certificate_config.json does not exist. Safe fallback to SPIFFE/well-known paths. + // Well-known gcloud certificate_config.json does not exist. Safe fallback to + // SPIFFE/well-known paths. } catch (Exception e) { // Ignore parsing errors for well-known config fallback } @@ -138,18 +140,17 @@ public static boolean useMtlsClientCertificate( if (bundleFile.exists()) { return bundleFile.getAbsolutePath(); } - - File certFile = new File(gkePath, "certificates.pem"); - File keyFile = new File(gkePath, "private_key.pem"); - if (certFile.exists() && keyFile.exists()) { - return certFile.getAbsolutePath(); - } return null; } /** Dedicated GCE Fallback Resolution Path */ public static @Nullable String getGceWorkloadCertPath() { - // Isolated GCE workload credentials fallback for independent rollout phase + String gcePath = "/var/run/secrets/workload-spiffe-credentials"; + File certFile = new File(gcePath, "certificates.pem"); + File keyFile = new File(gcePath, "private_key.pem"); + if (certFile.exists() && keyFile.exists()) { + return certFile.getAbsolutePath(); + } return null; } @@ -189,23 +190,39 @@ public static boolean useMtlsClientCertificate( * @throws IOException if the certificate configuration cannot be found or loaded. */ public static String getCertificatePath( - EnvironmentProvider envProvider, PropertyProvider propProvider, @Nullable String certConfigPathOverride) + EnvironmentProvider envProvider, + PropertyProvider propProvider, + @Nullable String certConfigPathOverride) throws IOException { String certPath = getWorkloadCertificateConfiguration(envProvider, propProvider, certConfigPathOverride) .getCertPath(); if (Strings.isNullOrEmpty(certPath)) { throw new CertificateSourceUnavailableException( - "Certificate configuration loaded successfully, but does not contain a 'certificate_file' path."); + "Certificate configuration loaded successfully, but does not contain a 'certificate_file'" + + " path."); } return certPath; } /** * Resolves and loads the workload certificate configuration. + * + *

The configuration file is resolved in the following order of precedence: 1. The provided + * certConfigPathOverride (if not null). 2. The path specified by the + * GOOGLE_API_CERTIFICATE_CONFIG environment variable. 3. The well-known certificate configuration + * file in the gcloud config directory. + * + * @param envProvider the environment provider to use for resolving environment variables + * @param propProvider the property provider to use for resolving system properties + * @param certConfigPathOverride optional override path for the configuration file + * @return the loaded WorkloadCertificateConfiguration + * @throws IOException if the configuration file cannot be found, read, or parsed */ static WorkloadCertificateConfiguration getWorkloadCertificateConfiguration( - EnvironmentProvider envProvider, PropertyProvider propProvider, @Nullable String certConfigPathOverride) + EnvironmentProvider envProvider, + PropertyProvider propProvider, + @Nullable String certConfigPathOverride) throws IOException { File certConfig; if (certConfigPathOverride != null) { diff --git a/google-auth-library-java/oauth2_http/javatests/com/google/auth/mtls/MtlsUtilsTest.java b/google-auth-library-java/oauth2_http/javatests/com/google/auth/mtls/MtlsUtilsTest.java index b1b2c7853192..2e4f4c94c15d 100644 --- a/google-auth-library-java/oauth2_http/javatests/com/google/auth/mtls/MtlsUtilsTest.java +++ b/google-auth-library-java/oauth2_http/javatests/com/google/auth/mtls/MtlsUtilsTest.java @@ -246,7 +246,8 @@ public String getProperty(String name, String def) { @Test void useMtlsClientCertificate_trueWithNoCertsOnDisk_returnsFalseWithoutThrowing() { - EnvironmentProvider envProvider = name -> "GOOGLE_API_USE_CLIENT_CERTIFICATE".equals(name) ? "true" : null; + EnvironmentProvider envProvider = + name -> "GOOGLE_API_USE_CLIENT_CERTIFICATE".equals(name) ? "true" : null; PropertyProvider propProvider = (name, def) -> def; assertFalse(MtlsUtils.useMtlsClientCertificate(envProvider, propProvider)); @@ -255,7 +256,8 @@ void useMtlsClientCertificate_trueWithNoCertsOnDisk_returnsFalseWithoutThrowing( @Test void useMtlsClientCertificate_false_returnsFalse() { - EnvironmentProvider envProvider = name -> "GOOGLE_API_USE_CLIENT_CERTIFICATE".equals(name) ? "false" : null; + EnvironmentProvider envProvider = + name -> "GOOGLE_API_USE_CLIENT_CERTIFICATE".equals(name) ? "false" : null; PropertyProvider propProvider = (name, def) -> def; assertFalse(MtlsUtils.useMtlsClientCertificate(envProvider, propProvider)); @@ -263,26 +265,25 @@ void useMtlsClientCertificate_false_returnsFalse() { } @Test - void getWorkloadCertPath_brokenConfigPath_throwsIllegalStateException() { - EnvironmentProvider envProvider = name -> "GOOGLE_API_CERTIFICATE_CONFIG".equals(name) ? "/nonexistent/config.json" : null; + void getWorkloadCertPath_missingConfigFile_returnsNullSafely() { + EnvironmentProvider envProvider = + name -> "GOOGLE_API_CERTIFICATE_CONFIG".equals(name) ? "/nonexistent/config.json" : null; PropertyProvider propProvider = (name, def) -> def; - IllegalStateException exception = - assertThrows( - IllegalStateException.class, - () -> MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); - assertTrue(exception.getMessage().contains("Certificate config is configured but file does not exist")); + assertNull(MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); } @Test - void getWorkloadCertPath_configPointsToMissingCertFiles_throwsIllegalStateException() throws IOException { + void getWorkloadCertPath_configPointsToMissingCertFiles_throwsIllegalStateException() + throws IOException { Path configFile = tempDir.resolve("config.json"); Files.write( configFile, "{\"cert_configs\":{\"workload\":{\"cert_path\":\"/nonexistent/cert.pem\",\"key_path\":\"/nonexistent/key.pem\"}}}" .getBytes()); - EnvironmentProvider envProvider = name -> "GOOGLE_API_CERTIFICATE_CONFIG".equals(name) ? configFile.toString() : null; + EnvironmentProvider envProvider = + name -> "GOOGLE_API_CERTIFICATE_CONFIG".equals(name) ? configFile.toString() : null; PropertyProvider propProvider = (name, def) -> def; IllegalStateException exception = @@ -300,13 +301,14 @@ void getWorkloadCertPath_validConfig_returnsCertPath() throws IOException { Files.write(keyFile, "dummy key".getBytes()); Path configFile = tempDir.resolve("config.json"); - String configJson = String.format( - "{\"cert_configs\":{\"workload\":{\"cert_path\":\"%s\",\"key_path\":\"%s\"}}}", - certFile.toString().replace("\\", "\\\\"), - keyFile.toString().replace("\\", "\\\\")); + String configJson = + String.format( + "{\"cert_configs\":{\"workload\":{\"cert_path\":\"%s\",\"key_path\":\"%s\"}}}", + certFile.toString().replace("\\", "\\\\"), keyFile.toString().replace("\\", "\\\\")); Files.write(configFile, configJson.getBytes()); - EnvironmentProvider envProvider = name -> "GOOGLE_API_CERTIFICATE_CONFIG".equals(name) ? configFile.toString() : null; + EnvironmentProvider envProvider = + name -> "GOOGLE_API_CERTIFICATE_CONFIG".equals(name) ? configFile.toString() : null; PropertyProvider propProvider = (name, def) -> def; assertTrue(MtlsUtils.useMtlsClientCertificate(envProvider, propProvider)); @@ -322,4 +324,14 @@ void getCertificateFingerprint_validFile_returnsSha256() throws IOException { assertNotNull(fingerprint); assertEquals(64, fingerprint.length()); // SHA-256 hex string length } + + @Test + void getGkeWorkloadCertPath_nonexistent_returnsNull() { + assertNull(MtlsUtils.getGkeWorkloadCertPath()); + } + + @Test + void getGceWorkloadCertPath_nonexistent_returnsNull() { + assertNull(MtlsUtils.getGceWorkloadCertPath()); + } } diff --git a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/GrpcLoggingInterceptorTest.java b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/GrpcLoggingInterceptorTest.java index fad4cd468b95..c93db599d575 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/GrpcLoggingInterceptorTest.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/GrpcLoggingInterceptorTest.java @@ -32,7 +32,6 @@ import static org.mockito.ArgumentMatchers.any; import static org.mockito.Mockito.mock; -import static org.mockito.Mockito.spy; import static org.mockito.Mockito.verify; import static org.mockito.Mockito.when; @@ -83,7 +82,7 @@ void testInterceptor_basic() { void testInterceptor_responseListener() { when(channel.newCall(Mockito.>any(), any(CallOptions.class))) .thenReturn(call); - GrpcLoggingInterceptor interceptor = spy(new GrpcLoggingInterceptor()); + GrpcLoggingInterceptor interceptor = new GrpcLoggingInterceptor(); Channel intercepted = ClientInterceptors.intercept(channel, interceptor); @SuppressWarnings("unchecked") ClientCall.Listener listener = mock(ClientCall.Listener.class); diff --git a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/InstantiatingGrpcChannelProviderTest.java b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/InstantiatingGrpcChannelProviderTest.java index be0365866615..c2127aede3b8 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/InstantiatingGrpcChannelProviderTest.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/InstantiatingGrpcChannelProviderTest.java @@ -664,7 +664,9 @@ private void createAndCloseTransportChannel(InstantiatingGrpcChannelProvider pro createAndCloseTransportChannel(provider); assertThat(logHandler.getAllMessages()) .contains( - "DirectPath is misconfigured. The DirectPath XDS option was set, but the attemptDirectPath option was not. Please set both the attemptDirectPath and attemptDirectPathXds options."); + "DirectPath is misconfigured. The DirectPath XDS option was set, but the" + + " attemptDirectPath option was not. Please set both the attemptDirectPath and" + + " attemptDirectPathXds options."); InstantiatingGrpcChannelProvider.LOG.removeHandler(logHandler); } @@ -682,8 +684,10 @@ void testLogDirectPathMisconfig_AttemptDirectPathNotSetAndAttemptDirectPathXdsSe createAndCloseTransportChannel(provider); assertThat(logHandler.getAllMessages()) .contains( - "Env var GOOGLE_CLOUD_ENABLE_DIRECT_PATH_XDS was found and set to TRUE, but DirectPath was not enabled for this client. If this is intended for " - + "this client, please note that this is a misconfiguration and set the attemptDirectPath option as well."); + "Env var GOOGLE_CLOUD_ENABLE_DIRECT_PATH_XDS was found and set to TRUE, but DirectPath" + + " was not enabled for this client. If this is intended for this client, please" + + " note that this is a misconfiguration and set the attemptDirectPath option as" + + " well."); InstantiatingGrpcChannelProvider.LOG.removeHandler(logHandler); } @@ -711,6 +715,7 @@ void testLogDirectPathMisconfigWrongCredential() throws Exception { InstantiatingGrpcChannelProvider.newBuilder() .setAttemptDirectPathXds() .setAttemptDirectPath(true) + .setEnvProvider(name -> null) .setHeaderProvider( mock(HeaderProvider.class, Mockito.withSettings().withoutAnnotations())) .setExecutor(mock(Executor.class)) @@ -877,12 +882,14 @@ public void canUseDirectPath_directPathEnvVarDisabled() throws IOException { @Test public void canUseDirectPath_directPathEnvVarNotSet_attemptDirectPathIsTrue() { System.setProperty("os.name", "Linux"); + EnvironmentProvider envProvider = name -> null; InstantiatingGrpcChannelProvider.Builder builder = InstantiatingGrpcChannelProvider.newBuilder() .setCertificateBasedAccess(certificateBasedAccess) .setAttemptDirectPath(true) .setCredentials(computeEngineCredentials) - .setEndpoint(DEFAULT_ENDPOINT); + .setEndpoint(DEFAULT_ENDPOINT) + .setEnvProvider(envProvider); InstantiatingGrpcChannelProvider provider = new InstantiatingGrpcChannelProvider(builder, GCE_PRODUCTION_NAME_AFTER_2016); Truth.assertThat(provider.canUseDirectPath()).isTrue(); @@ -891,12 +898,14 @@ public void canUseDirectPath_directPathEnvVarNotSet_attemptDirectPathIsTrue() { @Test public void canUseDirectPath_directPathEnvVarNotSet_attemptDirectPathIsFalse() { System.setProperty("os.name", "Linux"); + EnvironmentProvider envProvider = name -> null; InstantiatingGrpcChannelProvider.Builder builder = InstantiatingGrpcChannelProvider.newBuilder() .setCertificateBasedAccess(certificateBasedAccess) .setAttemptDirectPath(false) .setCredentials(computeEngineCredentials) - .setEndpoint(DEFAULT_ENDPOINT); + .setEndpoint(DEFAULT_ENDPOINT) + .setEnvProvider(envProvider); InstantiatingGrpcChannelProvider provider = new InstantiatingGrpcChannelProvider(builder, GCE_PRODUCTION_NAME_AFTER_2016); Truth.assertThat(provider.canUseDirectPath()).isFalse(); @@ -1201,7 +1210,8 @@ void createS2ASecuredChannelCredentials_bothS2AAddressesNull_returnsNull() { assertThat(provider.createS2ASecuredChannelCredentials()).isNotNull(); assertThat(logHandler.getAllMessages()) .contains( - "Cannot establish an mTLS connection to S2A because autoconfig endpoint did not return a mtls address to reach S2A."); + "Cannot establish an mTLS connection to S2A because autoconfig endpoint did not return" + + " a mtls address to reach S2A."); InstantiatingGrpcChannelProvider.LOG.removeHandler(logHandler); } @@ -1247,7 +1257,8 @@ void createS2ASecuredChannelCredentials_returnsPlaintextToS2AS2AChannelCredentia assertThat(provider.createS2ASecuredChannelCredentials()).isNotNull(); assertThat(logHandler.getAllMessages()) .contains( - "Cannot establish an mTLS connection to S2A because MTLS to MDS credentials do not exist on filesystem, falling back to plaintext connection to S2A"); + "Cannot establish an mTLS connection to S2A because MTLS to MDS credentials do not" + + " exist on filesystem, falling back to plaintext connection to S2A"); InstantiatingGrpcChannelProvider.LOG.removeHandler(logHandler); } diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/CertificateBasedAccess.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/CertificateBasedAccess.java index c3e1fced9774..0dbdf60723d2 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/CertificateBasedAccess.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/CertificateBasedAccess.java @@ -58,7 +58,13 @@ public interface FileContentReader { } public CertificateBasedAccess(EnvironmentProvider envProvider) { - this(envProvider, path -> new java.io.File(path).isFile(), path -> new String(java.nio.file.Files.readAllBytes(java.nio.file.Paths.get(path)), java.nio.charset.StandardCharsets.UTF_8)); + this( + envProvider, + path -> new java.io.File(path).isFile(), + path -> + new String( + java.nio.file.Files.readAllBytes(java.nio.file.Paths.get(path)), + java.nio.charset.StandardCharsets.UTF_8)); } CertificateBasedAccess( diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/CertificateBasedAccessTest.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/CertificateBasedAccessTest.java index ee01c89a3695..a6ba55fba6a5 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/CertificateBasedAccessTest.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/CertificateBasedAccessTest.java @@ -33,10 +33,7 @@ import static org.junit.jupiter.api.Assertions.assertEquals; import static org.junit.jupiter.api.Assertions.assertFalse; import static org.junit.jupiter.api.Assertions.assertNull; -import static org.junit.jupiter.api.Assertions.assertThrows; -import static org.junit.jupiter.api.Assertions.assertTrue; -import java.io.IOException; import java.util.HashMap; import java.util.Map; import org.junit.jupiter.api.Test; @@ -104,7 +101,8 @@ void testUseMtlsClientCertificateExplicitTrueNoCredentials() { TestEnv env = new TestEnv(); env.set("GOOGLE_API_USE_CLIENT_CERTIFICATE", "true"); CertificateBasedAccess cba = createCba(env); - // Explicit 'true' permits mTLS if certs exist, but if no certs are present, returns false/null cleanly (Row 3) + // Explicit 'true' permits mTLS if certs exist, but if no certs are present, returns false/null + // cleanly (Row 3) assertFalse(cba.useMtlsClientCertificate()); assertNull(cba.getWorkloadCertPath()); } From 9be88f63765ad314a3938fb46c9c553939ecfb26 Mon Sep 17 00:00:00 2001 From: Matt Castelaz Date: Tue, 18 Aug 2026 13:46:41 +0000 Subject: [PATCH 09/29] fix(auth,gax): align mTLS certificate discovery and error handling with go/sdk-mtls-by-default-cert-discovery Address PR 13995 review feedback from @nbayati: - Align discovery and error behavior with go/sdk-mtls-by-default-cert-discovery: - Fail closed (IllegalStateException) when GOOGLE_API_CERTIFICATE_CONFIG points to a missing, unreadable, malformed, or missing cert/key configuration. - Safe fallback (return null) when implicit default gcloud config is missing or is an ECP-only configuration without a workload block. - Fail closed with clear source identification if default gcloud config is unreadable, malformed, or points to missing cert/key files. - Replace .exists() with .isFile() && .canRead() checks across config, certificate, and key paths. - Make getGkeWorkloadCertPath and getGceWorkloadCertPath package-private stubs returning null with explanatory comments for phased rollout. - Explicitly identify the resolution source (GOOGLE_API_CERTIFICATE_CONFIG vs default gcloud location) in all error messages. - Update getCertificatePath exception message to reference 'cert_configs.workload.cert_path' rather than legacy 'certificate_file'. - Add comprehensive test coverage in MtlsUtilsTest and CertificateBasedAccessTest. --- .../java/com/google/auth/mtls/MtlsUtils.java | 131 ++++--- .../com/google/auth/mtls/MtlsUtilsTest.java | 345 +++++++++++++++++- .../rpc/mtls/CertificateBasedAccessTest.java | 15 +- 3 files changed, 435 insertions(+), 56 deletions(-) diff --git a/google-auth-library-java/oauth2_http/java/com/google/auth/mtls/MtlsUtils.java b/google-auth-library-java/oauth2_http/java/com/google/auth/mtls/MtlsUtils.java index 9cfa1f46d6eb..0f2518d51972 100644 --- a/google-auth-library-java/oauth2_http/java/com/google/auth/mtls/MtlsUtils.java +++ b/google-auth-library-java/oauth2_http/java/com/google/auth/mtls/MtlsUtils.java @@ -79,47 +79,78 @@ public static boolean useMtlsClientCertificate( return null; } - String certConfigPath = envProvider.getEnv(CERTIFICATE_CONFIGURATION_ENV_VARIABLE); - if (!Strings.isNullOrEmpty(certConfigPath)) { + String explicitConfigPath = envProvider.getEnv(CERTIFICATE_CONFIGURATION_ENV_VARIABLE); + + // 1. Explicit Configuration Path (Fail Closed) + if (!Strings.isNullOrEmpty(explicitConfigPath)) { + File configFile = new File(explicitConfigPath); + if (!configFile.exists()) { + throw new IllegalStateException( + "Certificate configuration file specified via GOOGLE_API_CERTIFICATE_CONFIG at '" + + explicitConfigPath + + "' does not exist."); + } + if (!configFile.isFile() || !configFile.canRead()) { + throw new IllegalStateException( + "Failed to read certificate configuration file specified via" + + " GOOGLE_API_CERTIFICATE_CONFIG at '" + + explicitConfigPath + + "'."); + } try { WorkloadCertificateConfiguration config = - getWorkloadCertificateConfiguration(envProvider, propProvider, certConfigPath); - - File certFile = new File(config.getCertPath()); - File keyFile = new File(config.getPrivateKeyPath()); - if (!certFile.exists() || !keyFile.exists()) { - throw new IllegalStateException( - "Certificate config points to certificate/key files that do not exist on disk: " - + "cert_path=" - + config.getCertPath() - + ", key_path=" - + config.getPrivateKeyPath()); - } + getWorkloadCertificateConfiguration(envProvider, propProvider, explicitConfigPath); + validateCertAndKeyFiles(config, explicitConfigPath, false); return config.getCertPath(); } catch (CertificateSourceUnavailableException e) { - // Certificate config file does not exist on disk -> safe fallback + // ECP / PKCS11 configuration without workload section; safe fallback + return null; } catch (IllegalStateException e) { throw e; } catch (Exception e) { - throw new IllegalStateException("Failed to parse certificate config: " + certConfigPath, e); + throw new IllegalStateException( + "Certificate configuration file specified via GOOGLE_API_CERTIFICATE_CONFIG at '" + + explicitConfigPath + + "' is malformed: " + + e.getMessage(), + e); + } + } + + // 2. Implicit / Default gcloud Configuration Path + File defaultConfigFile = null; + try { + defaultConfigFile = getWellKnownCertificateConfigFile(envProvider, propProvider); + } catch (IOException e) { + // APPDATA missing on Windows, etc. Safe fallback. + } + if (defaultConfigFile != null && defaultConfigFile.exists()) { + if (!defaultConfigFile.isFile() || !defaultConfigFile.canRead()) { + throw new IllegalStateException( + "Default certificate configuration file at '" + + defaultConfigFile.getAbsolutePath() + + "' exists but could not be read."); } - } else { try { WorkloadCertificateConfiguration config = getWorkloadCertificateConfiguration(envProvider, propProvider, null); - File certFile = new File(config.getCertPath()); - File keyFile = new File(config.getPrivateKeyPath()); - if (certFile.exists() && keyFile.exists()) { - return config.getCertPath(); - } + validateCertAndKeyFiles(config, defaultConfigFile.getAbsolutePath(), true); + return config.getCertPath(); } catch (CertificateSourceUnavailableException e) { - // Well-known gcloud certificate_config.json does not exist. Safe fallback to - // SPIFFE/well-known paths. + // ECP-only configuration without workload section; safe fallback + } catch (IllegalStateException e) { + throw e; } catch (Exception e) { - // Ignore parsing errors for well-known config fallback + throw new IllegalStateException( + "Default certificate configuration file at '" + + defaultConfigFile.getAbsolutePath() + + "' is malformed: " + + e.getMessage(), + e); } } + // 3. Platform SPIFFE Fallbacks (Stubs) String gkeCertPath = getGkeWorkloadCertPath(); if (gkeCertPath != null) { return gkeCertPath; @@ -133,24 +164,40 @@ public static boolean useMtlsClientCertificate( return null; } - /** Dedicated GKE Fallback Resolution Path */ - public static @Nullable String getGkeWorkloadCertPath() { - String gkePath = "/var/run/secrets/workload-spiffe-credentials"; - File bundleFile = new File(gkePath, "credentialbundle.pem"); - if (bundleFile.exists()) { - return bundleFile.getAbsolutePath(); + private static void validateCertAndKeyFiles( + WorkloadCertificateConfiguration config, String configPath, boolean isDefaultConfig) { + File certFile = new File(config.getCertPath()); + File keyFile = new File(config.getPrivateKeyPath()); + if (!certFile.isFile() || !certFile.canRead() || !keyFile.isFile() || !keyFile.canRead()) { + String sourcePrefix = + isDefaultConfig + ? "referenced by default configuration '" + : "referenced by configuration '"; + throw new IllegalStateException( + "Failed to read certificate/key file at '" + + config.getCertPath() + + "' or '" + + config.getPrivateKeyPath() + + "' " + + sourcePrefix + + configPath + + "'."); } + } + + /** Dedicated GKE Fallback Resolution Path */ + static @Nullable String getGkeWorkloadCertPath() { + // GKE workload certificate resolution is temporarily disabled (returns null) + // pending Phase 1 rollout of bound token support on GKE + // (go/agentic-bound-token-sdk-rollout-plan). return null; } /** Dedicated GCE Fallback Resolution Path */ - public static @Nullable String getGceWorkloadCertPath() { - String gcePath = "/var/run/secrets/workload-spiffe-credentials"; - File certFile = new File(gcePath, "certificates.pem"); - File keyFile = new File(gcePath, "private_key.pem"); - if (certFile.exists() && keyFile.exists()) { - return certFile.getAbsolutePath(); - } + static @Nullable String getGceWorkloadCertPath() { + // GCE workload certificate resolution is temporarily disabled (returns null) + // pending Phase 2 rollout of bound token support on GCE + // (go/agentic-bound-token-sdk-rollout-plan). return null; } @@ -160,7 +207,7 @@ public static boolean useMtlsClientCertificate( return null; } File file = new File(certPath); - if (!file.exists()) { + if (!file.isFile() || !file.canRead()) { return null; } try { @@ -199,8 +246,8 @@ public static String getCertificatePath( .getCertPath(); if (Strings.isNullOrEmpty(certPath)) { throw new CertificateSourceUnavailableException( - "Certificate configuration loaded successfully, but does not contain a 'certificate_file'" - + " path."); + "Certificate configuration loaded successfully, but does not contain a" + + " 'cert_configs.workload.cert_path' path."); } return certPath; } @@ -236,7 +283,7 @@ static WorkloadCertificateConfiguration getWorkloadCertificateConfiguration( } } - if (!certConfig.isFile()) { + if (!certConfig.isFile() || !certConfig.canRead()) { throw new CertificateSourceUnavailableException( "Certificate configuration file does not exist or is not a file: " + certConfig.getAbsolutePath()); diff --git a/google-auth-library-java/oauth2_http/javatests/com/google/auth/mtls/MtlsUtilsTest.java b/google-auth-library-java/oauth2_http/javatests/com/google/auth/mtls/MtlsUtilsTest.java index 2e4f4c94c15d..32f6c9741711 100644 --- a/google-auth-library-java/oauth2_http/javatests/com/google/auth/mtls/MtlsUtilsTest.java +++ b/google-auth-library-java/oauth2_http/javatests/com/google/auth/mtls/MtlsUtilsTest.java @@ -264,23 +264,177 @@ void useMtlsClientCertificate_false_returnsFalse() { assertNull(MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); } + // --- Explicit GOOGLE_API_CERTIFICATE_CONFIG Tests (Fail Closed) --- + @Test - void getWorkloadCertPath_missingConfigFile_returnsNullSafely() { + void getWorkloadCertPath_explicitConfigMissing_throwsIllegalStateException() { EnvironmentProvider envProvider = name -> "GOOGLE_API_CERTIFICATE_CONFIG".equals(name) ? "/nonexistent/config.json" : null; PropertyProvider propProvider = (name, def) -> def; + IllegalStateException exception = + assertThrows( + IllegalStateException.class, + () -> MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); + assertTrue( + exception + .getMessage() + .contains( + "specified via GOOGLE_API_CERTIFICATE_CONFIG at '/nonexistent/config.json' does not" + + " exist")); + } + + @Test + void getWorkloadCertPath_explicitConfigIsDirectory_throwsIllegalStateException() + throws IOException { + Path configDir = tempDir.resolve("config_dir"); + Files.createDirectory(configDir); + + EnvironmentProvider envProvider = + name -> "GOOGLE_API_CERTIFICATE_CONFIG".equals(name) ? configDir.toString() : null; + PropertyProvider propProvider = (name, def) -> def; + + IllegalStateException exception = + assertThrows( + IllegalStateException.class, + () -> MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); + assertTrue( + exception + .getMessage() + .contains( + "Failed to read certificate configuration file specified via" + + " GOOGLE_API_CERTIFICATE_CONFIG")); + } + + @Test + void getWorkloadCertPath_explicitConfigUnreadable_throwsIllegalStateException() + throws IOException { + Path configFile = tempDir.resolve("unreadable_config.json"); + Files.write(configFile, "{}".getBytes()); + File file = configFile.toFile(); + if (file.setReadable(false)) { + try { + EnvironmentProvider envProvider = + name -> "GOOGLE_API_CERTIFICATE_CONFIG".equals(name) ? configFile.toString() : null; + PropertyProvider propProvider = (name, def) -> def; + + IllegalStateException exception = + assertThrows( + IllegalStateException.class, + () -> MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); + assertTrue( + exception + .getMessage() + .contains( + "Failed to read certificate configuration file specified via" + + " GOOGLE_API_CERTIFICATE_CONFIG")); + } finally { + file.setReadable(true); + } + } + } + + @Test + void getWorkloadCertPath_explicitConfigMalformedJson_throwsIllegalStateException() + throws IOException { + Path configFile = tempDir.resolve("malformed.json"); + Files.write(configFile, "{ invalid json".getBytes()); + + EnvironmentProvider envProvider = + name -> "GOOGLE_API_CERTIFICATE_CONFIG".equals(name) ? configFile.toString() : null; + PropertyProvider propProvider = (name, def) -> def; + + IllegalStateException exception = + assertThrows( + IllegalStateException.class, + () -> MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); + assertTrue( + exception + .getMessage() + .contains( + "specified via GOOGLE_API_CERTIFICATE_CONFIG at '" + + configFile.toString() + + "' is malformed")); + } + + @Test + void getWorkloadCertPath_explicitConfigOnlyEcp_returnsNullSafely() throws IOException { + Path configFile = tempDir.resolve("ecp_config.json"); + Files.write(configFile, "{\"cert_configs\":{\"enterprise_certificates\":{}}}".getBytes()); + + EnvironmentProvider envProvider = + name -> "GOOGLE_API_CERTIFICATE_CONFIG".equals(name) ? configFile.toString() : null; + PropertyProvider propProvider = (name, def) -> def; + assertNull(MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); } @Test - void getWorkloadCertPath_configPointsToMissingCertFiles_throwsIllegalStateException() + void getWorkloadCertPath_explicitConfigCertFileMissing_throwsIllegalStateException() + throws IOException { + Path keyFile = tempDir.resolve("key.pem"); + Files.write(keyFile, "dummy key".getBytes()); + + Path configFile = tempDir.resolve("config.json"); + String configJson = + String.format( + "{\"cert_configs\":{\"workload\":{\"cert_path\":\"/nonexistent/cert.pem\",\"key_path\":\"%s\"}}}", + keyFile.toString().replace("\\", "\\\\")); + Files.write(configFile, configJson.getBytes()); + + EnvironmentProvider envProvider = + name -> "GOOGLE_API_CERTIFICATE_CONFIG".equals(name) ? configFile.toString() : null; + PropertyProvider propProvider = (name, def) -> def; + + IllegalStateException exception = + assertThrows( + IllegalStateException.class, + () -> MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); + assertTrue(exception.getMessage().contains("Failed to read certificate/key file")); + assertTrue( + exception + .getMessage() + .contains("referenced by configuration '" + configFile.toString() + "'")); + } + + @Test + void getWorkloadCertPath_explicitConfigCertFileIsDirectory_throwsIllegalStateException() throws IOException { + Path certDir = tempDir.resolve("cert_dir"); + Files.createDirectory(certDir); + Path keyFile = tempDir.resolve("key.pem"); + Files.write(keyFile, "dummy key".getBytes()); + Path configFile = tempDir.resolve("config.json"); - Files.write( - configFile, - "{\"cert_configs\":{\"workload\":{\"cert_path\":\"/nonexistent/cert.pem\",\"key_path\":\"/nonexistent/key.pem\"}}}" - .getBytes()); + String configJson = + String.format( + "{\"cert_configs\":{\"workload\":{\"cert_path\":\"%s\",\"key_path\":\"%s\"}}}", + certDir.toString().replace("\\", "\\\\"), keyFile.toString().replace("\\", "\\\\")); + Files.write(configFile, configJson.getBytes()); + + EnvironmentProvider envProvider = + name -> "GOOGLE_API_CERTIFICATE_CONFIG".equals(name) ? configFile.toString() : null; + PropertyProvider propProvider = (name, def) -> def; + + IllegalStateException exception = + assertThrows( + IllegalStateException.class, + () -> MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); + assertTrue(exception.getMessage().contains("Failed to read certificate/key file")); + } + + @Test + void getWorkloadCertPath_explicitConfigKeyFileMissing_throwsIllegalStateException() + throws IOException { + Path certFile = tempDir.resolve("cert.pem"); + Files.write(certFile, "dummy cert".getBytes()); + + Path configFile = tempDir.resolve("config.json"); + String configJson = + String.format( + "{\"cert_configs\":{\"workload\":{\"cert_path\":\"%s\",\"key_path\":\"/nonexistent/key.pem\"}}}", + certFile.toString().replace("\\", "\\\\")); + Files.write(configFile, configJson.getBytes()); EnvironmentProvider envProvider = name -> "GOOGLE_API_CERTIFICATE_CONFIG".equals(name) ? configFile.toString() : null; @@ -290,11 +444,15 @@ void getWorkloadCertPath_configPointsToMissingCertFiles_throwsIllegalStateExcept assertThrows( IllegalStateException.class, () -> MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); - assertTrue(exception.getMessage().contains("files that do not exist on disk")); + assertTrue(exception.getMessage().contains("Failed to read certificate/key file")); + assertTrue( + exception + .getMessage() + .contains("referenced by configuration '" + configFile.toString() + "'")); } @Test - void getWorkloadCertPath_validConfig_returnsCertPath() throws IOException { + void getWorkloadCertPath_explicitConfigValid_returnsCertPath() throws IOException { Path certFile = tempDir.resolve("cert.pem"); Path keyFile = tempDir.resolve("key.pem"); Files.write(certFile, "dummy cert".getBytes()); @@ -315,6 +473,166 @@ void getWorkloadCertPath_validConfig_returnsCertPath() throws IOException { assertEquals(certFile.toString(), MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); } + // --- Implicit / Default gcloud Config Tests --- + + @Test + void getWorkloadCertPath_defaultConfigMissing_returnsNullSafely() { + EnvironmentProvider envProvider = name -> null; + PropertyProvider propProvider = + (name, def) -> { + if ("user.home".equals(name)) return tempDir.toString(); + if ("os.name".equals(name)) return "Linux"; + return def; + }; + + assertNull(MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); + } + + @Test + void getWorkloadCertPath_defaultConfigIsDirectory_throwsIllegalStateException() + throws IOException { + Path gcloudDir = tempDir.resolve(".config/gcloud"); + Files.createDirectories(gcloudDir); + Path defaultConfigFile = gcloudDir.resolve("certificate_config.json"); + Files.createDirectory(defaultConfigFile); + + EnvironmentProvider envProvider = name -> null; + PropertyProvider propProvider = + (name, def) -> { + if ("user.home".equals(name)) return tempDir.toString(); + if ("os.name".equals(name)) return "Linux"; + return def; + }; + + IllegalStateException exception = + assertThrows( + IllegalStateException.class, + () -> MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); + assertTrue( + exception + .getMessage() + .contains( + "Default certificate configuration file at '" + + defaultConfigFile.toFile().getAbsolutePath() + + "' exists but could not be read")); + } + + @Test + void getWorkloadCertPath_defaultConfigMalformedJson_throwsIllegalStateException() + throws IOException { + Path gcloudDir = tempDir.resolve(".config/gcloud"); + Files.createDirectories(gcloudDir); + Path defaultConfigFile = gcloudDir.resolve("certificate_config.json"); + Files.write(defaultConfigFile, "{ malformed json".getBytes()); + + EnvironmentProvider envProvider = name -> null; + PropertyProvider propProvider = + (name, def) -> { + if ("user.home".equals(name)) return tempDir.toString(); + if ("os.name".equals(name)) return "Linux"; + return def; + }; + + IllegalStateException exception = + assertThrows( + IllegalStateException.class, + () -> MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); + assertTrue( + exception + .getMessage() + .contains( + "Default certificate configuration file at '" + + defaultConfigFile.toFile().getAbsolutePath() + + "' is malformed")); + } + + @Test + void getWorkloadCertPath_defaultConfigOnlyEcp_returnsNullSafely() throws IOException { + Path gcloudDir = tempDir.resolve(".config/gcloud"); + Files.createDirectories(gcloudDir); + Path defaultConfigFile = gcloudDir.resolve("certificate_config.json"); + Files.write( + defaultConfigFile, + "{\"cert_configs\":{\"enterprise_certificates\":{\"libs\":[]}}}".getBytes()); + + EnvironmentProvider envProvider = name -> null; + PropertyProvider propProvider = + (name, def) -> { + if ("user.home".equals(name)) return tempDir.toString(); + if ("os.name".equals(name)) return "Linux"; + return def; + }; + + assertNull(MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); + } + + @Test + void getWorkloadCertPath_defaultConfigCertFileMissing_throwsIllegalStateException() + throws IOException { + Path keyFile = tempDir.resolve("key.pem"); + Files.write(keyFile, "dummy key".getBytes()); + + Path gcloudDir = tempDir.resolve(".config/gcloud"); + Files.createDirectories(gcloudDir); + Path defaultConfigFile = gcloudDir.resolve("certificate_config.json"); + String configJson = + String.format( + "{\"cert_configs\":{\"workload\":{\"cert_path\":\"/nonexistent/cert.pem\",\"key_path\":\"%s\"}}}", + keyFile.toString().replace("\\", "\\\\")); + Files.write(defaultConfigFile, configJson.getBytes()); + + EnvironmentProvider envProvider = name -> null; + PropertyProvider propProvider = + (name, def) -> { + if ("user.home".equals(name)) return tempDir.toString(); + if ("os.name".equals(name)) return "Linux"; + return def; + }; + + IllegalStateException exception = + assertThrows( + IllegalStateException.class, + () -> MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); + assertTrue(exception.getMessage().contains("Failed to read certificate/key file")); + assertTrue( + exception + .getMessage() + .contains( + "referenced by default configuration '" + + defaultConfigFile.toFile().getAbsolutePath() + + "'")); + } + + @Test + void getWorkloadCertPath_defaultConfigValid_returnsCertPath() throws IOException { + Path certFile = tempDir.resolve("cert.pem"); + Path keyFile = tempDir.resolve("key.pem"); + Files.write(certFile, "dummy cert".getBytes()); + Files.write(keyFile, "dummy key".getBytes()); + + Path gcloudDir = tempDir.resolve(".config/gcloud"); + Files.createDirectories(gcloudDir); + Path defaultConfigFile = gcloudDir.resolve("certificate_config.json"); + String configJson = + String.format( + "{\"cert_configs\":{\"workload\":{\"cert_path\":\"%s\",\"key_path\":\"%s\"}}}", + certFile.toString().replace("\\", "\\\\"), keyFile.toString().replace("\\", "\\\\")); + Files.write(defaultConfigFile, configJson.getBytes()); + + EnvironmentProvider envProvider = name -> null; + PropertyProvider propProvider = + (name, def) -> { + if ("user.home".equals(name)) return tempDir.toString(); + if ("os.name".equals(name)) return "Linux"; + return def; + }; + + assertTrue(MtlsUtils.useMtlsClientCertificate(envProvider, propProvider)); + assertEquals(certFile.toString(), MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); + } + + // --- General Helpers & Stubs Tests --- + @Test void getCertificateFingerprint_validFile_returnsSha256() throws IOException { Path file = tempDir.resolve("test.crt"); @@ -326,12 +644,19 @@ void getCertificateFingerprint_validFile_returnsSha256() throws IOException { } @Test - void getGkeWorkloadCertPath_nonexistent_returnsNull() { + void getCertificateFingerprint_invalidOrNull_returnsNull() { + assertNull(MtlsUtils.getCertificateFingerprint(null)); + assertNull(MtlsUtils.getCertificateFingerprint("/nonexistent/file.crt")); + assertNull(MtlsUtils.getCertificateFingerprint(tempDir.toString())); // Directory + } + + @Test + void getGkeWorkloadCertPath_returnsNull() { assertNull(MtlsUtils.getGkeWorkloadCertPath()); } @Test - void getGceWorkloadCertPath_nonexistent_returnsNull() { + void getGceWorkloadCertPath_returnsNull() { assertNull(MtlsUtils.getGceWorkloadCertPath()); } } diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/CertificateBasedAccessTest.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/CertificateBasedAccessTest.java index a6ba55fba6a5..0e7055232298 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/CertificateBasedAccessTest.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/CertificateBasedAccessTest.java @@ -33,6 +33,7 @@ import static org.junit.jupiter.api.Assertions.assertEquals; import static org.junit.jupiter.api.Assertions.assertFalse; import static org.junit.jupiter.api.Assertions.assertNull; +import static org.junit.jupiter.api.Assertions.assertThrows; import java.util.HashMap; import java.util.Map; @@ -43,6 +44,11 @@ class CertificateBasedAccessTest { private static class TestEnv { private final Map env = new HashMap<>(); + TestEnv() { + // Hermetically isolate tests from the host's ~/.config/gcloud/certificate_config.json + env.put("CLOUDSDK_CONFIG", "/nonexistent/test/gcloud"); + } + void set(String key, String val) { env.put(key, val); } @@ -126,14 +132,15 @@ void testUseMtlsClientCertificateUnsetNoFiles() { } @Test - void testUseMtlsClientCertificateConfigMissingConfigFile_returnsNullSafely() { + void testUseMtlsClientCertificateConfigMissingConfigFile_throwsIllegalStateException() { TestEnv env = new TestEnv(); env.set("GOOGLE_API_CERTIFICATE_CONFIG", "/nonexistent/config.json"); CertificateBasedAccess cba = createCba(env); - // Non-existent config file on disk returns false/null safely per Row 3 - assertFalse(cba.useMtlsClientCertificate()); - assertNull(cba.getWorkloadCertPath()); + // Non-existent config file on disk specified via explicit env var throws IllegalStateException + // (Fail Closed) + assertThrows(IllegalStateException.class, () -> cba.useMtlsClientCertificate()); + assertThrows(IllegalStateException.class, () -> cba.getWorkloadCertPath()); } } From 765eb3b60204f463ea97ee10e7df40a35bb28fc3 Mon Sep 17 00:00:00 2001 From: Matt Castelaz Date: Tue, 18 Aug 2026 14:20:16 +0000 Subject: [PATCH 10/29] test(auth): add unit test in MtlsUtilsTest to provide coverage for ECP flow in getCertificatePath --- .../com/google/auth/mtls/MtlsUtilsTest.java | 15 +++++++++++++++ 1 file changed, 15 insertions(+) diff --git a/google-auth-library-java/oauth2_http/javatests/com/google/auth/mtls/MtlsUtilsTest.java b/google-auth-library-java/oauth2_http/javatests/com/google/auth/mtls/MtlsUtilsTest.java index 32f6c9741711..3de35dfeed12 100644 --- a/google-auth-library-java/oauth2_http/javatests/com/google/auth/mtls/MtlsUtilsTest.java +++ b/google-auth-library-java/oauth2_http/javatests/com/google/auth/mtls/MtlsUtilsTest.java @@ -101,6 +101,21 @@ public String getProperty(String name, String def) { () -> MtlsUtils.getCertificatePath(envProvider, propProvider, configFile.toString())); } + @Test + void getCertificatePath_ecpOnlyConfig_throwsCertificateSourceUnavailableException() + throws IOException { + Path configFile = tempDir.resolve("ecp_config.json"); + Files.write( + configFile, "{\"cert_configs\":{\"enterprise_certificates\":{\"libs\":[]}}}".getBytes()); + + EnvironmentProvider envProvider = name -> null; + PropertyProvider propProvider = (name, def) -> def; + + assertThrows( + CertificateSourceUnavailableException.class, + () -> MtlsUtils.getCertificatePath(envProvider, propProvider, configFile.toString())); + } + @Test void getWorkloadCertificateConfiguration_overridePath() throws IOException { Path configFile = tempDir.resolve("custom_config.json"); From 4904aadf3d37586598ea14c1e45588602a59f593 Mon Sep 17 00:00:00 2001 From: Matt Castelaz Date: Thu, 27 Aug 2026 15:47:11 +0000 Subject: [PATCH 11/29] fix(auth,gax-grpc): address PR 13995 review feedback on cert discovery and channel refresh - Rename MtlsUtils.validateCertAndKeyFiles to checkCertAndKeyFilesReadable. - Move file readability check outside try-catch in MtlsUtils to clearly separate parsing errors from file existence errors. - Remove GKE/GCE placeholder stubs and internal doc references from MtlsUtils. - Simplify MtlsUtils.getCertificateFingerprint using Files.readAllBytes and Guava BaseEncoding. - Defer activeCertFingerprint mutation in ChannelPool until after channel creation succeeds in refreshAll(). - Add unit test in ChannelPoolTest verifying failed refresh attempts do not mutate fingerprint or prevent subsequent retries. --- .../java/com/google/auth/mtls/MtlsUtils.java | 69 +++++-------------- .../com/google/auth/mtls/MtlsUtilsTest.java | 20 +++--- .../com/google/api/gax/grpc/ChannelPool.java | 17 +++-- .../google/api/gax/grpc/ChannelPoolTest.java | 59 ++++++++++++++++ 4 files changed, 97 insertions(+), 68 deletions(-) diff --git a/google-auth-library-java/oauth2_http/java/com/google/auth/mtls/MtlsUtils.java b/google-auth-library-java/oauth2_http/java/com/google/auth/mtls/MtlsUtils.java index 0f2518d51972..f5c5fd694836 100644 --- a/google-auth-library-java/oauth2_http/java/com/google/auth/mtls/MtlsUtils.java +++ b/google-auth-library-java/oauth2_http/java/com/google/auth/mtls/MtlsUtils.java @@ -34,10 +34,12 @@ import com.google.auth.oauth2.EnvironmentProvider; import com.google.auth.oauth2.PropertyProvider; import com.google.common.base.Strings; +import com.google.common.io.BaseEncoding; import java.io.File; import java.io.FileInputStream; import java.io.IOException; import java.io.InputStream; +import java.nio.file.Files; import java.security.MessageDigest; import java.util.Locale; import org.jspecify.annotations.NullMarked; @@ -97,16 +99,12 @@ public static boolean useMtlsClientCertificate( + explicitConfigPath + "'."); } + WorkloadCertificateConfiguration config; try { - WorkloadCertificateConfiguration config = - getWorkloadCertificateConfiguration(envProvider, propProvider, explicitConfigPath); - validateCertAndKeyFiles(config, explicitConfigPath, false); - return config.getCertPath(); + config = getWorkloadCertificateConfiguration(envProvider, propProvider, explicitConfigPath); } catch (CertificateSourceUnavailableException e) { // ECP / PKCS11 configuration without workload section; safe fallback return null; - } catch (IllegalStateException e) { - throw e; } catch (Exception e) { throw new IllegalStateException( "Certificate configuration file specified via GOOGLE_API_CERTIFICATE_CONFIG at '" @@ -115,6 +113,8 @@ public static boolean useMtlsClientCertificate( + e.getMessage(), e); } + checkCertAndKeyFilesReadable(config, explicitConfigPath, false); + return config.getCertPath(); } // 2. Implicit / Default gcloud Configuration Path @@ -131,15 +131,11 @@ public static boolean useMtlsClientCertificate( + defaultConfigFile.getAbsolutePath() + "' exists but could not be read."); } + WorkloadCertificateConfiguration config = null; try { - WorkloadCertificateConfiguration config = - getWorkloadCertificateConfiguration(envProvider, propProvider, null); - validateCertAndKeyFiles(config, defaultConfigFile.getAbsolutePath(), true); - return config.getCertPath(); + config = getWorkloadCertificateConfiguration(envProvider, propProvider, null); } catch (CertificateSourceUnavailableException e) { // ECP-only configuration without workload section; safe fallback - } catch (IllegalStateException e) { - throw e; } catch (Exception e) { throw new IllegalStateException( "Default certificate configuration file at '" @@ -148,23 +144,16 @@ public static boolean useMtlsClientCertificate( + e.getMessage(), e); } - } - - // 3. Platform SPIFFE Fallbacks (Stubs) - String gkeCertPath = getGkeWorkloadCertPath(); - if (gkeCertPath != null) { - return gkeCertPath; - } - - String gceCertPath = getGceWorkloadCertPath(); - if (gceCertPath != null) { - return gceCertPath; + if (config != null) { + checkCertAndKeyFilesReadable(config, defaultConfigFile.getAbsolutePath(), true); + return config.getCertPath(); + } } return null; } - private static void validateCertAndKeyFiles( + private static void checkCertAndKeyFilesReadable( WorkloadCertificateConfiguration config, String configPath, boolean isDefaultConfig) { File certFile = new File(config.getCertPath()); File keyFile = new File(config.getPrivateKeyPath()); @@ -185,22 +174,6 @@ private static void validateCertAndKeyFiles( } } - /** Dedicated GKE Fallback Resolution Path */ - static @Nullable String getGkeWorkloadCertPath() { - // GKE workload certificate resolution is temporarily disabled (returns null) - // pending Phase 1 rollout of bound token support on GKE - // (go/agentic-bound-token-sdk-rollout-plan). - return null; - } - - /** Dedicated GCE Fallback Resolution Path */ - static @Nullable String getGceWorkloadCertPath() { - // GCE workload certificate resolution is temporarily disabled (returns null) - // pending Phase 2 rollout of bound token support on GCE - // (go/agentic-bound-token-sdk-rollout-plan). - return null; - } - /** Centralized SHA-256 Fingerprint Calculator */ public static @Nullable String getCertificateFingerprint(@Nullable String certPath) { if (certPath == null) { @@ -211,19 +184,9 @@ private static void validateCertAndKeyFiles( return null; } try { - MessageDigest digest = MessageDigest.getInstance("SHA-256"); - try (FileInputStream fis = new FileInputStream(file)) { - byte[] byteArray = new byte[1024]; - int bytesCount; - while ((bytesCount = fis.read(byteArray)) != -1) { - digest.update(byteArray, 0, bytesCount); - } - } - StringBuilder sb = new StringBuilder(); - for (byte b : digest.digest()) { - sb.append(String.format("%02x", b)); - } - return sb.toString(); + byte[] certBytes = Files.readAllBytes(file.toPath()); + byte[] digest = MessageDigest.getInstance("SHA-256").digest(certBytes); + return BaseEncoding.base16().lowerCase().encode(digest); } catch (Exception e) { return null; } diff --git a/google-auth-library-java/oauth2_http/javatests/com/google/auth/mtls/MtlsUtilsTest.java b/google-auth-library-java/oauth2_http/javatests/com/google/auth/mtls/MtlsUtilsTest.java index 3de35dfeed12..3c1a06614705 100644 --- a/google-auth-library-java/oauth2_http/javatests/com/google/auth/mtls/MtlsUtilsTest.java +++ b/google-auth-library-java/oauth2_http/javatests/com/google/auth/mtls/MtlsUtilsTest.java @@ -656,22 +656,22 @@ void getCertificateFingerprint_validFile_returnsSha256() throws IOException { String fingerprint = MtlsUtils.getCertificateFingerprint(file.toString()); assertNotNull(fingerprint); assertEquals(64, fingerprint.length()); // SHA-256 hex string length + assertEquals("b94d27b9934d3e08a52e52d7da7dabfac484efe37a5380ee9088f7ace2efcde9", fingerprint); } @Test - void getCertificateFingerprint_invalidOrNull_returnsNull() { - assertNull(MtlsUtils.getCertificateFingerprint(null)); - assertNull(MtlsUtils.getCertificateFingerprint("/nonexistent/file.crt")); - assertNull(MtlsUtils.getCertificateFingerprint(tempDir.toString())); // Directory - } + void getCertificateFingerprint_emptyFile_returnsValidSha256() throws IOException { + Path emptyFile = tempDir.resolve("empty.crt"); + Files.write(emptyFile, new byte[0]); - @Test - void getGkeWorkloadCertPath_returnsNull() { - assertNull(MtlsUtils.getGkeWorkloadCertPath()); + String fingerprint = MtlsUtils.getCertificateFingerprint(emptyFile.toString()); + assertEquals("e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855", fingerprint); } @Test - void getGceWorkloadCertPath_returnsNull() { - assertNull(MtlsUtils.getGceWorkloadCertPath()); + void getCertificateFingerprint_invalidOrNull_returnsNull() { + assertNull(MtlsUtils.getCertificateFingerprint(null)); + assertNull(MtlsUtils.getCertificateFingerprint("/nonexistent/file.crt")); + assertNull(MtlsUtils.getCertificateFingerprint(tempDir.toString())); // Directory } } diff --git a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java index f0e0db3d3017..d1198058f31d 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java @@ -449,13 +449,12 @@ private void expand(int desiredSize) { private void refreshSafely() { try { synchronized (entryWriteLock) { - if (workloadCertPath != null) { + if (refreshAll() && workloadCertPath != null) { String currentDiskFingerprint = getOrUpdateDiskFingerprint(workloadCertPath); if (!currentDiskFingerprint.isEmpty()) { this.activeCertFingerprint = currentDiskFingerprint; } } - refreshAll(); } } catch (Exception e) { LOG.log(Level.WARNING, "Failed to pre-emptively refresh channels", e); @@ -534,13 +533,14 @@ void refresh() { return; } - this.activeCertFingerprint = currentDiskFingerprint; - refreshAll(); + if (refreshAll()) { + this.activeCertFingerprint = currentDiskFingerprint; + } } } @InternalApi("Visible for testing") - void refreshAll() { + boolean refreshAll() { synchronized (entryWriteLock) { LOG.fine( "Refreshing all channels" @@ -548,15 +548,21 @@ void refreshAll() { ? "" : " with certificate fingerprint: " + activeCertFingerprint)); ArrayList newEntries = new ArrayList<>(entries.get()); + boolean anyCreated = false; for (int i = 0; i < newEntries.size(); i++) { try { newEntries.set(i, new Entry(channelFactory.createSingleChannel())); + anyCreated = true; } catch (IOException e) { LOG.log(Level.WARNING, "Failed to refresh channel, leaving old channel", e); } } + if (!anyCreated && !newEntries.isEmpty()) { + return false; + } + ImmutableList replacedEntries = entries.getAndSet(ImmutableList.copyOf(newEntries)); // Shutdown the channels that were cycled out. @@ -565,6 +571,7 @@ void refreshAll() { e.requestShutdown(); } } + return true; } } diff --git a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java index 6cb7be99c9eb..496451aa2dfd 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java @@ -506,6 +506,65 @@ void channelReactiveMTlsRefreshShouldConditionallySwapChannels() .newCall(Mockito.>any(), Mockito.any(CallOptions.class)); } + @Test + void channelReactiveMTlsRefresh_failedCreationDoesNotMutateFingerprintAndAllowsRetry() + throws IOException { + ManagedChannel channel1 = Mockito.mock(ManagedChannel.class); + ManagedChannel channel2 = Mockito.mock(ManagedChannel.class); + ChannelFactory channelFactory = Mockito.mock(ChannelFactory.class); + + // Initial creation returns channel1, refresh attempt 1 throws IOException, refresh attempt 2 + // returns channel2 + Mockito.when(channelFactory.createSingleChannel()) + .thenReturn(channel1) + .thenThrow(new IOException("Transient channel creation error")) + .thenReturn(channel2); + + tempCert = java.nio.file.Files.createTempFile("cert", ".pem"); + java.nio.file.Path clientCert = + java.nio.file.Paths.get("src", "test", "resources", "client_cert.pem"); + java.nio.file.Files.copy( + clientCert, tempCert, java.nio.file.StandardCopyOption.REPLACE_EXISTING); + + ChannelPoolSettings channelPoolSettings = + ChannelPoolSettings.builder().setInitialChannelCount(1).build(); + + pool = ChannelPool.create(channelPoolSettings, channelFactory, null, tempCert.toString()); + + // Initially uses channel1 + pool.newCall(FakeMethodDescriptor.create(), CallOptions.DEFAULT); + Mockito.verify(channel1, Mockito.times(1)) + .newCall(Mockito.>any(), Mockito.any(CallOptions.class)); + + // Rotate cert on disk + pool.invalidateDiskFingerprintCache(); + java.nio.file.Path rootCert = + java.nio.file.Paths.get("src", "test", "resources", "root_cert.pem"); + java.nio.file.Files.copy(rootCert, tempCert, java.nio.file.StandardCopyOption.REPLACE_EXISTING); + + // Refresh attempt 1: createSingleChannel throws IOException. + // Refresh should fail to replace channel and MUST NOT record the new cert fingerprint as + // active. + pool.refresh(); + + // Verify still channel1 + pool.newCall(FakeMethodDescriptor.create(), CallOptions.DEFAULT); + Mockito.verify(channel1, Mockito.times(2)) + .newCall(Mockito.>any(), Mockito.any(CallOptions.class)); + + // Refresh attempt 2: with the same cert file on disk (cache expired), channelFactory now + // succeeds. + // If the fingerprint had been mutated on the failed attempt, this call would be skipped as a + // duplicate! + pool.invalidateDiskFingerprintCache(); + pool.refresh(); + + // Verify it has now swapped to channel2! + pool.newCall(FakeMethodDescriptor.create(), CallOptions.DEFAULT); + Mockito.verify(channel2, Mockito.times(1)) + .newCall(Mockito.>any(), Mockito.any(CallOptions.class)); + } + @Test void channelRefreshShouldSwapChannels() throws IOException { ManagedChannel underlyingChannel1 = mock(ManagedChannel.class); From 10535d62a652b1fc3e4ae0c1fe8abd2f32dc4deb Mon Sep 17 00:00:00 2001 From: Matt Castelaz Date: Fri, 28 Aug 2026 19:35:28 +0000 Subject: [PATCH 12/29] fix(gax,gax-grpc,gax-httpjson): address PR 13995 review feedback on rotation retries - Remove unused FileExistenceProvider/FileContentReader and 3-arg constructor from CertificateBasedAccess. - In ServerStreamingAttemptCallable, BidiStreamingCallable, and ClientStreamingCallable, wrap transportChannel.refresh() in try-catch with warning logging and propagate original exception without marking isRetryable=true. - Add getGeneration() to TransportChannel, ChannelPool, and RefreshingHttpJsonChannel; update AttemptCallable to track attemptGeneration so sibling in-flight requests that failed on the stale connection are retried without redundant channel recreation. - Guard ChannelPool.refresh() and refreshAll() against invocation on shut-down pool and synchronize isShutdown state across shutdown methods. - Add delegating protected constructor in ManagedHttpJsonChannel so RefreshingHttpJsonChannel and ManagedHttpJsonInterceptorChannel do not leak unused parent scheduled executors and default HTTP transports. - Only wrap HTTP/JSON channels with RefreshingHttpJsonChannel when workloadCertPath is not null. - Configure Conscrypt security provider prior to calling NetHttpTransport.Builder.trustCertificates in InstantiatingHttpJsonChannelProvider. - Clear stale transportChannel reference in HttpJsonCallContext.withChannel() and merge() when channel changes. - Add comprehensive unit tests across gax, gax-grpc, and gax-httpjson modules. --- .../com/google/api/gax/grpc/ChannelPool.java | 71 ++++++--- .../api/gax/grpc/GrpcTransportChannel.java | 9 ++ .../google/api/gax/grpc/ChannelPoolTest.java | 35 +++++ .../api/gax/httpjson/HttpJsonCallContext.java | 5 +- .../httpjson/HttpJsonTransportChannel.java | 5 + .../InstantiatingHttpJsonChannelProvider.java | 7 +- .../gax/httpjson/ManagedHttpJsonChannel.java | 13 ++ .../ManagedHttpJsonInterceptorChannel.java | 7 +- .../httpjson/RefreshingHttpJsonChannel.java | 9 ++ .../gax/httpjson/HttpJsonCallContextTest.java | 43 ++++++ ...tantiatingHttpJsonChannelProviderTest.java | 50 ++++++- .../RefreshingHttpJsonChannelTest.java | 21 +++ .../google/api/gax/rpc/AttemptCallable.java | 52 +++++-- .../api/gax/rpc/BidiStreamingCallable.java | 23 ++- .../api/gax/rpc/ClientStreamingCallable.java | 23 ++- .../rpc/ServerStreamingAttemptCallable.java | 26 ++-- .../google/api/gax/rpc/TransportChannel.java | 8 + .../gax/rpc/mtls/CertificateBasedAccess.java | 24 --- .../api/gax/rpc/AttemptCallableTest.java | 138 ++++++++++++++++++ .../ServerStreamingAttemptCallableTest.java | 38 ++++- .../api/gax/rpc/testing/FakeChannel.java | 11 ++ .../gax/rpc/testing/FakeTransportChannel.java | 15 ++ 22 files changed, 523 insertions(+), 110 deletions(-) diff --git a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java index d1198058f31d..8056ad734bbb 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java @@ -35,6 +35,7 @@ import com.google.common.annotations.VisibleForTesting; import com.google.common.base.Preconditions; import com.google.common.collect.ImmutableList; +import com.google.errorprone.annotations.concurrent.GuardedBy; import io.grpc.CallOptions; import io.grpc.Channel; import io.grpc.ClientCall; @@ -54,6 +55,7 @@ import java.util.concurrent.TimeUnit; import java.util.concurrent.atomic.AtomicBoolean; import java.util.concurrent.atomic.AtomicInteger; +import java.util.concurrent.atomic.AtomicLong; import java.util.concurrent.atomic.AtomicReference; import java.util.logging.Level; import java.util.logging.Logger; @@ -102,6 +104,11 @@ private static class DiskCheckResult { private final java.util.concurrent.locks.ReentrantLock diskCheckLock = new java.util.concurrent.locks.ReentrantLock(); private final Object entryWriteLock = new Object(); + + @GuardedBy("entryWriteLock") + private boolean isShutdown = false; + + private final AtomicLong generation = new AtomicLong(0); private volatile String activeCertFingerprint = ""; @VisibleForTesting final AtomicReference> entries = new AtomicReference<>(); private final AtomicInteger indexTicker = new AtomicInteger(); @@ -212,19 +219,22 @@ Channel getChannel(int affinity) { public ManagedChannel shutdown() { LOG.fine("Initiating graceful shutdown due to explicit request"); - // Resize and refresh tasks can block on channel priming. We don't need - // to wait for the channels to be ready since we're shutting down the - // pool. Allowing interrupt to speed it up. - if (resizeFuture != null) { - resizeFuture.cancel(true); - } - if (refreshFuture != null) { - refreshFuture.cancel(true); - } + synchronized (entryWriteLock) { + isShutdown = true; + // Resize and refresh tasks can block on channel priming. We don't need + // to wait for the channels to be ready since we're shutting down the + // pool. Allowing interrupt to speed it up. + if (resizeFuture != null) { + resizeFuture.cancel(true); + } + if (refreshFuture != null) { + refreshFuture.cancel(true); + } - List localEntries = entries.get(); - for (Entry entry : localEntries) { - entry.channel.shutdown(); + List localEntries = entries.get(); + for (Entry entry : localEntries) { + entry.channel.shutdown(); + } } if (backgroundExecutorProvider.shouldAutoClose()) { @@ -237,6 +247,11 @@ public ManagedChannel shutdown() { /** {@inheritDoc} */ @Override public boolean isShutdown() { + synchronized (entryWriteLock) { + if (isShutdown) { + return true; + } + } List localEntries = entries.get(); for (Entry entry : localEntries) { if (!entry.channel.isShutdown()) { @@ -263,16 +278,19 @@ public boolean isTerminated() { public ManagedChannel shutdownNow() { LOG.fine("Initiating immediate shutdown due to explicit request"); - if (resizeFuture != null) { - resizeFuture.cancel(true); - } - if (refreshFuture != null) { - refreshFuture.cancel(true); - } + synchronized (entryWriteLock) { + isShutdown = true; + if (resizeFuture != null) { + resizeFuture.cancel(true); + } + if (refreshFuture != null) { + refreshFuture.cancel(true); + } - List localEntries = entries.get(); - for (Entry entry : localEntries) { - entry.channel.shutdownNow(); + List localEntries = entries.get(); + for (Entry entry : localEntries) { + entry.channel.shutdownNow(); + } } if (backgroundExecutorProvider.shouldAutoClose()) { @@ -516,6 +534,9 @@ void refresh() { // - then thread2 will shut down channel that thread1 will put back into circulation (after it // replaces the list) synchronized (entryWriteLock) { + if (isShutdown) { + return; + } if (workloadCertPath == null) { refreshAll(); return; @@ -542,6 +563,9 @@ void refresh() { @InternalApi("Visible for testing") boolean refreshAll() { synchronized (entryWriteLock) { + if (isShutdown) { + return false; + } LOG.fine( "Refreshing all channels" + (activeCertFingerprint == null @@ -571,10 +595,15 @@ boolean refreshAll() { e.requestShutdown(); } } + generation.incrementAndGet(); return true; } } + public long getGeneration() { + return generation.get(); + } + /** * Get and retain a Channel Entry. The returned Entry will have its rpc count incremented, * preventing it from getting recycled. diff --git a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/GrpcTransportChannel.java b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/GrpcTransportChannel.java index 31ede726f3f3..63180a8149f6 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/GrpcTransportChannel.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/GrpcTransportChannel.java @@ -85,6 +85,15 @@ public boolean shouldRefresh() { return false; } + @Override + public long getGeneration() { + Channel channel = getChannel(); + if (channel instanceof ChannelPool) { + return ((ChannelPool) channel).getGeneration(); + } + return 0; + } + @Override public void shutdown() { getManagedChannel().shutdown(); diff --git a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java index 496451aa2dfd..97f27f77513f 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java @@ -565,6 +565,41 @@ void channelReactiveMTlsRefresh_failedCreationDoesNotMutateFingerprintAndAllowsR .newCall(Mockito.>any(), Mockito.any(CallOptions.class)); } + @Test + void refresh_onShutdownPool_noOpsAndCreatesNoChannels() throws IOException { + ManagedChannel channel1 = mock(ManagedChannel.class); + ManagedChannel channel2 = mock(ManagedChannel.class); + ChannelFactory channelFactory = Mockito.mock(ChannelFactory.class); + Mockito.when(channelFactory.createSingleChannel()).thenReturn(channel1, channel2); + + pool = ChannelPool.create(ChannelPoolSettings.staticallySized(1), channelFactory, null, null); + Mockito.verify(channelFactory, Mockito.times(1)).createSingleChannel(); + + pool.shutdown(); + assertThat(pool.isShutdown()).isTrue(); + + // Invoking refresh or refreshAll on shut down pool must no-op and never create new subchannels + pool.refresh(); + boolean refreshed = pool.refreshAll(); + assertThat(refreshed).isFalse(); + Mockito.verify(channelFactory, Mockito.times(1)).createSingleChannel(); + assertThat(pool.isShutdown()).isTrue(); + } + + @Test + void generationCounterIncrementsOnRefresh() throws IOException { + ManagedChannel channel1 = mock(ManagedChannel.class); + ManagedChannel channel2 = mock(ManagedChannel.class); + ChannelFactory channelFactory = Mockito.mock(ChannelFactory.class); + Mockito.when(channelFactory.createSingleChannel()).thenReturn(channel1, channel2); + + pool = ChannelPool.create(ChannelPoolSettings.staticallySized(1), channelFactory, null, null); + assertThat(pool.getGeneration()).isEqualTo(0); + + pool.refreshAll(); + assertThat(pool.getGeneration()).isEqualTo(1); + } + @Test void channelRefreshShouldSwapChannels() throws IOException { ManagedChannel underlyingChannel1 = mock(ManagedChannel.class); diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/HttpJsonCallContext.java b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/HttpJsonCallContext.java index 5a2c739ac345..47e3d31ee6e7 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/HttpJsonCallContext.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/HttpJsonCallContext.java @@ -226,7 +226,8 @@ public HttpJsonCallContext merge(ApiCallContext inputCallContext) { } TransportChannel newTransportChannel = httpJsonCallContext.transportChannel; - if (newTransportChannel == null) { + if (newTransportChannel == null + && (httpJsonCallContext.channel == null || httpJsonCallContext.channel.equals(channel))) { newTransportChannel = this.transportChannel; } @@ -600,7 +601,7 @@ public HttpJsonCallContext withChannel(@Nullable HttpJsonChannel newChannel) { this.retrySettings, this.retryableCodes, this.endpointContext, - this.transportChannel); + (newChannel == null || newChannel.equals(this.channel)) ? this.transportChannel : null); } public HttpJsonCallContext withCallOptions(HttpJsonCallOptions newCallOptions) { diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/HttpJsonTransportChannel.java b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/HttpJsonTransportChannel.java index a8333a589a4a..fde0673f650d 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/HttpJsonTransportChannel.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/HttpJsonTransportChannel.java @@ -74,6 +74,11 @@ public boolean shouldRefresh() { return getManagedChannel().shouldRefresh(); } + @Override + public long getGeneration() { + return getManagedChannel().getGeneration(); + } + @Override public void shutdown() { getManagedChannel().shutdown(); diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProvider.java b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProvider.java index 90ce27c2879d..01a1b9b6b625 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProvider.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProvider.java @@ -200,8 +200,8 @@ public TransportChannelProvider withCredentials(Credentials credentials) { KeyStore mtlsKeyStore = mtlsProvider.getKeyStore(); if (mtlsKeyStore != null) { NetHttpTransport.Builder builder = new NetHttpTransport.Builder(); - builder.trustCertificates(null, mtlsKeyStore, ""); HttpJsonConscryptUtils.configureConscryptSecurityProvider(builder); + builder.trustCertificates(null, mtlsKeyStore, ""); return builder.build(); } } @@ -228,8 +228,11 @@ private HttpJsonTransportChannel createChannel() throws IOException, GeneralSecu } }; + String workloadCertPath = certificateBasedAccess.getWorkloadCertPath(); ManagedHttpJsonChannel channel = - new RefreshingHttpJsonChannel(channelFactory, certificateBasedAccess.getWorkloadCertPath()); + workloadCertPath != null + ? new RefreshingHttpJsonChannel(channelFactory, workloadCertPath) + : channelFactory.get(); HttpJsonClientInterceptor headerInterceptor = new HttpJsonHeaderInterceptor(headerProvider.getHeaders()); diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/ManagedHttpJsonChannel.java b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/ManagedHttpJsonChannel.java index f83f09bac486..86863e4ce51c 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/ManagedHttpJsonChannel.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/ManagedHttpJsonChannel.java @@ -60,6 +60,19 @@ protected ManagedHttpJsonChannel() { this(null, true, null, null, true); } + protected ManagedHttpJsonChannel(boolean isDelegatingWrapper) { + this.executor = null; + this.usingDefaultExecutor = false; + this.endpoint = null; + this.httpTransport = null; + this.usingDefaultTransport = false; + this.deadlineScheduledExecutorService = null; + } + + public long getGeneration() { + return 0; + } + String getEndpoint() { return endpoint; } diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/ManagedHttpJsonInterceptorChannel.java b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/ManagedHttpJsonInterceptorChannel.java index e552608b6529..9ccbbc4ec1a2 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/ManagedHttpJsonInterceptorChannel.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/ManagedHttpJsonInterceptorChannel.java @@ -43,11 +43,16 @@ class ManagedHttpJsonInterceptorChannel extends ManagedHttpJsonChannel { ManagedHttpJsonInterceptorChannel( ManagedHttpJsonChannel channel, HttpJsonClientInterceptor interceptor) { - super(); + super(true); this.channel = channel; this.interceptor = interceptor; } + @Override + public long getGeneration() { + return channel.getGeneration(); + } + @VisibleForTesting ManagedHttpJsonChannel getChannel() { return channel; diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java index c0855b16871b..2fec3ba02253 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java @@ -75,10 +75,13 @@ private static class DiskCheckResult { private final java.util.concurrent.ConcurrentLinkedQueue allEntries = new java.util.concurrent.ConcurrentLinkedQueue<>(); private final Object refreshLock = new Object(); + private final java.util.concurrent.atomic.AtomicLong generation = + new java.util.concurrent.atomic.AtomicLong(0); private volatile String activeCertFingerprint = ""; public RefreshingHttpJsonChannel( Supplier channelFactory, String workloadCertPath) { + super(true); this.channelFactory = channelFactory; this.workloadCertPath = workloadCertPath; ChannelEntry initial = new ChannelEntry(channelFactory.get()); @@ -167,6 +170,7 @@ public void refresh() { allEntries.add(newEntry); ChannelEntry oldEntry = activeEntry.getAndSet(newEntry); this.activeCertFingerprint = currentDiskFingerprint; + generation.incrementAndGet(); if (oldEntry != null) { oldEntry.requestShutdown(); @@ -174,6 +178,11 @@ public void refresh() { } } + @Override + public long getGeneration() { + return generation.get(); + } + private ChannelEntry getRetainedEntry() { while (true) { ChannelEntry entry = activeEntry.get(); diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/HttpJsonCallContextTest.java b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/HttpJsonCallContextTest.java index 08044522e729..15b3df3b86ce 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/HttpJsonCallContextTest.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/HttpJsonCallContextTest.java @@ -334,4 +334,47 @@ void testMergeOptions() { assertEquals(testContext2, mergedContext.getOption(contextKey2)); assertEquals(testContext3, mergedContext.getOption(contextKey3)); } + + @Test + void testWithChannelClearsStaleTransportChannel() { + ManagedHttpJsonChannel channel1 = + mock(ManagedHttpJsonChannel.class, Mockito.withSettings().withoutAnnotations()); + ManagedHttpJsonChannel channel2 = + mock(ManagedHttpJsonChannel.class, Mockito.withSettings().withoutAnnotations()); + + HttpJsonTransportChannel transportChannel1 = + HttpJsonTransportChannel.newBuilder().setManagedChannel(channel1).build(); + + HttpJsonCallContext context = + HttpJsonCallContext.createDefault().withTransportChannel(transportChannel1); + Truth.assertThat(context.getTransportChannel()).isSameInstanceAs(transportChannel1); + + // Retains transportChannel when setting same channel or null + Truth.assertThat(context.withChannel(channel1).getTransportChannel()) + .isSameInstanceAs(transportChannel1); + Truth.assertThat(context.withChannel(null).getTransportChannel()) + .isSameInstanceAs(transportChannel1); + + // Clears transportChannel to null when setting a different channel + Truth.assertThat(context.withChannel(channel2).getTransportChannel()).isNull(); + } + + @Test + void testMergeClearsStaleTransportChannel() { + ManagedHttpJsonChannel channel1 = + mock(ManagedHttpJsonChannel.class, Mockito.withSettings().withoutAnnotations()); + ManagedHttpJsonChannel channel2 = + mock(ManagedHttpJsonChannel.class, Mockito.withSettings().withoutAnnotations()); + + HttpJsonTransportChannel transportChannel1 = + HttpJsonTransportChannel.newBuilder().setManagedChannel(channel1).build(); + + HttpJsonCallContext context1 = + HttpJsonCallContext.createDefault().withTransportChannel(transportChannel1); + HttpJsonCallContext context2 = HttpJsonCallContext.createDefault().withChannel(channel2); + + HttpJsonCallContext merged = context1.merge(context2); + Truth.assertThat(merged.getChannel()).isSameInstanceAs(channel2); + Truth.assertThat(merged.getTransportChannel()).isNull(); + } } diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProviderTest.java b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProviderTest.java index 4482f2367a4a..59fe17f1621c 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProviderTest.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProviderTest.java @@ -197,14 +197,13 @@ void managedChannelDoesNotShutdownCustomHttpTransport() throws IOException { HttpJsonTransportChannel httpJsonTransportChannel = provider.getTransportChannel(); - // Verify custom transport is injected + // Verify custom transport is injected (direct ManagedHttpJsonChannel when workloadCertPath is + // null) ManagedHttpJsonInterceptorChannel interceptorChannel = (ManagedHttpJsonInterceptorChannel) httpJsonTransportChannel.getManagedChannel(); ManagedHttpJsonInterceptorChannel managedHttpJsonChannel = (ManagedHttpJsonInterceptorChannel) interceptorChannel.getChannel(); - RefreshingHttpJsonChannel refreshingHttpJsonChannel = - (RefreshingHttpJsonChannel) managedHttpJsonChannel.getChannel(); - ManagedHttpJsonChannel channel = refreshingHttpJsonChannel.getActiveChannel(); + ManagedHttpJsonChannel channel = managedHttpJsonChannel.getChannel(); assertThat(channel.getHttpTransport()).isEqualTo(mockHttpTransport); @@ -215,6 +214,49 @@ void managedChannelDoesNotShutdownCustomHttpTransport() throws IOException { org.mockito.Mockito.verify(mockHttpTransport, org.mockito.Mockito.never()).shutdown(); } + @Test + void channelCreation_withWorkloadCertPath_wrapsWithRefreshingHttpJsonChannel() + throws IOException { + Mockito.when(certificateBasedAccess.getWorkloadCertPath()).thenReturn("fake/cert/path.json"); + + InstantiatingHttpJsonChannelProvider provider = + InstantiatingHttpJsonChannelProvider.newBuilder() + .setEndpoint(DEFAULT_ENDPOINT) + .setCertificateBasedAccess(certificateBasedAccess) + .build(); + provider = (InstantiatingHttpJsonChannelProvider) provider.withHeaders(DEFAULT_HEADER_MAP); + + HttpJsonTransportChannel httpJsonTransportChannel = provider.getTransportChannel(); + + ManagedHttpJsonInterceptorChannel interceptorChannel = + (ManagedHttpJsonInterceptorChannel) httpJsonTransportChannel.getManagedChannel(); + ManagedHttpJsonInterceptorChannel managedHttpJsonChannel = + (ManagedHttpJsonInterceptorChannel) interceptorChannel.getChannel(); + assertThat(managedHttpJsonChannel.getChannel()).isInstanceOf(RefreshingHttpJsonChannel.class); + + provider.getTransportChannel().shutdownNow(); + } + + @Test + void createHttpTransport_withMtlsAndConscrypt_configuresSecurityProvider() + throws IOException, GeneralSecurityException { + Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); + com.google.auth.mtls.MtlsProvider provider = + new com.google.api.gax.rpc.testing.FakeMtlsProvider( + com.google.api.gax.rpc.testing.FakeMtlsProvider.createTestMtlsKeyStore(), "", false); + + InstantiatingHttpJsonChannelProvider channelProvider = + InstantiatingHttpJsonChannelProvider.newBuilder() + .setEndpoint(DEFAULT_ENDPOINT) + .setMtlsProvider(provider) + .setCertificateBasedAccess(certificateBasedAccess) + .build(); + + com.google.api.client.http.HttpTransport transport = channelProvider.createHttpTransport(); + assertThat(transport).isNotNull(); + assertThat(transport).isInstanceOf(com.google.api.client.http.javanet.NetHttpTransport.class); + } + @Override protected Object getMtlsObjectFromTransportChannel( MtlsProvider provider, CertificateBasedAccess certificateBasedAccess) diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java index ea147deb74bc..2ac955592c9d 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java @@ -389,4 +389,25 @@ void testConcurrentNewCallDuringRefresh() throws InterruptedException { assertEquals(threadCount, successCount.get()); } + + @Test + void testGenerationIncrementAndLifecycleOnDelegatingWrapper() throws Exception { + RefreshingHttpJsonChannel channel = createTestChannel(); + assertEquals(0, channel.getGeneration()); + + channel.invalidateDiskFingerprintCache(); + testFingerprint = "fingerprint2"; + channel.refresh(); + + assertEquals(1, channel.getGeneration()); + + // Verify lifecycle methods on delegating wrapper do not throw NullPointerException + assertFalse(channel.isShutdown()); + assertFalse(channel.isTerminated()); + channel.shutdown(); + assertTrue(channel.isShutdown()); + channel.shutdownNow(); + assertTrue(channel.isTerminated()); + assertTrue(channel.awaitTermination(1, TimeUnit.SECONDS)); + } } diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/AttemptCallable.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/AttemptCallable.java index e9978959f85a..66dbcb2dcf14 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/AttemptCallable.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/AttemptCallable.java @@ -35,6 +35,8 @@ import com.google.api.gax.retrying.RetryingFuture; import com.google.common.base.Preconditions; import java.util.concurrent.Callable; +import java.util.logging.Level; +import java.util.logging.Logger; import org.jspecify.annotations.NullMarked; /** @@ -48,6 +50,7 @@ */ @NullMarked class AttemptCallable implements Callable { + private static final Logger LOG = Logger.getLogger(AttemptCallable.class.getName()); private final UnaryCallable callable; private final RequestT request; private final ApiCallContext originalCallContext; @@ -85,6 +88,10 @@ public ResponseT call() { .getTracer() .attemptStarted(request, externalFuture.getAttemptSettings().getOverallAttemptCount()); + TransportChannel transportChannel = callContext.getTransportChannel(); + final long attemptGeneration = + transportChannel != null ? transportChannel.getGeneration() : 0; + ApiFuture internalFuture = callable.futureCall(request, callContext); final ApiCallContext finalContext = callContext; ApiFuture mappedFuture = @@ -92,21 +99,38 @@ public ResponseT call() { internalFuture, UnauthenticatedException.class, unauthenticatedException -> { - TransportChannel transportChannel = finalContext.getTransportChannel(); - if (transportChannel != null && transportChannel.shouldRefresh()) { - transportChannel.refresh(); - UnauthenticatedException newEx = - new UnauthenticatedException( - unauthenticatedException.getMessage(), - unauthenticatedException.getCause(), - unauthenticatedException.getStatusCode(), - true, // isRetryable = true - unauthenticatedException.getErrorDetails()); - newEx.setStackTrace(unauthenticatedException.getStackTrace()); - for (Throwable suppressed : unauthenticatedException.getSuppressed()) { - newEx.addSuppressed(suppressed); + TransportChannel channel = finalContext.getTransportChannel(); + if (channel != null) { + boolean shouldRetry = false; + if (channel.shouldRefresh()) { + try { + channel.refresh(); + } catch (Exception e) { + LOG.log( + Level.WARNING, + "Failed to refresh transport channel after authentication error", + e); + } + shouldRetry = true; + } else if (channel.getGeneration() > attemptGeneration) { + // Channel was rotated by a concurrent request while this call was in flight + shouldRetry = true; + } + + if (shouldRetry) { + UnauthenticatedException newEx = + new UnauthenticatedException( + unauthenticatedException.getMessage(), + unauthenticatedException.getCause(), + unauthenticatedException.getStatusCode(), + true, // isRetryable = true + unauthenticatedException.getErrorDetails()); + newEx.setStackTrace(unauthenticatedException.getStackTrace()); + for (Throwable suppressed : unauthenticatedException.getSuppressed()) { + newEx.addSuppressed(suppressed); + } + throw newEx; } - throw newEx; } throw unauthenticatedException; }, diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/BidiStreamingCallable.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/BidiStreamingCallable.java index 980aa5b9ac26..c8fe6b2a59d3 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/BidiStreamingCallable.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/BidiStreamingCallable.java @@ -29,6 +29,8 @@ */ package com.google.api.gax.rpc; +import java.util.logging.Level; +import java.util.logging.Logger; import org.jspecify.annotations.NullMarked; import org.jspecify.annotations.Nullable; @@ -42,6 +44,7 @@ */ @NullMarked public abstract class BidiStreamingCallable { + private static final Logger LOG = Logger.getLogger(BidiStreamingCallable.class.getName()); protected BidiStreamingCallable() {} @@ -262,20 +265,14 @@ public void onError(Throwable t) { if (t instanceof UnauthenticatedException) { TransportChannel transportChannel = mergedContext.getTransportChannel(); if (transportChannel != null && transportChannel.shouldRefresh()) { - transportChannel.refresh(); - UnauthenticatedException causeEx = (UnauthenticatedException) t; - UnauthenticatedException newEx = - new UnauthenticatedException( - causeEx.getMessage(), - causeEx.getCause(), - causeEx.getStatusCode(), - true, // isRetryable = true - causeEx.getErrorDetails()); - newEx.setStackTrace(causeEx.getStackTrace()); - for (Throwable suppressed : causeEx.getSuppressed()) { - newEx.addSuppressed(suppressed); + try { + transportChannel.refresh(); + } catch (Exception e) { + LOG.log( + Level.WARNING, + "Failed to refresh transport channel after authentication error", + e); } - t = newEx; } } responseObserver.onError(t); diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ClientStreamingCallable.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ClientStreamingCallable.java index 9bff9209fbe6..5c87e4c558b7 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ClientStreamingCallable.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ClientStreamingCallable.java @@ -29,6 +29,8 @@ */ package com.google.api.gax.rpc; +import java.util.logging.Level; +import java.util.logging.Logger; import org.jspecify.annotations.NullMarked; import org.jspecify.annotations.Nullable; @@ -42,6 +44,7 @@ */ @NullMarked public abstract class ClientStreamingCallable { + private static final Logger LOG = Logger.getLogger(ClientStreamingCallable.class.getName()); protected ClientStreamingCallable() {} @@ -91,20 +94,14 @@ public void onError(Throwable t) { if (t instanceof UnauthenticatedException) { TransportChannel transportChannel = mergedContext.getTransportChannel(); if (transportChannel != null && transportChannel.shouldRefresh()) { - transportChannel.refresh(); - UnauthenticatedException causeEx = (UnauthenticatedException) t; - UnauthenticatedException newEx = - new UnauthenticatedException( - causeEx.getMessage(), - causeEx.getCause(), - causeEx.getStatusCode(), - true, // isRetryable = true - causeEx.getErrorDetails()); - newEx.setStackTrace(causeEx.getStackTrace()); - for (Throwable suppressed : causeEx.getSuppressed()) { - newEx.addSuppressed(suppressed); + try { + transportChannel.refresh(); + } catch (Exception e) { + LOG.log( + Level.WARNING, + "Failed to refresh transport channel after authentication error", + e); } - t = newEx; } } responseObserver.onError(t); diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ServerStreamingAttemptCallable.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ServerStreamingAttemptCallable.java index 26951d355f79..7ef3c8491cb7 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ServerStreamingAttemptCallable.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ServerStreamingAttemptCallable.java @@ -37,6 +37,8 @@ import com.google.errorprone.annotations.concurrent.GuardedBy; import java.util.concurrent.Callable; import java.util.concurrent.CancellationException; +import java.util.logging.Level; +import java.util.logging.Logger; import org.jspecify.annotations.NullMarked; /** @@ -96,6 +98,8 @@ */ @NullMarked final class ServerStreamingAttemptCallable implements Callable { + private static final Logger LOG = + Logger.getLogger(ServerStreamingAttemptCallable.class.getName()); private final Object lock = new Object(); private final ServerStreamingCallable innerCallable; @@ -244,22 +248,14 @@ public void onErrorImpl(Throwable t) { if (cause instanceof UnauthenticatedException) { TransportChannel transportChannel = finalContext.getTransportChannel(); if (transportChannel != null && transportChannel.shouldRefresh()) { - transportChannel.refresh(); - UnauthenticatedException causeEx = (UnauthenticatedException) cause; - UnauthenticatedException newEx = - new UnauthenticatedException( - causeEx.getMessage(), - causeEx.getCause(), - causeEx.getStatusCode(), - true, // isRetryable = true - causeEx.getErrorDetails()); - newEx.setStackTrace(causeEx.getStackTrace()); - for (Throwable suppressed : causeEx.getSuppressed()) { - newEx.addSuppressed(suppressed); + try { + transportChannel.refresh(); + } catch (Exception e) { + LOG.log( + Level.WARNING, + "Failed to refresh transport channel after authentication error", + e); } - cause = newEx; - - t = cause; } } onAttemptError(t); diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/TransportChannel.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/TransportChannel.java index de83dbc73861..cad80d4f87b3 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/TransportChannel.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/TransportChannel.java @@ -63,4 +63,12 @@ default void refresh() {} default boolean shouldRefresh() { return false; } + + /** + * Returns a monotonic generation counter tracking the number of successful refreshes or channel + * rotations performed by this transport channel. + */ + default long getGeneration() { + return 0; + } } diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/CertificateBasedAccess.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/CertificateBasedAccess.java index 0dbdf60723d2..d77378ad9612 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/CertificateBasedAccess.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/CertificateBasedAccess.java @@ -34,7 +34,6 @@ import com.google.api.gax.rpc.internal.EnvironmentProvider; import com.google.auth.mtls.MtlsUtils; import com.google.auth.oauth2.PropertyProvider; -import java.io.IOException; /** * Utility class for handling certificate-based access configurations. @@ -47,30 +46,7 @@ public class CertificateBasedAccess { private final EnvironmentProvider envProvider; - @InternalApi - public interface FileExistenceProvider { - boolean exists(String path); - } - - @InternalApi - public interface FileContentReader { - String read(String path) throws IOException; - } - public CertificateBasedAccess(EnvironmentProvider envProvider) { - this( - envProvider, - path -> new java.io.File(path).isFile(), - path -> - new String( - java.nio.file.Files.readAllBytes(java.nio.file.Paths.get(path)), - java.nio.charset.StandardCharsets.UTF_8)); - } - - CertificateBasedAccess( - EnvironmentProvider envProvider, - FileExistenceProvider fileExistenceProvider, - FileContentReader fileContentReader) { this.envProvider = envProvider; } diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/AttemptCallableTest.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/AttemptCallableTest.java index 2916d8a1d019..baeb35630014 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/AttemptCallableTest.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/AttemptCallableTest.java @@ -183,4 +183,142 @@ void testUnauthenticatedExceptionReThrowPreservesContext() { assertThat(rethrown.getSuppressed().length).isEqualTo(1); assertThat(rethrown.getSuppressed()[0]).isInstanceOf(RuntimeException.class); } + + @Test + void testSiblingInFlightRequest_channelRotatedInFlight_markedRetryableWithoutDuplicateRefresh() { + FakeChannel innerChannel = new FakeChannel(); + innerChannel.setGeneration(1); + // Initially shouldRefresh is false because sibling request already completed the refresh + innerChannel.setShouldRefresh(false); + FakeTransportChannel transportChannel = FakeTransportChannel.create(innerChannel); + + ApiCallContext callContext = + FakeCallContext.createDefault().withTransportChannel(transportChannel); + + UnauthenticatedException originalEx = + new UnauthenticatedException( + "Expired cert", null, FakeStatusCode.of(StatusCode.Code.UNAUTHENTICATED), false); + + SettableApiFuture failedFuture = SettableApiFuture.create(); + failedFuture.setException(originalEx); + when(mockInnerCallable.futureCall(Mockito.anyString(), Mockito.any())) + .thenAnswer( + invocation -> { + // While request was in flight, sibling finished rotation and bumped generation to 2 + innerChannel.setGeneration(2); + return failedFuture; + }); + + AttemptCallable callable = + new AttemptCallable<>(mockInnerCallable, "fake-request", callContext); + callable.setExternalFuture(mockExternalFuture); + + callable.call(); + + ArgumentCaptor futureCaptor = ArgumentCaptor.forClass(ApiFuture.class); + Mockito.verify(mockExternalFuture, Mockito.times(2)).setAttemptFuture(futureCaptor.capture()); + + Throwable thrown = null; + try { + futureCaptor.getValue().get(); + } catch (Exception e) { + thrown = e.getCause(); + } + + assertThat(thrown).isInstanceOf(UnauthenticatedException.class); + UnauthenticatedException rethrown = (UnauthenticatedException) thrown; + // Sibling request should be marked retryable to run on the new channel + assertThat(rethrown.isRetryable()).isTrue(); + // But should NOT have triggered a second refresh call + assertThat(innerChannel.getRefreshCount()).isEqualTo(0); + } + + @Test + void testPermanentUnauthenticatedFailure_sameGeneration_notMarkedRetryable() { + FakeChannel innerChannel = new FakeChannel(); + innerChannel.setGeneration(1); + innerChannel.setShouldRefresh(false); + FakeTransportChannel transportChannel = FakeTransportChannel.create(innerChannel); + + ApiCallContext callContext = + FakeCallContext.createDefault().withTransportChannel(transportChannel); + + UnauthenticatedException originalEx = + new UnauthenticatedException( + "Invalid credentials", null, FakeStatusCode.of(StatusCode.Code.UNAUTHENTICATED), false); + + SettableApiFuture failedFuture = SettableApiFuture.create(); + failedFuture.setException(originalEx); + when(mockInnerCallable.futureCall(Mockito.anyString(), Mockito.any())).thenReturn(failedFuture); + + AttemptCallable callable = + new AttemptCallable<>(mockInnerCallable, "fake-request", callContext); + callable.setExternalFuture(mockExternalFuture); + + callable.call(); + + ArgumentCaptor futureCaptor = ArgumentCaptor.forClass(ApiFuture.class); + Mockito.verify(mockExternalFuture, Mockito.times(2)).setAttemptFuture(futureCaptor.capture()); + + Throwable thrown = null; + try { + futureCaptor.getValue().get(); + } catch (Exception e) { + thrown = e.getCause(); + } + + assertThat(thrown).isInstanceOf(UnauthenticatedException.class); + UnauthenticatedException rethrown = (UnauthenticatedException) thrown; + // Genuine permanent error on same generation is NOT retryable + assertThat(rethrown.isRetryable()).isFalse(); + assertThat(innerChannel.getRefreshCount()).isEqualTo(0); + } + + @Test + void testRefreshThrowsException_originalUnauthenticatedPropagated() { + FakeChannel fakeChannel = + new FakeChannel() { + @Override + public boolean shouldRefresh() { + return true; + } + + @Override + public void refresh() { + throw new RuntimeException("Refresh error"); + } + }; + FakeTransportChannel transportChannel = FakeTransportChannel.create(fakeChannel); + + ApiCallContext callContext = + FakeCallContext.createDefault().withTransportChannel(transportChannel); + + UnauthenticatedException originalEx = + new UnauthenticatedException( + "Expired cert", null, FakeStatusCode.of(StatusCode.Code.UNAUTHENTICATED), false); + + SettableApiFuture failedFuture = SettableApiFuture.create(); + failedFuture.setException(originalEx); + when(mockInnerCallable.futureCall(Mockito.anyString(), Mockito.any())).thenReturn(failedFuture); + + AttemptCallable callable = + new AttemptCallable<>(mockInnerCallable, "fake-request", callContext); + callable.setExternalFuture(mockExternalFuture); + + callable.call(); + + ArgumentCaptor futureCaptor = ArgumentCaptor.forClass(ApiFuture.class); + Mockito.verify(mockExternalFuture, Mockito.times(2)).setAttemptFuture(futureCaptor.capture()); + + Throwable thrown = null; + try { + futureCaptor.getValue().get(); + } catch (Exception e) { + thrown = e.getCause(); + } + + assertThat(thrown).isInstanceOf(UnauthenticatedException.class); + UnauthenticatedException rethrown = (UnauthenticatedException) thrown; + assertThat(rethrown.isRetryable()).isTrue(); + } } diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ServerStreamingAttemptCallableTest.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ServerStreamingAttemptCallableTest.java index 2a52cb2b2572..ca6ea284d845 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ServerStreamingAttemptCallableTest.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ServerStreamingAttemptCallableTest.java @@ -286,7 +286,43 @@ void testUnauthenticatedRefresh() { Truth.assertThat(((ServerStreamingAttemptException) outerError).hasSeenResponses()).isFalse(); Truth.assertThat(((ServerStreamingAttemptException) outerError).canResume()).isTrue(); Truth.assertThat(outerError.getCause()).isInstanceOf(UnauthenticatedException.class); - Truth.assertThat(((UnauthenticatedException) outerError.getCause()).isRetryable()).isTrue(); + Truth.assertThat(((UnauthenticatedException) outerError.getCause()).isRetryable()).isFalse(); + } + + @Test + @SuppressWarnings("ConstantConditions") + void testRefreshThrowsException_originalErrorNotLost() { + TransportChannel transportChannel = Mockito.mock(TransportChannel.class); + Mockito.when(transportChannel.shouldRefresh()).thenReturn(true); + Mockito.doThrow(new RuntimeException("Refresh error")).when(transportChannel).refresh(); + + ApiCallContext context = Mockito.mock(ApiCallContext.class); + Mockito.when(context.getTransportChannel()).thenReturn(transportChannel); + Mockito.when(context.getTracer()).thenReturn(BaseApiTracer.getInstance()); + Mockito.when(context.getTimeoutDuration()).thenReturn(java.time.Duration.ofHours(5)); + + resumptionStrategy = new MyStreamResumptionStrategy(); + ServerStreamingAttemptCallable callable = createCallable(context); + callable.start(); + + MockServerStreamingCall call = innerCallable.popLastCall(); + + UnauthenticatedException initialError = + new UnauthenticatedException( + "test", + null, + com.google.api.gax.rpc.testing.FakeStatusCode.of(Code.UNAUTHENTICATED), + false); + call.getController().getObserver().onError(initialError); + + ExecutionException ee = + assertThrows( + ExecutionException.class, + () -> fakeRetryingFuture.getAttemptResult().get(1, TimeUnit.SECONDS)); + Throwable outerError = ee.getCause(); + Mockito.verify(transportChannel).refresh(); + Truth.assertThat(outerError).isInstanceOf(ServerStreamingAttemptException.class); + Truth.assertThat(outerError.getCause()).isEqualTo(initialError); } @Test diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/testing/FakeChannel.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/testing/FakeChannel.java index 6d676db40032..3f77f9affe93 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/testing/FakeChannel.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/testing/FakeChannel.java @@ -52,4 +52,15 @@ public void refresh() { public int getRefreshCount() { return refreshCount; } + + private volatile long generation = 0; + + public FakeChannel setGeneration(long generation) { + this.generation = generation; + return this; + } + + public long getGeneration() { + return generation; + } } diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/testing/FakeTransportChannel.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/testing/FakeTransportChannel.java index bbde9feb4807..7d92556b88d5 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/testing/FakeTransportChannel.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/testing/FakeTransportChannel.java @@ -69,6 +69,21 @@ public int getRefreshCount() { return channel != null ? channel.getRefreshCount() : refreshCount; } + private volatile long generation = 0; + + public FakeTransportChannel setGeneration(long generation) { + if (channel != null) { + channel.setGeneration(generation); + } + this.generation = generation; + return this; + } + + @Override + public long getGeneration() { + return channel != null ? channel.getGeneration() : generation; + } + private FakeTransportChannel(FakeChannel channel) { this.channel = channel; } From a97680b714691628942f14822d66a8fea3d25910 Mon Sep 17 00:00:00 2001 From: Matt Castelaz Date: Fri, 28 Aug 2026 21:12:52 +0000 Subject: [PATCH 13/29] fix(gax-grpc): use javax.annotation.concurrent.GuardedBy to satisfy dependency analyzer --- .../src/main/java/com/google/api/gax/grpc/ChannelPool.java | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java index 8056ad734bbb..a743dd1a071f 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java @@ -35,7 +35,7 @@ import com.google.common.annotations.VisibleForTesting; import com.google.common.base.Preconditions; import com.google.common.collect.ImmutableList; -import com.google.errorprone.annotations.concurrent.GuardedBy; +import javax.annotation.concurrent.GuardedBy; import io.grpc.CallOptions; import io.grpc.Channel; import io.grpc.ClientCall; From 79213966eb12379e77ab1f51170d4870bfc20106 Mon Sep 17 00:00:00 2001 From: Matt Castelaz Date: Fri, 18 Sep 2026 01:45:35 +0000 Subject: [PATCH 14/29] fix(gax): address review feedback for mTLS certificate rotation retries (#13995) - CertificateRotationTracker: extract shared mTLS disk fingerprint rotation tracking and 1s positive rotation cache into core gax, using monotonic sequence numbers incremented before disk I/O to coalesce concurrent lock waiters without caching unchanged checks - WorkloadCertificateUtils & MtlsUtils: simplify getCertificateFingerprint(), document return/throw contracts, treat 0-byte truncated certificate files mid-write as empty string, and return true in useMtlsClientCertificate() when workloadCertPath is present while preserving ECP support - ChannelPool: avoid marking pool rotated on partial refresh failure in refreshAll(), clean up newly created entries if any creation fails, guard refreshSafely() against mid-write empty fingerprints, and make getGeneration() package-private - GrpcCallContext & HttpJsonCallContext: allow withChannel(null) to clear the channel - AttemptCallable & ServerStreamingAttemptCallable: check channel.getGeneration() > attemptGeneration after refresh() so failed refreshes do not loop retries on unrotated channels, and mark server-streaming UnauthenticatedException retryable when channel rotates - ApiResultRetryAlgorithm: grant one immediate free retry on retryable UnauthenticatedException even when maxAttempts is 1 or totalTimeout is 0 - ChannelPool & RefreshingHttpJsonChannel: synchronize start() and cancel() on a per-call lock, guard against duplicate start(), only release immediately on cancel() exception if call was not started, catch Throwable in newCall()/start()/cancel(), and re-check outstandingCalls.get() == 0 after shutdownRequested.get() - InstantiatingHttpJsonChannelProvider: gate workloadCertPath on active mTLS without custom HttpTransport, pass initialChannel directly to RefreshingHttpJsonChannel to preserve checked IOException on startup, and guard against leaks and null keystore fallback - InstantiatingGrpcChannelProvider: gate workloadCertPath on !canUseDirectPath() && active mTLS, and fail fast with IOException if mTLS channel credentials cannot be initialized when mTLS is active --- .../java/com/google/auth/mtls/MtlsUtils.java | 50 ++- .../com/google/auth/mtls/MtlsUtilsTest.java | 68 +++- .../com/google/api/gax/grpc/ChannelPool.java | 270 ++++++------- .../google/api/gax/grpc/GrpcCallContext.java | 2 +- .../InstantiatingGrpcChannelProvider.java | 10 +- .../google/api/gax/grpc/ChannelPoolTest.java | 355 ++++++++++++++++++ .../api/gax/grpc/GrpcCallContextTest.java | 10 + .../InstantiatingGrpcChannelProviderTest.java | 91 +++++ .../api/gax/httpjson/HttpJsonCallContext.java | 2 +- .../InstantiatingHttpJsonChannelProvider.java | 89 +++-- .../httpjson/RefreshingHttpJsonChannel.java | 195 +++++----- .../gax/httpjson/HttpJsonCallContextTest.java | 15 +- ...tantiatingHttpJsonChannelProviderTest.java | 80 +++- .../RefreshingHttpJsonChannelTest.java | 214 ++++++++++- .../api/gax/rpc/ApiResultRetryAlgorithm.java | 43 ++- .../google/api/gax/rpc/AttemptCallable.java | 6 +- .../rpc/ServerStreamingAttemptCallable.java | 45 ++- .../rpc/mtls/CertificateRotationTracker.java | 197 ++++++++++ .../rpc/mtls/WorkloadCertificateUtils.java | 21 +- .../gax/rpc/ApiResultRetryAlgorithmTest.java | 137 +++++++ .../api/gax/rpc/AttemptCallableTest.java | 101 ++++- .../ServerStreamingAttemptCallableTest.java | 173 +++++++++ .../rpc/mtls/CertificateBasedAccessTest.java | 22 +- .../api/gax/rpc/testing/FakeChannel.java | 1 + 24 files changed, 1892 insertions(+), 305 deletions(-) create mode 100644 sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/CertificateRotationTracker.java diff --git a/google-auth-library-java/oauth2_http/java/com/google/auth/mtls/MtlsUtils.java b/google-auth-library-java/oauth2_http/java/com/google/auth/mtls/MtlsUtils.java index f5c5fd694836..d46d3822f827 100644 --- a/google-auth-library-java/oauth2_http/java/com/google/auth/mtls/MtlsUtils.java +++ b/google-auth-library-java/oauth2_http/java/com/google/auth/mtls/MtlsUtils.java @@ -40,6 +40,7 @@ import java.io.IOException; import java.io.InputStream; import java.nio.file.Files; +import java.nio.file.Paths; import java.security.MessageDigest; import java.util.Locale; import org.jspecify.annotations.NullMarked; @@ -62,17 +63,43 @@ private MtlsUtils() { } /** - * Returns if mutual TLS client certificate should be used. Delegates directly to - * getWorkloadCertPath to avoid duplicate logic. + * Returns if mutual TLS client certificate should be used. Returns true if valid workload + * certificates are configured or if GOOGLE_API_USE_CLIENT_CERTIFICATE is explicitly set to true + * (e.g. for Enterprise Certificate Proxy or custom MtlsProviders), unless explicitly disabled via + * GOOGLE_API_USE_CLIENT_CERTIFICATE=false. */ public static boolean useMtlsClientCertificate( EnvironmentProvider envProvider, PropertyProvider propProvider) { - return getWorkloadCertPath(envProvider, propProvider) != null; + String useClientCertificate = envProvider.getEnv("GOOGLE_API_USE_CLIENT_CERTIFICATE"); + if ("false".equalsIgnoreCase(useClientCertificate)) { + return false; + } + if (getWorkloadCertPath(envProvider, propProvider) != null) { + return true; + } + return "true".equalsIgnoreCase(useClientCertificate); } /** * Resolves and returns the path to the mutual TLS client certificate, or null if none should be * used. + * + *

Possible outcomes: + * + *

    + *
  1. Non-null {@link String} (Valid happy path): A valid workload certificate + * configuration was found and both the certificate and private key files exist and are + * readable. + *
  2. {@link IllegalStateException} (Invalid state - fail closed): An explicit {@code + * GOOGLE_API_CERTIFICATE_CONFIG} path or an existing default well-known certificate + * configuration file is missing, unreadable, malformed, or references missing/unreadable + * certificate or private key files. This is treated as an unrecoverable misconfiguration. + *
  3. {@code null} (Safe fallback / fail open): Client certificates are explicitly + * disabled via {@code GOOGLE_API_USE_CLIENT_CERTIFICATE=false}, no explicit configuration + * is set and the default well-known configuration file does not exist on disk, or the + * configuration specifies an non-workload source (e.g., ECP/PKCS11 without a {@code + * workload} section). Callers can proceed without workload certificate file polling. + *
*/ public static @Nullable String getWorkloadCertPath( EnvironmentProvider envProvider, PropertyProvider propProvider) { @@ -174,17 +201,22 @@ private static void checkCertAndKeyFilesReadable( } } - /** Centralized SHA-256 Fingerprint Calculator */ + /** + * Computes the lower-case SHA-256 hex fingerprint of the certificate file at {@code certPath}. + * + *

Unlike {@link #getWorkloadCertPath}, which validates configuration at channel initialization + * and fails closed on errors, this method is called dynamically at runtime during active RPCs to + * detect certificate rotations on disk. External certificate rotators may temporarily delete, + * truncate, or rewrite the certificate file mid-RPC. Returning {@code null} on read/digest + * exceptions (which callers normalize to {@code ""}) allows runtime refresh checks to ignore + * transient mid-write states and keep the active healthy channel without failing in-flight RPCs. + */ public static @Nullable String getCertificateFingerprint(@Nullable String certPath) { if (certPath == null) { return null; } - File file = new File(certPath); - if (!file.isFile() || !file.canRead()) { - return null; - } try { - byte[] certBytes = Files.readAllBytes(file.toPath()); + byte[] certBytes = Files.readAllBytes(Paths.get(certPath)); byte[] digest = MessageDigest.getInstance("SHA-256").digest(certBytes); return BaseEncoding.base16().lowerCase().encode(digest); } catch (Exception e) { diff --git a/google-auth-library-java/oauth2_http/javatests/com/google/auth/mtls/MtlsUtilsTest.java b/google-auth-library-java/oauth2_http/javatests/com/google/auth/mtls/MtlsUtilsTest.java index 3c1a06614705..d5fdc840d804 100644 --- a/google-auth-library-java/oauth2_http/javatests/com/google/auth/mtls/MtlsUtilsTest.java +++ b/google-auth-library-java/oauth2_http/javatests/com/google/auth/mtls/MtlsUtilsTest.java @@ -260,12 +260,36 @@ public String getProperty(String name, String def) { } @Test - void useMtlsClientCertificate_trueWithNoCertsOnDisk_returnsFalseWithoutThrowing() { + void + useMtlsClientCertificate_trueWithNoCertsOnDisk_returnsTrueWhileWorkloadCertPathReturnsNull() { EnvironmentProvider envProvider = name -> "GOOGLE_API_USE_CLIENT_CERTIFICATE".equals(name) ? "true" : null; + PropertyProvider propProvider = + (name, def) -> { + if ("user.home".equals(name)) return tempDir.toString(); + if ("os.name".equals(name)) return "Linux"; + return def; + }; + + assertTrue(MtlsUtils.useMtlsClientCertificate(envProvider, propProvider)); + assertNull(MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); + } + + @Test + void useMtlsClientCertificate_trueWithEcpOnlyConfig_returnsTrueAndWorkloadCertPathReturnsNull() + throws IOException { + Path configFile = tempDir.resolve("ecp_config.json"); + Files.write(configFile, "{\"cert_configs\":{\"enterprise_certificates\":{}}}".getBytes()); + + EnvironmentProvider envProvider = + name -> { + if ("GOOGLE_API_USE_CLIENT_CERTIFICATE".equals(name)) return "true"; + if ("GOOGLE_API_CERTIFICATE_CONFIG".equals(name)) return configFile.toString(); + return null; + }; PropertyProvider propProvider = (name, def) -> def; - assertFalse(MtlsUtils.useMtlsClientCertificate(envProvider, propProvider)); + assertTrue(MtlsUtils.useMtlsClientCertificate(envProvider, propProvider)); assertNull(MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); } @@ -279,6 +303,46 @@ void useMtlsClientCertificate_false_returnsFalse() { assertNull(MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); } + @Test + void useMtlsClientCertificate_falseEvenWhenWorkloadCertsExist_returnsFalse() throws IOException { + Path certFile = tempDir.resolve("cert.pem"); + Path keyFile = tempDir.resolve("key.pem"); + Files.write(certFile, "dummy cert".getBytes()); + Files.write(keyFile, "dummy key".getBytes()); + + Path configFile = tempDir.resolve("config.json"); + String configJson = + String.format( + "{\"cert_configs\":{\"workload\":{\"cert_path\":\"%s\",\"key_path\":\"%s\"}}}", + certFile.toString().replace("\\", "\\\\"), keyFile.toString().replace("\\", "\\\\")); + Files.write(configFile, configJson.getBytes()); + + EnvironmentProvider envProvider = + name -> { + if ("GOOGLE_API_USE_CLIENT_CERTIFICATE".equals(name)) return "false"; + if ("GOOGLE_API_CERTIFICATE_CONFIG".equals(name)) return configFile.toString(); + return null; + }; + PropertyProvider propProvider = (name, def) -> def; + + assertFalse(MtlsUtils.useMtlsClientCertificate(envProvider, propProvider)); + assertNull(MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); + } + + @Test + void useMtlsClientCertificate_unsetWithNoCertsOnDisk_returnsFalse() { + EnvironmentProvider envProvider = name -> null; + PropertyProvider propProvider = + (name, def) -> { + if ("user.home".equals(name)) return tempDir.toString(); + if ("os.name".equals(name)) return "Linux"; + return def; + }; + + assertFalse(MtlsUtils.useMtlsClientCertificate(envProvider, propProvider)); + assertNull(MtlsUtils.getWorkloadCertPath(envProvider, propProvider)); + } + // --- Explicit GOOGLE_API_CERTIFICATE_CONFIG Tests (Fail Closed) --- @Test diff --git a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java index a743dd1a071f..8004593a2673 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java @@ -31,11 +31,11 @@ import com.google.api.core.InternalApi; import com.google.api.gax.core.FixedExecutorProvider; -import com.google.api.gax.rpc.mtls.WorkloadCertificateUtils; +import com.google.api.gax.rpc.mtls.CertificateRotationTracker; import com.google.common.annotations.VisibleForTesting; import com.google.common.base.Preconditions; +import com.google.common.base.Strings; import com.google.common.collect.ImmutableList; -import javax.annotation.concurrent.GuardedBy; import io.grpc.CallOptions; import io.grpc.Channel; import io.grpc.ClientCall; @@ -59,6 +59,7 @@ import java.util.concurrent.atomic.AtomicReference; import java.util.logging.Level; import java.util.logging.Logger; +import javax.annotation.concurrent.GuardedBy; import org.jspecify.annotations.NullMarked; import org.jspecify.annotations.Nullable; @@ -90,26 +91,13 @@ class ChannelPool extends ManagedChannel { private @Nullable ScheduledFuture refreshFuture = null; private @Nullable ScheduledFuture resizeFuture = null; - private static class DiskCheckResult { - final String fingerprint; - final long timestampNanos; - - DiskCheckResult(String fingerprint, long timestampNanos) { - this.fingerprint = fingerprint; - this.timestampNanos = timestampNanos; - } - } - - private volatile DiskCheckResult lastDiskCheck = null; - private final java.util.concurrent.locks.ReentrantLock diskCheckLock = - new java.util.concurrent.locks.ReentrantLock(); + private final CertificateRotationTracker rotationTracker; private final Object entryWriteLock = new Object(); @GuardedBy("entryWriteLock") private boolean isShutdown = false; private final AtomicLong generation = new AtomicLong(0); - private volatile String activeCertFingerprint = ""; @VisibleForTesting final AtomicReference> entries = new AtomicReference<>(); private final AtomicInteger indexTicker = new AtomicInteger(); private final String authority; @@ -155,6 +143,7 @@ static ChannelPool create( this.channelFactory = channelFactory; this.backgroundExecutorProvider = executorProvider; this.workloadCertPath = workloadCertPath; + this.rotationTracker = new CertificateRotationTracker(workloadCertPath); ImmutableList.Builder initialListBuilder = ImmutableList.builder(); @@ -165,11 +154,6 @@ static ChannelPool create( entries.set(initialListBuilder.build()); authority = entries.get().get(0).channel.authority(); - if (workloadCertPath != null) { - this.activeCertFingerprint = - WorkloadCertificateUtils.getCertificateFingerprint(workloadCertPath); - } - if (!settings.isStaticSize()) { resizeFuture = backgroundExecutorProvider @@ -464,14 +448,22 @@ private void expand(int desiredSize) { entries.set(newEntries.build()); } + /** + * Periodically refreshes all channels when {@link + * ChannelPoolSettings#isPreemptiveRefreshEnabled()} is enabled (to mitigate hourly GFE + * disconnects). This applies to all channels even when {@code workloadCertPath == null}. If + * {@code workloadCertPath} is configured, also updates the tracked certificate fingerprint on + * success (or skips if the certificate file is currently unreadable or mid-write on disk). + */ private void refreshSafely() { try { synchronized (entryWriteLock) { - if (refreshAll() && workloadCertPath != null) { - String currentDiskFingerprint = getOrUpdateDiskFingerprint(workloadCertPath); - if (!currentDiskFingerprint.isEmpty()) { - this.activeCertFingerprint = currentDiskFingerprint; - } + String currentDiskFingerprint = rotationTracker.readDiskFingerprint(); + if (workloadCertPath != null && currentDiskFingerprint.isEmpty()) { + return; + } + if (refreshAll() && !currentDiskFingerprint.isEmpty()) { + rotationTracker.markRefreshed(currentDiskFingerprint); } } } catch (Exception e) { @@ -479,43 +471,13 @@ private void refreshSafely() { } } - private String getOrUpdateDiskFingerprint(String certPath) { - long now = System.nanoTime(); - DiskCheckResult cached = lastDiskCheck; - if (cached != null - && (now - cached.timestampNanos < java.util.concurrent.TimeUnit.SECONDS.toNanos(1))) { - return cached.fingerprint; - } - - diskCheckLock.lock(); - try { - cached = lastDiskCheck; - if (cached != null - && (now - cached.timestampNanos < java.util.concurrent.TimeUnit.SECONDS.toNanos(1))) { - return cached.fingerprint; - } - String fingerprint = WorkloadCertificateUtils.getCertificateFingerprint(certPath); - lastDiskCheck = new DiskCheckResult(fingerprint, System.nanoTime()); - return fingerprint; - } finally { - diskCheckLock.unlock(); - } - } - @VisibleForTesting void invalidateDiskFingerprintCache() { - this.lastDiskCheck = null; + rotationTracker.invalidateCache(); } boolean shouldRefresh() { - if (workloadCertPath == null) { - return false; - } - String currentDiskFingerprint = getOrUpdateDiskFingerprint(workloadCertPath); - if (currentDiskFingerprint.isEmpty()) { - return false; - } - return !currentDiskFingerprint.equalsIgnoreCase(activeCertFingerprint); + return rotationTracker.shouldRefresh(); } /** @@ -541,13 +503,13 @@ void refresh() { refreshAll(); return; } - String currentDiskFingerprint = getOrUpdateDiskFingerprint(workloadCertPath); + String currentDiskFingerprint = rotationTracker.readDiskFingerprint(); if (currentDiskFingerprint.isEmpty()) { return; } // Double-check fingerprint inside the lock - if (currentDiskFingerprint.equalsIgnoreCase(this.activeCertFingerprint)) { + if (rotationTracker.isAlreadyActive(currentDiskFingerprint)) { LOG.fine( "Channel pool was already refreshed by a concurrent thread, skipping duplicate" + " refresh"); @@ -555,52 +517,80 @@ void refresh() { } if (refreshAll()) { - this.activeCertFingerprint = currentDiskFingerprint; + rotationTracker.markRefreshed(currentDiskFingerprint); } } } + @InternalApi("Visible for testing") + @Nullable String getWorkloadCertPath() { + return workloadCertPath; + } + @InternalApi("Visible for testing") boolean refreshAll() { synchronized (entryWriteLock) { if (isShutdown) { return false; } + String activeFingerprint = rotationTracker.getActiveCertFingerprint(); LOG.fine( "Refreshing all channels" - + (activeCertFingerprint == null + + (Strings.isNullOrEmpty(activeFingerprint) ? "" - : " with certificate fingerprint: " + activeCertFingerprint)); + : " with certificate fingerprint: " + activeFingerprint)); ArrayList newEntries = new ArrayList<>(entries.get()); boolean anyCreated = false; + boolean allCreated = !newEntries.isEmpty(); + List createdEntries = new ArrayList<>(); - for (int i = 0; i < newEntries.size(); i++) { - try { - newEntries.set(i, new Entry(channelFactory.createSingleChannel())); - anyCreated = true; - } catch (IOException e) { - LOG.log(Level.WARNING, "Failed to refresh channel, leaving old channel", e); + try { + for (int i = 0; i < newEntries.size(); i++) { + try { + Entry newEntry = new Entry(channelFactory.createSingleChannel()); + createdEntries.add(newEntry); + newEntries.set(i, newEntry); + anyCreated = true; + } catch (Exception e) { + allCreated = false; + LOG.log(Level.WARNING, "Failed to refresh channel, leaving old channel", e); + } } - } - if (!anyCreated && !newEntries.isEmpty()) { - return false; - } + if (!anyCreated) { + return false; + } - ImmutableList replacedEntries = entries.getAndSet(ImmutableList.copyOf(newEntries)); + ImmutableList replacedEntries = entries.getAndSet(ImmutableList.copyOf(newEntries)); + createdEntries.clear(); // Ownership transferred to pool - // Shutdown the channels that were cycled out. - for (Entry e : replacedEntries) { - if (!newEntries.contains(e)) { + // Shutdown the channels that were cycled out. + for (Entry e : replacedEntries) { + if (!newEntries.contains(e)) { + e.requestShutdown(); + } + } + generation.incrementAndGet(); + return allCreated; + } finally { + // If an Error aborted before getAndSet, shut down newly created channels so they don't leak + for (Entry e : createdEntries) { e.requestShutdown(); } } - generation.incrementAndGet(); - return true; } } - public long getGeneration() { + /** + * Returns the current channel pool generation counter. + * + *

The generation is a monotonically increasing counter incremented each time {@link + * #refreshAll()} replaces the channels in the pool. Retry loops ({@code AttemptCallable} and + * {@code ServerStreamingAttemptCallable}) snapshot the generation before starting an RPC attempt + * and compare it after an {@code UNAUTHENTICATED} failure to determine whether the pool rotated + * to a new certificate generation during or after the attempt. + */ + long getGeneration() { return generation.get(); } @@ -755,8 +745,13 @@ public ClientCall newCall( MethodDescriptor methodDescriptor, CallOptions callOptions) { Entry entry = getRetainedEntry(affinity); - - return new ReleasingClientCall<>(entry.channel.newCall(methodDescriptor, callOptions), entry); + try { + return new ReleasingClientCall<>( + entry.channel.newCall(methodDescriptor, callOptions), entry); + } catch (Throwable t) { + entry.release(); + throw t; + } } } @@ -768,7 +763,8 @@ public ClientCall newCall( * the reference count when {@code start()} is subsequently invoked. */ static class ReleasingClientCall extends SimpleForwardingClientCall { - private @Nullable CancellationException cancellationException; + private final Object callLock = new Object(); + private volatile @Nullable CancellationException cancellationException; final Entry entry; private final AtomicBoolean wasClosed = new AtomicBoolean(); private final AtomicBoolean wasReleased = new AtomicBoolean(); @@ -781,62 +777,80 @@ public ReleasingClientCall(ClientCall delegate, Entry entry) { @Override public void start(Listener responseListener, Metadata headers) { - wasStarted.set(true); - if (cancellationException != null) { - if (wasReleased.compareAndSet(false, true)) { - entry.release(); + synchronized (callLock) { + if (!wasStarted.compareAndSet(false, true)) { + throw new IllegalStateException("Call is already started"); } - throw new IllegalStateException("Call is already cancelled", cancellationException); - } - try { - super.start( - new SimpleForwardingClientCallListener(responseListener) { - @Override - public void onClose(Status status, Metadata trailers) { - if (!wasClosed.compareAndSet(false, true)) { - LOG.log( - Level.WARNING, - "Call is being closed more than once. Please make sure that onClose() is not" - + " being manually called."); - return; - } - try { - super.onClose(status, trailers); - } finally { - if (wasReleased.compareAndSet(false, true)) { - entry.release(); - } else { + if (cancellationException != null) { + if (wasReleased.compareAndSet(false, true)) { + entry.release(); + } + throw new IllegalStateException("Call is already cancelled", cancellationException); + } + try { + super.start( + new SimpleForwardingClientCallListener(responseListener) { + @Override + public void onClose(Status status, Metadata trailers) { + if (!wasClosed.compareAndSet(false, true)) { LOG.log( Level.WARNING, - "Entry was released before the call is closed. This may be due to an" - + " exception on start of the call."); + "Call is being closed more than once. Please make sure that onClose() is" + + " not being manually called."); + return; + } + try { + super.onClose(status, trailers); + } finally { + if (wasReleased.compareAndSet(false, true)) { + entry.release(); + } else { + LOG.log( + Level.WARNING, + "Entry was released before the call is closed. This may be due to an" + + " exception on start of the call."); + } } } - } - }, - headers); - } catch (Exception e) { - // In case start failed, make sure to release - if (wasReleased.compareAndSet(false, true)) { - entry.release(); - } else { - LOG.log( - Level.WARNING, - "The entry is already released. This indicates that onClose() has already been called" - + " previously"); + }, + headers); + } catch (Throwable t) { + // In case start failed, make sure to release + if (wasReleased.compareAndSet(false, true)) { + entry.release(); + } else { + LOG.log( + Level.WARNING, + "The entry is already released. This indicates that onClose() has already been" + + " called previously"); + } + throw t; } - throw e; } } @Override public void cancel(@Nullable String message, @Nullable Throwable cause) { - this.cancellationException = new CancellationException(message); - if (delegate() != null) { - super.cancel(message, cause); - } - if (!wasStarted.get() && wasReleased.compareAndSet(false, true)) { - entry.release(); + boolean releaseImmediately = false; + try { + synchronized (callLock) { + this.cancellationException = new CancellationException(message); + if (!wasStarted.get()) { + releaseImmediately = true; + } + if (delegate() != null) { + super.cancel(message, cause); + } + } + } catch (Throwable t) { + if (!wasStarted.get()) { + releaseImmediately = true; + } + throw t; + } finally { + if (releaseImmediately && wasReleased.compareAndSet(false, true)) { + entry.release(); + } } } } diff --git a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/GrpcCallContext.java b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/GrpcCallContext.java index 428531848b12..c0288a1ae382 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/GrpcCallContext.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/GrpcCallContext.java @@ -663,7 +663,7 @@ public GrpcCallContext withChannel(@Nullable Channel newChannel) { retryableCodes, endpointContext, isDirectPath, - (newChannel == null || newChannel.equals(channel)) ? transportChannel : null); + (newChannel != null && newChannel.equals(channel)) ? transportChannel : null); } /** Returns a new instance with the call options set to the given call options. */ diff --git a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/InstantiatingGrpcChannelProvider.java b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/InstantiatingGrpcChannelProvider.java index 07ffa3b86b5e..8a1c1fc8b0bb 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/InstantiatingGrpcChannelProvider.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/InstantiatingGrpcChannelProvider.java @@ -401,13 +401,19 @@ public TransportChannel getTransportChannel() throws IOException { } private TransportChannel createChannel() throws IOException { + String workloadCertPath = + !this.canUseDirectPath() + && mtlsProvider != null + && certificateBasedAccess.useMtlsClientCertificate() + ? certificateBasedAccess.getWorkloadCertPath() + : null; return GrpcTransportChannel.newBuilder() .setManagedChannel( ChannelPool.create( channelPoolSettings, InstantiatingGrpcChannelProvider.this::createSingleChannel, backgroundExecutor, - certificateBasedAccess.getWorkloadCertPath())) + workloadCertPath)) .setDirectPath(this.canUseDirectPath()) .build(); } @@ -762,6 +768,8 @@ public ManagedChannelBuilder createChannelBuilder() throws IOException { if (channelCredentials != null) { // Create the channel using channel credentials created via DCA. builder = Grpc.newChannelBuilder(endpoint, channelCredentials); + } else if (mtlsProvider != null && certificateBasedAccess.useMtlsClientCertificate()) { + throw new IOException("Failed to initialize mTLS channel credentials"); } else { // Could not create channel credentials via DCA. In accordance with // https://google.aip.dev/auth/4115, if credentials not available through diff --git a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java index 97f27f77513f..4dde16f9aedf 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java @@ -565,6 +565,96 @@ void channelReactiveMTlsRefresh_failedCreationDoesNotMutateFingerprintAndAllowsR .newCall(Mockito.>any(), Mockito.any(CallOptions.class)); } + @Test + void + channelReactiveMTlsRefresh_partialFailureInMultiChannelPool_retainsShouldRefreshAndCompletesOnSubsequentRefresh() + throws IOException { + ManagedChannel initial1 = Mockito.mock(ManagedChannel.class); + ManagedChannel initial2 = Mockito.mock(ManagedChannel.class); + ManagedChannel rotated1 = Mockito.mock(ManagedChannel.class); + ManagedChannel rotated1SecondPass = Mockito.mock(ManagedChannel.class); + ManagedChannel rotated2 = Mockito.mock(ManagedChannel.class); + ChannelFactory channelFactory = Mockito.mock(ChannelFactory.class); + + // Initial creation: initial1, initial2 + // Refresh pass 1: rotated1 succeeds, second throws IOException + // Refresh pass 2: rotated1SecondPass, rotated2 both succeed + Mockito.when(channelFactory.createSingleChannel()) + .thenReturn(initial1, initial2) + .thenReturn(rotated1) + .thenThrow(new IOException("Transient failure on second sub-channel")) + .thenReturn(rotated1SecondPass, rotated2); + + tempCert = java.nio.file.Files.createTempFile("cert", ".pem"); + java.nio.file.Path clientCert = + java.nio.file.Paths.get("src", "test", "resources", "client_cert.pem"); + java.nio.file.Files.copy( + clientCert, tempCert, java.nio.file.StandardCopyOption.REPLACE_EXISTING); + + pool = + ChannelPool.create( + ChannelPoolSettings.staticallySized(2), channelFactory, null, tempCert.toString()); + + // Rotate cert on disk + pool.invalidateDiskFingerprintCache(); + java.nio.file.Path rootCert = + java.nio.file.Paths.get("src", "test", "resources", "root_cert.pem"); + java.nio.file.Files.copy(rootCert, tempCert, java.nio.file.StandardCopyOption.REPLACE_EXISTING); + + assertThat(pool.shouldRefresh()).isTrue(); + long genBefore = pool.getGeneration(); + + // First refresh: partial failure (channel 0 rotates to rotated1, channel 1 fails and keeps + // initial2) + pool.refresh(); + + // Generation should still increment since partial progress was committed + assertThat(pool.getGeneration()).isGreaterThan(genBefore); + // initial1 should have been shut down, initial2 should NOT be shut down yet + Mockito.verify(initial1).shutdown(); + Mockito.verify(initial2, Mockito.never()).shutdown(); + + // Crucial assertion: shouldRefresh() MUST remain true so subsequent 401s on unrotated channel 1 + // trigger retry/refresh + pool.invalidateDiskFingerprintCache(); + assertThat(pool.shouldRefresh()).isTrue(); + + // Second refresh: both channels succeed + pool.refresh(); + + assertThat(pool.shouldRefresh()).isFalse(); + Mockito.verify(initial2).shutdown(); + } + + @Test + void refreshAll_runtimeExceptionOrError_doesNotLeakCreatedChannels() throws IOException { + ManagedChannel initial1 = Mockito.mock(ManagedChannel.class); + ManagedChannel initial2 = Mockito.mock(ManagedChannel.class); + ManagedChannel createdBeforeRuntimeEx = Mockito.mock(ManagedChannel.class); + ManagedChannel createdBeforeError = Mockito.mock(ManagedChannel.class); + ChannelFactory channelFactory = Mockito.mock(ChannelFactory.class); + + Mockito.when(channelFactory.createSingleChannel()) + .thenReturn(initial1, initial2) + .thenReturn(createdBeforeRuntimeEx) + .thenThrow(new RuntimeException("Unchecked runtime exception")) + .thenReturn(createdBeforeError) + .thenThrow(new AssertionError("Simulated Error during refresh")); + + pool = ChannelPool.create(ChannelPoolSettings.staticallySized(2), channelFactory, null, null); + + // Case 1: RuntimeException on channel 1 after creating channel 0 -> caught as Exception, + // partial progress committed + boolean allCreated = pool.refreshAll(); + assertThat(allCreated).isFalse(); + Mockito.verify(initial1).shutdown(); + + // Case 2: Error on channel 1 after creating channel 0 -> aborts, finally block must shut down + // createdBeforeError + org.junit.jupiter.api.Assertions.assertThrows(AssertionError.class, () -> pool.refreshAll()); + Mockito.verify(createdBeforeError).shutdown(); + } + @Test void refresh_onShutdownPool_noOpsAndCreatesNoChannels() throws IOException { ManagedChannel channel1 = mock(ManagedChannel.class); @@ -1138,4 +1228,269 @@ void maxChannelsClampedToMinChannelCountUnderLowLoad() throws Exception { // Should be clamped to minChannelCount = 3 assertThat(pool.entries.get()).hasSize(3); } + + @Test + void shouldRefresh_doesNotCacheNegativeResultAndDetectsSubsequentRotationImmediately() + throws IOException { + ManagedChannel initial = Mockito.mock(ManagedChannel.class); + ManagedChannel rotated = Mockito.mock(ManagedChannel.class); + ChannelFactory channelFactory = Mockito.mock(ChannelFactory.class); + Mockito.when(channelFactory.createSingleChannel()).thenReturn(initial, rotated); + + tempCert = java.nio.file.Files.createTempFile("cert", ".pem"); + java.nio.file.Path clientCert = + java.nio.file.Paths.get("src", "test", "resources", "client_cert.pem"); + java.nio.file.Files.copy( + clientCert, tempCert, java.nio.file.StandardCopyOption.REPLACE_EXISTING); + + pool = + ChannelPool.create( + ChannelPoolSettings.staticallySized(1), channelFactory, null, tempCert.toString()); + + // First check returns false (unchanged disk cert) + assertThat(pool.shouldRefresh()).isFalse(); + + // Immediately rotate cert on disk WITHOUT invalidating the 1-second cache + java.nio.file.Path rootCert = + java.nio.file.Paths.get("src", "test", "resources", "root_cert.pem"); + java.nio.file.Files.copy(rootCert, tempCert, java.nio.file.StandardCopyOption.REPLACE_EXISTING); + + // Must immediately detect rotation because negative/unchanged disk checks are not cached for 1s + assertThat(pool.shouldRefresh()).isTrue(); + + // Refresh should update activeCertFingerprint and clear any cached positive check + pool.refresh(); + assertThat(pool.shouldRefresh()).isFalse(); + } + + @Test + void newCall_whenDelegateThrowsError_releasesEntryAndShutsDownRetiredChannel() + throws IOException { + ManagedChannel initial = Mockito.mock(ManagedChannel.class); + ManagedChannel rotated = Mockito.mock(ManagedChannel.class); + ChannelFactory channelFactory = Mockito.mock(ChannelFactory.class); + Mockito.when(channelFactory.createSingleChannel()).thenReturn(initial, rotated); + Mockito.when(initial.newCall(Mockito.any(), Mockito.any())) + .thenThrow(new LinkageError("Simulated native/JNI linkage error")); + + pool = ChannelPool.create(ChannelPoolSettings.staticallySized(1), channelFactory, null, null); + + assertThrows(LinkageError.class, () -> pool.newCall(METHOD_RECOGNIZE, CallOptions.DEFAULT)); + + // Rotating the pool should immediately shut down initial channel because its ref count is 0 + pool.refresh(); + Mockito.verify(initial).shutdown(); + } + + @Test + @SuppressWarnings("unchecked") + void start_whenDelegateThrowsError_releasesEntryAndShutsDownRetiredChannel() throws IOException { + ManagedChannel initial = Mockito.mock(ManagedChannel.class); + ManagedChannel rotated = Mockito.mock(ManagedChannel.class); + ClientCall mockCall = Mockito.mock(ClientCall.class); + ChannelFactory channelFactory = Mockito.mock(ChannelFactory.class); + Mockito.when(channelFactory.createSingleChannel()).thenReturn(initial, rotated); + Mockito.when(initial.newCall(Mockito.any(), Mockito.any())).thenReturn((ClientCall) mockCall); + Mockito.doThrow(new AssertionError("Simulated Error in start")) + .when(mockCall) + .start(Mockito.any(), Mockito.any()); + + pool = ChannelPool.create(ChannelPoolSettings.staticallySized(1), channelFactory, null, null); + + ClientCall call = pool.newCall(METHOD_RECOGNIZE, CallOptions.DEFAULT); + // Rotate pool while call is retained + pool.refresh(); + Mockito.verify(initial, Mockito.never()).shutdown(); + + // Calling start() throws Error, which must release the retained entry and trigger shutdown + assertThrows( + AssertionError.class, + () -> call.start(new ClientCall.Listener() {}, new io.grpc.Metadata())); + Mockito.verify(initial).shutdown(); + } + + @Test + @SuppressWarnings("unchecked") + void cancel_whenDelegateThrowsException_releasesEntryAndShutsDownRetiredChannel() + throws IOException { + ManagedChannel initial = Mockito.mock(ManagedChannel.class); + ManagedChannel rotated = Mockito.mock(ManagedChannel.class); + ClientCall mockCall = Mockito.mock(ClientCall.class); + ChannelFactory channelFactory = Mockito.mock(ChannelFactory.class); + Mockito.when(channelFactory.createSingleChannel()).thenReturn(initial, rotated); + Mockito.when(initial.newCall(Mockito.any(), Mockito.any())).thenReturn((ClientCall) mockCall); + Mockito.doThrow(new RuntimeException("Simulated cancel exception")) + .when(mockCall) + .cancel(Mockito.any(), Mockito.any()); + + pool = ChannelPool.create(ChannelPoolSettings.staticallySized(1), channelFactory, null, null); + + ClientCall call = pool.newCall(METHOD_RECOGNIZE, CallOptions.DEFAULT); + pool.refresh(); + Mockito.verify(initial, Mockito.never()).shutdown(); + + assertThrows(RuntimeException.class, () -> call.cancel("cancelled", null)); + Mockito.verify(initial).shutdown(); + } + + @Test + @SuppressWarnings("unchecked") + void concurrentStartAndCancel_neverLeaksOrDoubleReleasesEntry() throws Exception { + ManagedChannel initial = Mockito.mock(ManagedChannel.class); + ManagedChannel rotated = Mockito.mock(ManagedChannel.class); + ChannelFactory channelFactory = Mockito.mock(ChannelFactory.class); + Mockito.when(channelFactory.createSingleChannel()).thenReturn(initial, rotated); + + Mockito.when(initial.newCall(Mockito.any(), Mockito.any())) + .thenAnswer( + invocation -> + new ClientCall() { + private Listener listener; + private boolean cancelled; + + @Override + public synchronized void start( + Listener responseListener, io.grpc.Metadata headers) { + this.listener = responseListener; + if (cancelled) { + responseListener.onClose(io.grpc.Status.CANCELLED, new io.grpc.Metadata()); + } + } + + @Override + public synchronized void cancel(String message, Throwable cause) { + cancelled = true; + if (listener != null) { + listener.onClose(io.grpc.Status.CANCELLED, new io.grpc.Metadata()); + } + } + + @Override + public void request(int numMessages) {} + + @Override + public void halfClose() {} + + @Override + public void sendMessage(Color message) {} + }); + + pool = ChannelPool.create(ChannelPoolSettings.staticallySized(1), channelFactory, null, null); + + int iterations = 100; + java.util.concurrent.ExecutorService executor = + java.util.concurrent.Executors.newFixedThreadPool(2); + try { + for (int i = 0; i < iterations; i++) { + ClientCall call = pool.newCall(METHOD_RECOGNIZE, CallOptions.DEFAULT); + java.util.concurrent.CyclicBarrier barrier = new java.util.concurrent.CyclicBarrier(2); + java.util.concurrent.Future f1 = + executor.submit( + () -> { + try { + barrier.await(); + call.start(new ClientCall.Listener() {}, new io.grpc.Metadata()); + } catch (Exception ignored) { + } + }); + java.util.concurrent.Future f2 = + executor.submit( + () -> { + try { + barrier.await(); + call.cancel("cancel", null); + } catch (Exception ignored) { + } + }); + f1.get(5, java.util.concurrent.TimeUnit.SECONDS); + f2.get(5, java.util.concurrent.TimeUnit.SECONDS); + } + } finally { + executor.shutdownNow(); + } + + // Rotate pool: initial channel must shut down cleanly, proving outstandingRpcs == 0 (no leaks + // or negative counts) + pool.refresh(); + Mockito.verify(initial).shutdown(); + } + + @Test + void cancel_whenStartedAndSuperCancelThrows_doesNotReleasePrematurelyUntilOnClose() + throws Exception { + ManagedChannel initial = Mockito.mock(ManagedChannel.class); + ManagedChannel replacement = Mockito.mock(ManagedChannel.class); + @SuppressWarnings("unchecked") + ClientCall delegateCall = Mockito.mock(ClientCall.class); + @SuppressWarnings("unchecked") + ArgumentCaptor> listenerCaptor = + ArgumentCaptor.forClass(ClientCall.Listener.class); + Mockito.doThrow(new RuntimeException("cancel failure")) + .when(delegateCall) + .cancel(Mockito.any(), Mockito.any()); + Mockito.when(initial.newCall(Mockito.eq(METHOD_RECOGNIZE), Mockito.any())) + .thenReturn(delegateCall); + + java.util.concurrent.atomic.AtomicInteger createCount = + new java.util.concurrent.atomic.AtomicInteger(0); + pool = + new ChannelPool( + ChannelPoolSettings.staticallySized(1), + () -> createCount.getAndIncrement() == 0 ? initial : replacement, + FixedExecutorProvider.create(Mockito.mock(ScheduledExecutorService.class)), + null); + + ClientCall call = pool.newCall(METHOD_RECOGNIZE, CallOptions.DEFAULT); + call.start(new ClientCall.Listener() {}, new Metadata()); + Mockito.verify(delegateCall).start(listenerCaptor.capture(), Mockito.any()); + + assertThrows(RuntimeException.class, () -> call.cancel("abort", null)); + + // Rotate pool while call is still active (onClose hasn't fired yet): + // initial channel must NOT be shut down yet because call is still active + pool.refreshAll(); + Mockito.verify(initial, Mockito.never()).shutdown(); + + // Once onClose fires, entry is released and initial channel shuts down + listenerCaptor.getValue().onClose(Status.CANCELLED, new Metadata()); + Mockito.verify(initial).shutdown(); + } + + @Test + void start_whenCalledTwice_throwsIllegalStateExceptionAndDoesNotReleaseFirstCallEntry() + throws Exception { + ManagedChannel initial = Mockito.mock(ManagedChannel.class); + ManagedChannel replacement = Mockito.mock(ManagedChannel.class); + @SuppressWarnings("unchecked") + ClientCall delegateCall = Mockito.mock(ClientCall.class); + @SuppressWarnings("unchecked") + ArgumentCaptor> listenerCaptor = + ArgumentCaptor.forClass(ClientCall.Listener.class); + Mockito.when(initial.newCall(Mockito.eq(METHOD_RECOGNIZE), Mockito.any())) + .thenReturn(delegateCall); + + java.util.concurrent.atomic.AtomicInteger createCount = + new java.util.concurrent.atomic.AtomicInteger(0); + pool = + new ChannelPool( + ChannelPoolSettings.staticallySized(1), + () -> createCount.getAndIncrement() == 0 ? initial : replacement, + FixedExecutorProvider.create(Mockito.mock(ScheduledExecutorService.class)), + null); + + ClientCall call = pool.newCall(METHOD_RECOGNIZE, CallOptions.DEFAULT); + call.start(new ClientCall.Listener() {}, new Metadata()); + Mockito.verify(delegateCall).start(listenerCaptor.capture(), Mockito.any()); + + // Duplicate start() must throw IllegalStateException without releasing the entry + assertThrows( + IllegalStateException.class, + () -> call.start(new ClientCall.Listener() {}, new Metadata())); + + pool.refreshAll(); + Mockito.verify(initial, Mockito.never()).shutdown(); + + listenerCaptor.getValue().onClose(Status.OK, new Metadata()); + Mockito.verify(initial).shutdown(); + } } diff --git a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/GrpcCallContextTest.java b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/GrpcCallContextTest.java index 59d5bbf568be..cbaa7af2475c 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/GrpcCallContextTest.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/GrpcCallContextTest.java @@ -543,5 +543,15 @@ public void testWithChannelWithCustomChannelClearsTransportChannel() { assertEquals(customChannel, updatedContext.getChannel()); assertNull(updatedContext.getTransportChannel()); + + // Clearing channel via withChannel(null) also clears transportChannel + GrpcCallContext nullChannelContext = baseContext.withChannel(null); + assertNull(nullChannelContext.getChannel()); + assertNull(nullChannelContext.getTransportChannel()); + + // Merging a cleared context into defaultContext falls back to defaultContext's transportChannel + GrpcCallContext mergedWithNullChannel = (GrpcCallContext) baseContext.merge(nullChannelContext); + assertEquals(defaultChannel, mergedWithNullChannel.getChannel()); + assertEquals(transportChannel, mergedWithNullChannel.getTransportChannel()); } } diff --git a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/InstantiatingGrpcChannelProviderTest.java b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/InstantiatingGrpcChannelProviderTest.java index c2127aede3b8..d083c01ad845 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/InstantiatingGrpcChannelProviderTest.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/InstantiatingGrpcChannelProviderTest.java @@ -48,6 +48,7 @@ import com.google.api.gax.rpc.internal.EnvironmentProvider; import com.google.api.gax.rpc.mtls.AbstractMtlsTransportChannelTest; import com.google.api.gax.rpc.mtls.CertificateBasedAccess; +import com.google.api.gax.rpc.testing.FakeMtlsProvider; import com.google.auth.ApiKeyCredentials; import com.google.auth.Credentials; import com.google.auth.http.AuthHttpConstants; @@ -1353,6 +1354,96 @@ void testSettingBackgroundExecutor() { assertThat(provider.getBackgroundExecutor()).isEqualTo(mockExecutor); } + @Test + void createChannel_whenDirectPathEnabled_ignoresWorkloadCertPath() throws Exception { + System.setProperty("os.name", "Linux"); + EnvironmentProvider envProvider = + mock(EnvironmentProvider.class, Mockito.withSettings().withoutAnnotations()); + Mockito.when( + envProvider.getenv( + InstantiatingGrpcChannelProvider.DIRECT_PATH_ENV_DISABLE_DIRECT_PATH)) + .thenReturn("false"); + CertificateBasedAccess mtlsCertificateBasedAccess = + mock(CertificateBasedAccess.class, Mockito.withSettings().withoutAnnotations()); + Mockito.when(mtlsCertificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); + Mockito.when(mtlsCertificateBasedAccess.getWorkloadCertPath()) + .thenReturn("/path/to/workload/cert.pem"); + MtlsProvider mtlsProvider = + new FakeMtlsProvider(FakeMtlsProvider.createTestMtlsKeyStore(), "", false); + + InstantiatingGrpcChannelProvider.Builder builder = + InstantiatingGrpcChannelProvider.newBuilder() + .setCertificateBasedAccess(mtlsCertificateBasedAccess) + .setMtlsProvider(mtlsProvider) + .setAttemptDirectPath(true) + .setCredentials(computeEngineCredentials) + .setEndpoint(DEFAULT_ENDPOINT) + .setEnvProvider(envProvider) + .setHeaderProvider( + mock(HeaderProvider.class, Mockito.withSettings().withoutAnnotations())); + InstantiatingGrpcChannelProvider provider = + new InstantiatingGrpcChannelProvider(builder, GCE_PRODUCTION_NAME_AFTER_2016); + Truth.assertThat(provider.canUseDirectPath()).isTrue(); + + TransportChannel transportChannel = provider.getTransportChannel(); + try { + ChannelPool pool = (ChannelPool) ((GrpcTransportChannel) transportChannel).getChannel(); + assertThat(pool.getWorkloadCertPath()).isNull(); + } finally { + transportChannel.shutdownNow(); + } + } + + @Test + void createChannel_whenMtlsActive_passesWorkloadCertPathToChannelPool() throws Exception { + CertificateBasedAccess mtlsCertificateBasedAccess = + mock(CertificateBasedAccess.class, Mockito.withSettings().withoutAnnotations()); + Mockito.when(mtlsCertificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); + Mockito.when(mtlsCertificateBasedAccess.getWorkloadCertPath()) + .thenReturn("/path/to/workload/cert.pem"); + MtlsProvider mtlsProvider = + new FakeMtlsProvider(FakeMtlsProvider.createTestMtlsKeyStore(), "", false); + + InstantiatingGrpcChannelProvider provider = + InstantiatingGrpcChannelProvider.newBuilder() + .setCertificateBasedAccess(mtlsCertificateBasedAccess) + .setMtlsProvider(mtlsProvider) + .setAttemptDirectPath(false) + .setEndpoint(DEFAULT_ENDPOINT) + .setHeaderProvider( + mock(HeaderProvider.class, Mockito.withSettings().withoutAnnotations())) + .build(); + + TransportChannel transportChannel = provider.getTransportChannel(); + try { + ChannelPool pool = (ChannelPool) ((GrpcTransportChannel) transportChannel).getChannel(); + assertThat(pool.getWorkloadCertPath()).isEqualTo("/path/to/workload/cert.pem"); + } finally { + transportChannel.shutdownNow(); + } + } + + @Test + void createChannelBuilder_whenMtlsActiveAndCredentialsNull_throwsIOException() { + CertificateBasedAccess mtlsCertificateBasedAccess = + mock(CertificateBasedAccess.class, Mockito.withSettings().withoutAnnotations()); + Mockito.when(mtlsCertificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); + MtlsProvider mtlsProviderWithNullKeyStore = new FakeMtlsProvider(null, "", false); + + InstantiatingGrpcChannelProvider provider = + InstantiatingGrpcChannelProvider.newBuilder() + .setCertificateBasedAccess(mtlsCertificateBasedAccess) + .setMtlsProvider(mtlsProviderWithNullKeyStore) + .setAttemptDirectPath(false) + .setEndpoint(DEFAULT_ENDPOINT) + .setHeaderProvider( + mock(HeaderProvider.class, Mockito.withSettings().withoutAnnotations())) + .build(); + + IOException thrown = assertThrows(IOException.class, provider::createChannelBuilder); + assertThat(thrown).hasMessageThat().contains("Failed to initialize mTLS channel credentials"); + } + private static class FakeLogHandler extends Handler { List records = new ArrayList<>(); diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/HttpJsonCallContext.java b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/HttpJsonCallContext.java index 47e3d31ee6e7..1944b6c9a6b3 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/HttpJsonCallContext.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/HttpJsonCallContext.java @@ -601,7 +601,7 @@ public HttpJsonCallContext withChannel(@Nullable HttpJsonChannel newChannel) { this.retrySettings, this.retryableCodes, this.endpointContext, - (newChannel == null || newChannel.equals(this.channel)) ? this.transportChannel : null); + (newChannel != null && newChannel.equals(this.channel)) ? this.transportChannel : null); } public HttpJsonCallContext withCallOptions(HttpJsonCallOptions newCallOptions) { diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProvider.java b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProvider.java index 01a1b9b6b625..730ae80bb8b9 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProvider.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProvider.java @@ -208,44 +208,65 @@ public TransportChannelProvider withCredentials(Credentials credentials) { return null; } + private ManagedHttpJsonChannel createSingleManagedChannel() + throws IOException, GeneralSecurityException { + HttpTransport httpTransportToUse = httpTransport; + if (httpTransportToUse == null) { + httpTransportToUse = createHttpTransport(); + if (httpTransportToUse == null + && mtlsProvider != null + && certificateBasedAccess.useMtlsClientCertificate()) { + throw new IOException("Failed to initialize mTLS HttpTransport"); + } + } + return ManagedHttpJsonChannel.newBuilder() + .setEndpoint(endpoint) + .setExecutor(executor) + .setHttpTransport(httpTransportToUse) + .setManageHttpTransport(httpTransport == null) + .build(); + } + private HttpJsonTransportChannel createChannel() throws IOException, GeneralSecurityException { - java.util.function.Supplier channelFactory = - () -> { - try { - HttpTransport httpTransportToUse = httpTransport; - if (httpTransportToUse == null) { - httpTransportToUse = createHttpTransport(); + boolean isMtlsActive = + httpTransport == null + && mtlsProvider != null + && certificateBasedAccess.useMtlsClientCertificate(); + String workloadCertPath = isMtlsActive ? certificateBasedAccess.getWorkloadCertPath() : null; + + ManagedHttpJsonChannel initialChannel = createSingleManagedChannel(); + try { + java.util.function.Supplier channelFactory = + () -> { + try { + return createSingleManagedChannel(); + } catch (Exception e) { + throw new java.lang.RuntimeException( + "Failed to create fresh ManagedHttpJsonChannel", e); } - return ManagedHttpJsonChannel.newBuilder() - .setEndpoint(endpoint) - .setExecutor(executor) - .setHttpTransport(httpTransportToUse) - .setManageHttpTransport(httpTransport == null) - .build(); - } catch (Exception e) { - throw new java.lang.RuntimeException( - "Failed to create fresh ManagedHttpJsonChannel", e); - } - }; - - String workloadCertPath = certificateBasedAccess.getWorkloadCertPath(); - ManagedHttpJsonChannel channel = - workloadCertPath != null - ? new RefreshingHttpJsonChannel(channelFactory, workloadCertPath) - : channelFactory.get(); - - HttpJsonClientInterceptor headerInterceptor = - new HttpJsonHeaderInterceptor(headerProvider.getHeaders()); - - channel = new ManagedHttpJsonInterceptorChannel(channel, new HttpJsonLoggingInterceptor()); - channel = new ManagedHttpJsonInterceptorChannel(channel, headerInterceptor); - if (interceptorProvider != null && interceptorProvider.getInterceptors() != null) { - for (HttpJsonClientInterceptor interceptor : interceptorProvider.getInterceptors()) { - channel = new ManagedHttpJsonInterceptorChannel(channel, interceptor); + }; + + ManagedHttpJsonChannel channel = + workloadCertPath != null + ? new RefreshingHttpJsonChannel(initialChannel, channelFactory, workloadCertPath) + : initialChannel; + + HttpJsonClientInterceptor headerInterceptor = + new HttpJsonHeaderInterceptor(headerProvider.getHeaders()); + + channel = new ManagedHttpJsonInterceptorChannel(channel, new HttpJsonLoggingInterceptor()); + channel = new ManagedHttpJsonInterceptorChannel(channel, headerInterceptor); + if (interceptorProvider != null && interceptorProvider.getInterceptors() != null) { + for (HttpJsonClientInterceptor interceptor : interceptorProvider.getInterceptors()) { + channel = new ManagedHttpJsonInterceptorChannel(channel, interceptor); + } } - } - return HttpJsonTransportChannel.newBuilder().setManagedChannel(channel).build(); + return HttpJsonTransportChannel.newBuilder().setManagedChannel(channel).build(); + } catch (Throwable t) { + initialChannel.shutdownNow(); + throw t; + } } /** The endpoint to be used for the channel. */ diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java index 2fec3ba02253..3d9557152845 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java @@ -33,12 +33,15 @@ import com.google.api.core.InternalApi; import com.google.api.gax.httpjson.ForwardingHttpJsonClientCall.SimpleForwardingHttpJsonClientCall; import com.google.api.gax.httpjson.ForwardingHttpJsonClientCallListener.SimpleForwardingHttpJsonClientCallListener; +import com.google.api.gax.rpc.mtls.CertificateRotationTracker; import com.google.api.gax.rpc.mtls.WorkloadCertificateUtils; import com.google.common.annotations.VisibleForTesting; import java.util.concurrent.CancellationException; +import java.util.concurrent.ConcurrentLinkedQueue; import java.util.concurrent.TimeUnit; import java.util.concurrent.atomic.AtomicBoolean; import java.util.concurrent.atomic.AtomicInteger; +import java.util.concurrent.atomic.AtomicLong; import java.util.concurrent.atomic.AtomicReference; import java.util.function.Supplier; import java.util.logging.Level; @@ -55,87 +58,53 @@ public class RefreshingHttpJsonChannel extends ManagedHttpJsonChannel { private static final Logger LOG = Logger.getLogger(RefreshingHttpJsonChannel.class.getName()); - private static class DiskCheckResult { - final String fingerprint; - final long timestampNanos; - - DiskCheckResult(String fingerprint, long timestampNanos) { - this.fingerprint = fingerprint; - this.timestampNanos = timestampNanos; - } - } - - private volatile DiskCheckResult lastDiskCheck = null; - private final java.util.concurrent.locks.ReentrantLock diskCheckLock = - new java.util.concurrent.locks.ReentrantLock(); + private final CertificateRotationTracker rotationTracker; private final Supplier channelFactory; private final String workloadCertPath; private final AtomicReference activeEntry; // Keep track of all entries to properly await their termination - private final java.util.concurrent.ConcurrentLinkedQueue allEntries = - new java.util.concurrent.ConcurrentLinkedQueue<>(); + private final ConcurrentLinkedQueue allEntries = new ConcurrentLinkedQueue<>(); private final Object refreshLock = new Object(); - private final java.util.concurrent.atomic.AtomicLong generation = - new java.util.concurrent.atomic.AtomicLong(0); - private volatile String activeCertFingerprint = ""; + private final AtomicLong generation = new AtomicLong(0); public RefreshingHttpJsonChannel( Supplier channelFactory, String workloadCertPath) { + this(channelFactory.get(), channelFactory, workloadCertPath); + } + + public RefreshingHttpJsonChannel( + ManagedHttpJsonChannel initialChannel, + Supplier channelFactory, + String workloadCertPath) { super(true); this.channelFactory = channelFactory; this.workloadCertPath = workloadCertPath; - ChannelEntry initial = new ChannelEntry(channelFactory.get()); + ChannelEntry initial = new ChannelEntry(initialChannel); this.activeEntry = new AtomicReference<>(initial); this.allEntries.add(initial); - if (workloadCertPath != null) { - this.activeCertFingerprint = getCertificateFingerprint(workloadCertPath); - } - } - - private String getOrUpdateDiskFingerprint(String certPath) { - long now = System.nanoTime(); - DiskCheckResult cached = lastDiskCheck; - if (cached != null - && (now - cached.timestampNanos < java.util.concurrent.TimeUnit.SECONDS.toNanos(1))) { - return cached.fingerprint; - } - - diskCheckLock.lock(); try { - cached = lastDiskCheck; - if (cached != null - && (now - cached.timestampNanos < java.util.concurrent.TimeUnit.SECONDS.toNanos(1))) { - return cached.fingerprint; - } - String fingerprint = getCertificateFingerprint(certPath); - lastDiskCheck = new DiskCheckResult(fingerprint, System.nanoTime()); - return fingerprint; - } finally { - diskCheckLock.unlock(); + this.rotationTracker = + new CertificateRotationTracker( + this::getWorkloadCertPath, this::getCertificateFingerprint); + } catch (Throwable t) { + initialChannel.shutdownNow(); + throw t; } } // Visible for testing - protected String getWorkloadCertPath() { + String getWorkloadCertPath() { return workloadCertPath; } // Visible for testing - protected String getCertificateFingerprint(String certPath) { + String getCertificateFingerprint(String certPath) { return WorkloadCertificateUtils.getCertificateFingerprint(certPath); } @Override public boolean shouldRefresh() { - String certPath = getWorkloadCertPath(); - if (certPath == null) { - return false; - } - String currentDiskFingerprint = getOrUpdateDiskFingerprint(certPath); - if (currentDiskFingerprint.isEmpty()) { - return false; - } - return !currentDiskFingerprint.equalsIgnoreCase(activeCertFingerprint); + return rotationTracker.shouldRefresh(); } @Override @@ -144,17 +113,13 @@ public void refresh() { if (isShutdown()) { return; } - String certPath = getWorkloadCertPath(); - if (certPath == null) { - return; - } - String currentDiskFingerprint = getOrUpdateDiskFingerprint(certPath); + String currentDiskFingerprint = rotationTracker.readDiskFingerprint(); if (currentDiskFingerprint.isEmpty()) { return; } // Double-check inside refreshLock - if (currentDiskFingerprint.equalsIgnoreCase(this.activeCertFingerprint)) { + if (rotationTracker.isAlreadyActive(currentDiskFingerprint)) { LOG.fine( "HTTP/JSON channel was already refreshed by a concurrent thread, skipping duplicate" + " refresh"); @@ -163,13 +128,13 @@ public void refresh() { LOG.info("mTLS certificate rotation detected. Triggering HTTP/JSON channel pool refresh."); - // Prune terminated entries to prevent memory leak - allEntries.removeIf(entry -> entry.channel.isTerminated()); - ChannelEntry newEntry = new ChannelEntry(channelFactory.get()); allEntries.add(newEntry); + // Prune terminated entries after adding newEntry to ensure allEntries is never empty + allEntries.removeIf(entry -> entry != newEntry && entry.channel.isTerminated()); + ChannelEntry oldEntry = activeEntry.getAndSet(newEntry); - this.activeCertFingerprint = currentDiskFingerprint; + rotationTracker.markRefreshed(currentDiskFingerprint); generation.incrementAndGet(); if (oldEntry != null) { @@ -203,9 +168,9 @@ public HttpJsonClientCall newCall( HttpJsonClientCall delegateCall = entry.channel.newCall(methodDescriptor, callOptions); return new ReleasingHttpJsonClientCall<>(delegateCall, entry); - } catch (Exception e) { + } catch (Throwable t) { entry.release(); - throw e; + throw t; } } @@ -238,6 +203,9 @@ public boolean isShutdown() { @Override public boolean isTerminated() { + if (!isShuttingDown) { + return false; + } for (ChannelEntry entry : allEntries) { if (!entry.channel.isTerminated()) { return false; @@ -260,7 +228,7 @@ public void shutdownNow() { @VisibleForTesting void invalidateDiskFingerprintCache() { - this.lastDiskCheck = null; + rotationTracker.invalidateCache(); } @Override @@ -274,7 +242,8 @@ public boolean awaitTermination(long duration, TimeUnit unit) throws Interrupted if (remainingNanos <= 0) { return false; } - if (!entry.channel.awaitTermination(remainingNanos, TimeUnit.NANOSECONDS)) { + if (!entry.channel.awaitTermination(remainingNanos, TimeUnit.NANOSECONDS) + && !entry.channel.isTerminated()) { return false; } } @@ -319,7 +288,12 @@ boolean retain() { void release() { int count = outstandingCalls.decrementAndGet(); - if (shutdownRequested.get() && count == 0) { + if (count < 0) { + LOG.warning("Channel entry reference count dropped below 0"); + } + // Must check outstandingCalls after shutdownRequested (in reverse order of retain()) to + // ensure mutual exclusion. + if (shutdownRequested.get() && outstandingCalls.get() == 0) { shutdown(); } } @@ -346,7 +320,8 @@ private void shutdown() { private static class ReleasingHttpJsonClientCall extends SimpleForwardingHttpJsonClientCall { - private @Nullable CancellationException cancellationException; + private final Object callLock = new Object(); + private volatile @Nullable CancellationException cancellationException; private final ChannelEntry entry; private final AtomicBoolean wasClosed = new AtomicBoolean(false); private final AtomicBoolean wasReleased = new AtomicBoolean(false); @@ -359,45 +334,65 @@ private static class ReleasingHttpJsonClientCall @Override public void start(Listener responseListener, HttpJsonMetadata requestHeaders) { - wasStarted.set(true); - if (cancellationException != null) { - if (wasReleased.compareAndSet(false, true)) { - entry.release(); + synchronized (callLock) { + if (!wasStarted.compareAndSet(false, true)) { + throw new IllegalStateException("Call is already started"); } - throw new IllegalStateException("Call is already cancelled", cancellationException); - } - try { - super.start( - new SimpleForwardingHttpJsonClientCallListener(responseListener) { - @Override - public void onClose(int statusCode, HttpJsonMetadata trailers) { - if (!wasClosed.compareAndSet(false, true)) { - return; - } - try { - super.onClose(statusCode, trailers); - } finally { - if (wasReleased.compareAndSet(false, true)) { - entry.release(); + if (cancellationException != null) { + if (wasReleased.compareAndSet(false, true)) { + entry.release(); + } + throw new IllegalStateException("Call is already cancelled", cancellationException); + } + try { + super.start( + new SimpleForwardingHttpJsonClientCallListener(responseListener) { + @Override + public void onClose(int statusCode, HttpJsonMetadata trailers) { + if (!wasClosed.compareAndSet(false, true)) { + return; + } + try { + super.onClose(statusCode, trailers); + } finally { + if (wasReleased.compareAndSet(false, true)) { + entry.release(); + } } } - } - }, - requestHeaders); - } catch (Exception e) { - if (wasReleased.compareAndSet(false, true)) { - entry.release(); + }, + requestHeaders); + } catch (Throwable t) { + if (wasReleased.compareAndSet(false, true)) { + entry.release(); + } + throw t; } - throw e; } } @Override public void cancel(@Nullable String message, @Nullable Throwable cause) { - this.cancellationException = new CancellationException(message); - super.cancel(message, cause); - if (!wasStarted.get() && wasReleased.compareAndSet(false, true)) { - entry.release(); + boolean releaseImmediately = false; + try { + synchronized (callLock) { + this.cancellationException = new CancellationException(message); + if (!wasStarted.get()) { + releaseImmediately = true; + } + if (delegate() != null) { + super.cancel(message, cause); + } + } + } catch (Throwable t) { + if (!wasStarted.get()) { + releaseImmediately = true; + } + throw t; + } finally { + if (releaseImmediately && wasReleased.compareAndSet(false, true)) { + entry.release(); + } } } } diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/HttpJsonCallContextTest.java b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/HttpJsonCallContextTest.java index 15b3df3b86ce..7ca2de0775c0 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/HttpJsonCallContextTest.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/HttpJsonCallContextTest.java @@ -349,14 +349,21 @@ void testWithChannelClearsStaleTransportChannel() { HttpJsonCallContext.createDefault().withTransportChannel(transportChannel1); Truth.assertThat(context.getTransportChannel()).isSameInstanceAs(transportChannel1); - // Retains transportChannel when setting same channel or null + // Retains transportChannel when setting same channel Truth.assertThat(context.withChannel(channel1).getTransportChannel()) .isSameInstanceAs(transportChannel1); - Truth.assertThat(context.withChannel(null).getTransportChannel()) - .isSameInstanceAs(transportChannel1); - // Clears transportChannel to null when setting a different channel + // Clears transportChannel to null when setting null or a different channel + HttpJsonCallContext nullChannelContext = context.withChannel(null); + Truth.assertThat(nullChannelContext.getChannel()).isNull(); + Truth.assertThat(nullChannelContext.getTransportChannel()).isNull(); Truth.assertThat(context.withChannel(channel2).getTransportChannel()).isNull(); + + // Merging a cleared context into default context preserves default context's transportChannel + HttpJsonCallContext mergedWithNullChannel = context.merge(nullChannelContext); + Truth.assertThat(mergedWithNullChannel.getChannel()).isSameInstanceAs(channel1); + Truth.assertThat(mergedWithNullChannel.getTransportChannel()) + .isSameInstanceAs(transportChannel1); } @Test diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProviderTest.java b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProviderTest.java index 59fe17f1621c..833c5ff6f5e8 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProviderTest.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProviderTest.java @@ -216,12 +216,17 @@ void managedChannelDoesNotShutdownCustomHttpTransport() throws IOException { @Test void channelCreation_withWorkloadCertPath_wrapsWithRefreshingHttpJsonChannel() - throws IOException { + throws IOException, GeneralSecurityException { + Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); Mockito.when(certificateBasedAccess.getWorkloadCertPath()).thenReturn("fake/cert/path.json"); + com.google.auth.mtls.MtlsProvider mtlsProvider = + new com.google.api.gax.rpc.testing.FakeMtlsProvider( + com.google.api.gax.rpc.testing.FakeMtlsProvider.createTestMtlsKeyStore(), "", false); InstantiatingHttpJsonChannelProvider provider = InstantiatingHttpJsonChannelProvider.newBuilder() .setEndpoint(DEFAULT_ENDPOINT) + .setMtlsProvider(mtlsProvider) .setCertificateBasedAccess(certificateBasedAccess) .build(); provider = (InstantiatingHttpJsonChannelProvider) provider.withHeaders(DEFAULT_HEADER_MAP); @@ -237,6 +242,79 @@ void channelCreation_withWorkloadCertPath_wrapsWithRefreshingHttpJsonChannel() provider.getTransportChannel().shutdownNow(); } + @Test + void channelCreation_withCustomHttpTransport_ignoresWorkloadCertPathAndDoesNotWrap() + throws IOException { + com.google.api.client.http.HttpTransport mockHttpTransport = + org.mockito.Mockito.mock(com.google.api.client.http.HttpTransport.class); + + InstantiatingHttpJsonChannelProvider provider = + InstantiatingHttpJsonChannelProvider.newBuilder() + .setEndpoint(DEFAULT_ENDPOINT) + .setHttpTransport(mockHttpTransport) + .setCertificateBasedAccess(certificateBasedAccess) + .build(); + provider = (InstantiatingHttpJsonChannelProvider) provider.withHeaders(DEFAULT_HEADER_MAP); + + HttpJsonTransportChannel httpJsonTransportChannel = provider.getTransportChannel(); + + ManagedHttpJsonInterceptorChannel interceptorChannel = + (ManagedHttpJsonInterceptorChannel) httpJsonTransportChannel.getManagedChannel(); + ManagedHttpJsonInterceptorChannel managedHttpJsonChannel = + (ManagedHttpJsonInterceptorChannel) interceptorChannel.getChannel(); + assertThat(managedHttpJsonChannel.getChannel()) + .isNotInstanceOf(RefreshingHttpJsonChannel.class); + Mockito.verify(certificateBasedAccess, Mockito.never()).getWorkloadCertPath(); + + httpJsonTransportChannel.shutdownNow(); + } + + @Test + void getTransportChannel_whenMtlsKeyStoreThrowsIOException_throwsCheckedIOException() + throws Exception { + Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); + Mockito.when(certificateBasedAccess.getWorkloadCertPath()).thenReturn("fake/cert/path.json"); + MtlsProvider failingMtlsProvider = Mockito.mock(MtlsProvider.class); + Mockito.when(failingMtlsProvider.getKeyStore()) + .thenThrow(new IOException("Simulated keystore read failure")); + + InstantiatingHttpJsonChannelProvider provider = + InstantiatingHttpJsonChannelProvider.newBuilder() + .setEndpoint(DEFAULT_ENDPOINT) + .setMtlsProvider(failingMtlsProvider) + .setCertificateBasedAccess(certificateBasedAccess) + .build(); + final InstantiatingHttpJsonChannelProvider finalProvider = + (InstantiatingHttpJsonChannelProvider) provider.withHeaders(DEFAULT_HEADER_MAP); + + // Must throw checked IOException directly (not wrapped in RuntimeException) + IOException thrown = + org.junit.jupiter.api.Assertions.assertThrows( + IOException.class, finalProvider::getTransportChannel); + assertThat(thrown).hasMessageThat().contains("Simulated keystore read failure"); + } + + @Test + void getTransportChannel_whenMtlsActiveAndKeyStoreNull_throwsIOException() { + Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); + com.google.auth.mtls.MtlsProvider providerWithNullKeyStore = + new com.google.api.gax.rpc.testing.FakeMtlsProvider(null, "", false); + + InstantiatingHttpJsonChannelProvider provider = + InstantiatingHttpJsonChannelProvider.newBuilder() + .setEndpoint(DEFAULT_ENDPOINT) + .setMtlsProvider(providerWithNullKeyStore) + .setCertificateBasedAccess(certificateBasedAccess) + .build(); + InstantiatingHttpJsonChannelProvider finalProvider = + (InstantiatingHttpJsonChannelProvider) provider.withHeaders(DEFAULT_HEADER_MAP); + + IOException thrown = + org.junit.jupiter.api.Assertions.assertThrows( + IOException.class, finalProvider::getTransportChannel); + assertThat(thrown).hasMessageThat().contains("Failed to initialize mTLS HttpTransport"); + } + @Test void createHttpTransport_withMtlsAndConscrypt_configuresSecurityProvider() throws IOException, GeneralSecurityException { diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java index 2ac955592c9d..aaaacd3c53db 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java @@ -48,7 +48,7 @@ class RefreshingHttpJsonChannelTest { private static class FakeHttpJsonClientCall extends HttpJsonClientCall { - private Listener listener; + protected Listener listener; @Override public void start(Listener responseListener, HttpJsonMetadata requestHeaders) { @@ -153,12 +153,12 @@ private RefreshingHttpJsonChannel createTestChannel() { RefreshingHttpJsonChannel ch = new RefreshingHttpJsonChannel(channelFactory, "fake/cert/path.json") { @Override - protected String getWorkloadCertPath() { + String getWorkloadCertPath() { return testCertPath; } @Override - protected String getCertificateFingerprint(String certPath) { + String getCertificateFingerprint(String certPath) { return testFingerprint; } }; @@ -193,6 +193,24 @@ void testShouldRefreshTrueWhenChanged() throws InterruptedException { assertTrue(channel.shouldRefresh()); } + @Test + void shouldRefresh_doesNotCacheNegativeResultAndDetectsSubsequentRotationImmediately() { + RefreshingHttpJsonChannel channel = createTestChannel(); + + // First check returns false (unchanged fingerprint) + assertFalse(channel.shouldRefresh()); + + // Immediately change fingerprint WITHOUT invalidating cache + testFingerprint = "fingerprint2"; + + // Must immediately detect rotation because negative/unchanged checks are not cached for 1s + assertTrue(channel.shouldRefresh()); + + // Refresh updates activeCertFingerprint and clears cache + channel.refresh(); + assertFalse(channel.shouldRefresh()); + } + @Test void testRefreshSwapsChannel() throws InterruptedException { RefreshingHttpJsonChannel channel = createTestChannel(); @@ -410,4 +428,194 @@ void testGenerationIncrementAndLifecycleOnDelegatingWrapper() throws Exception { assertTrue(channel.isTerminated()); assertTrue(channel.awaitTermination(1, TimeUnit.SECONDS)); } + + @Test + void testNewCall_whenDelegateThrowsError_releasesEntryAndShutsDownRetiredChannel() { + RefreshingHttpJsonChannel channel = createTestChannel(); + FakeManagedHttpJsonChannel firstChannel = lastCreatedChannel; + firstChannel.nextCall = null; + // Configure firstChannel to throw an Error on newCall + FakeManagedHttpJsonChannel throwingChannel = + new FakeManagedHttpJsonChannel() { + @Override + public HttpJsonClientCall newCall( + ApiMethodDescriptor methodDescriptor, + HttpJsonCallOptions callOptions) { + throw new LinkageError("Simulated native error in newCall"); + } + }; + channelFactory = + () -> { + channelFactoryCount.incrementAndGet(); + lastCreatedChannel = throwingChannel; + return throwingChannel; + }; + RefreshingHttpJsonChannel testChannel = createTestChannel(); + + assertThrows(LinkageError.class, () -> testChannel.newCall(null, null)); + + // Refresh should immediately shut down throwingChannel since ref count returned to 0 + testFingerprint = "fingerprint2"; + testChannel.refresh(); + assertTrue(throwingChannel.isShutdown()); + } + + @Test + void testStart_whenDelegateThrowsError_releasesEntryAndShutsDownRetiredChannel() { + RefreshingHttpJsonChannel channel = createTestChannel(); + FakeManagedHttpJsonChannel firstChannel = lastCreatedChannel; + firstChannel.nextCall = + new FakeHttpJsonClientCall() { + @Override + public void start(Listener responseListener, HttpJsonMetadata requestHeaders) { + throw new AssertionError("Simulated Error in start"); + } + }; + + HttpJsonClientCall call = channel.newCall(null, null); + testFingerprint = "fingerprint2"; + channel.refresh(); + assertFalse(firstChannel.isShutdown()); + + assertThrows( + AssertionError.class, () -> call.start(new HttpJsonClientCall.Listener() {}, null)); + assertTrue(firstChannel.isShutdown()); + } + + @Test + void testCancel_whenDelegateThrowsException_releasesEntryAndShutsDownRetiredChannel() { + RefreshingHttpJsonChannel channel = createTestChannel(); + FakeManagedHttpJsonChannel firstChannel = lastCreatedChannel; + firstChannel.nextCall = + new FakeHttpJsonClientCall() { + @Override + public void cancel(String message, Throwable cause) { + throw new RuntimeException("Simulated cancel failure"); + } + }; + + HttpJsonClientCall call = channel.newCall(null, null); + testFingerprint = "fingerprint2"; + channel.refresh(); + assertFalse(firstChannel.isShutdown()); + + assertThrows(RuntimeException.class, () -> call.cancel("cancel", null)); + assertTrue(firstChannel.isShutdown()); + } + + @Test + void testConcurrentStartAndCancel_neverLeaksOrDoubleReleasesEntry() throws Exception { + RefreshingHttpJsonChannel channel = createTestChannel(); + FakeManagedHttpJsonChannel firstChannel = lastCreatedChannel; + + int iterations = 100; + java.util.concurrent.ExecutorService executor = + java.util.concurrent.Executors.newFixedThreadPool(2); + try { + for (int i = 0; i < iterations; i++) { + firstChannel.nextCall = + new FakeHttpJsonClientCall() { + private boolean closed = false; + + @Override + public synchronized void start( + Listener responseListener, HttpJsonMetadata requestHeaders) { + if (closed) { + // Models HttpJsonClientCallImpl returning early when closed + return; + } + super.start(responseListener, requestHeaders); + } + + @Override + public synchronized void cancel(String message, Throwable cause) { + closed = true; + if (listener != null) { + listener.onClose(499, null); + } + } + }; + + HttpJsonClientCall call = channel.newCall(null, null); + java.util.concurrent.CyclicBarrier barrier = new java.util.concurrent.CyclicBarrier(2); + java.util.concurrent.Future f1 = + executor.submit( + () -> { + try { + barrier.await(); + call.start(new HttpJsonClientCall.Listener() {}, null); + } catch (Exception ignored) { + } + }); + java.util.concurrent.Future f2 = + executor.submit( + () -> { + try { + barrier.await(); + call.cancel("cancel", null); + } catch (Exception ignored) { + } + }); + f1.get(5, TimeUnit.SECONDS); + f2.get(5, TimeUnit.SECONDS); + } + } finally { + executor.shutdownNow(); + } + + testFingerprint = "fingerprint2"; + channel.refresh(); + assertTrue(firstChannel.isShutdown()); + } + + @Test + void cancel_whenStartedAndSuperCancelThrows_doesNotReleasePrematurelyUntilOnClose() { + RefreshingHttpJsonChannel channel = createTestChannel(); + FakeManagedHttpJsonChannel firstChannel = lastCreatedChannel; + FakeHttpJsonClientCall delegateCall = + new FakeHttpJsonClientCall() { + @Override + public void cancel(String message, Throwable cause) { + throw new RuntimeException("Simulated cancel failure"); + } + }; + firstChannel.nextCall = delegateCall; + + HttpJsonClientCall call = channel.newCall(null, null); + call.start(new HttpJsonClientCall.Listener() {}, null); + + assertThrows(RuntimeException.class, () -> call.cancel("abort", null)); + + // Rotate pool while call is still active (onClose hasn't fired yet): + // firstChannel must NOT be shut down yet because call is still active + testFingerprint = "fingerprint2"; + channel.refresh(); + assertFalse(firstChannel.isShutdown()); + + // Once onClose fires, entry is released and firstChannel shuts down + delegateCall.listener.onClose(200, null); + assertTrue(firstChannel.isShutdown()); + } + + @Test + void start_whenCalledTwice_throwsIllegalStateExceptionAndDoesNotReleaseFirstCallEntry() { + RefreshingHttpJsonChannel channel = createTestChannel(); + FakeManagedHttpJsonChannel firstChannel = lastCreatedChannel; + FakeHttpJsonClientCall delegateCall = new FakeHttpJsonClientCall<>(); + firstChannel.nextCall = delegateCall; + + HttpJsonClientCall call = channel.newCall(null, null); + call.start(new HttpJsonClientCall.Listener() {}, null); + + assertThrows( + IllegalStateException.class, + () -> call.start(new HttpJsonClientCall.Listener() {}, null)); + + testFingerprint = "fingerprint2"; + channel.refresh(); + assertFalse(firstChannel.isShutdown()); + + delegateCall.listener.onClose(200, null); + assertTrue(firstChannel.isShutdown()); + } } diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ApiResultRetryAlgorithm.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ApiResultRetryAlgorithm.java index 97d24d441ad2..1c80f1d90d59 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ApiResultRetryAlgorithm.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ApiResultRetryAlgorithm.java @@ -30,15 +30,52 @@ package com.google.api.gax.rpc; import com.google.api.gax.retrying.BasicResultRetryAlgorithm; +import com.google.api.gax.retrying.RetrySettings; import com.google.api.gax.retrying.RetryingContext; +import com.google.api.gax.retrying.TimedAttemptSettings; import org.jspecify.annotations.NullMarked; +import org.jspecify.annotations.Nullable; /* Package-private for internal use. */ @NullMarked class ApiResultRetryAlgorithm extends BasicResultRetryAlgorithm { @Override - public boolean shouldRetry(Throwable previousThrowable, ResponseT previousResponse) { + public @Nullable TimedAttemptSettings createNextAttempt( + @Nullable Throwable previousThrowable, + @Nullable ResponseT previousResponse, + TimedAttemptSettings previousSettings) { + if (previousThrowable instanceof UnauthenticatedException + && ((UnauthenticatedException) previousThrowable).isRetryable() + && previousSettings.getOverallAttemptCount() == previousSettings.getAttemptCount()) { + RetrySettings globalSettings = previousSettings.getGlobalSettings(); + if (globalSettings.getMaxAttempts() == 0 + && globalSettings.getTotalTimeoutDuration().isZero()) { + globalSettings = globalSettings.toBuilder().setMaxAttempts(1).build(); + } + return previousSettings.toBuilder() + .setGlobalSettings(globalSettings) + .setRetryDelayDuration(java.time.Duration.ZERO) + .setRandomizedRetryDelayDuration(java.time.Duration.ZERO) + .setAttemptCount(previousSettings.getAttemptCount()) + .setOverallAttemptCount(previousSettings.getOverallAttemptCount() + 1) + .build(); + } + return null; + } + + @Override + public @Nullable TimedAttemptSettings createNextAttempt( + @Nullable RetryingContext context, + @Nullable Throwable previousThrowable, + @Nullable ResponseT previousResponse, + TimedAttemptSettings previousSettings) { + return createNextAttempt(previousThrowable, previousResponse, previousSettings); + } + + @Override + public boolean shouldRetry( + @Nullable Throwable previousThrowable, @Nullable ResponseT previousResponse) { return (previousThrowable instanceof ApiException) && ((ApiException) previousThrowable).isRetryable(); } @@ -51,7 +88,9 @@ public boolean shouldRetry(Throwable previousThrowable, ResponseT previousRespon */ @Override public boolean shouldRetry( - RetryingContext context, Throwable previousThrowable, ResponseT previousResponse) { + RetryingContext context, + @Nullable Throwable previousThrowable, + @Nullable ResponseT previousResponse) { // Check UnauthenticatedException retryability first to ensure mTLS certificate // rotation retries take precedence over static method retry codes. if (previousThrowable instanceof UnauthenticatedException diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/AttemptCallable.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/AttemptCallable.java index 66dbcb2dcf14..a7211fb46d3b 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/AttemptCallable.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/AttemptCallable.java @@ -101,7 +101,6 @@ public ResponseT call() { unauthenticatedException -> { TransportChannel channel = finalContext.getTransportChannel(); if (channel != null) { - boolean shouldRetry = false; if (channel.shouldRefresh()) { try { channel.refresh(); @@ -111,11 +110,8 @@ public ResponseT call() { "Failed to refresh transport channel after authentication error", e); } - shouldRetry = true; - } else if (channel.getGeneration() > attemptGeneration) { - // Channel was rotated by a concurrent request while this call was in flight - shouldRetry = true; } + boolean shouldRetry = channel.getGeneration() > attemptGeneration; if (shouldRetry) { UnauthenticatedException newEx = diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ServerStreamingAttemptCallable.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ServerStreamingAttemptCallable.java index 7ef3c8491cb7..c645f20f806c 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ServerStreamingAttemptCallable.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ServerStreamingAttemptCallable.java @@ -226,6 +226,9 @@ public Void call() { .attemptStarted(request, outerRetryingFuture.getAttemptSettings().getOverallAttemptCount()); final ApiCallContext finalContext = attemptContext; + TransportChannel channelForAttempt = finalContext.getTransportChannel(); + final long attemptGeneration = + channelForAttempt != null ? channelForAttempt.getGeneration() : 0; innerCallable.call( request, new StateCheckingResponseObserver() { @@ -242,19 +245,43 @@ public void onResponseImpl(ResponseT response) { @Override public void onErrorImpl(Throwable t) { Throwable cause = t; - if (cause instanceof com.google.api.gax.retrying.ServerStreamingAttemptException) { + if (cause instanceof ServerStreamingAttemptException) { cause = cause.getCause(); } if (cause instanceof UnauthenticatedException) { + UnauthenticatedException unauthenticatedException = (UnauthenticatedException) cause; TransportChannel transportChannel = finalContext.getTransportChannel(); - if (transportChannel != null && transportChannel.shouldRefresh()) { - try { - transportChannel.refresh(); - } catch (Exception e) { - LOG.log( - Level.WARNING, - "Failed to refresh transport channel after authentication error", - e); + if (transportChannel != null) { + if (transportChannel.shouldRefresh()) { + try { + transportChannel.refresh(); + } catch (Exception e) { + LOG.log( + Level.WARNING, + "Failed to refresh transport channel after authentication error", + e); + } + } + boolean shouldRetry = transportChannel.getGeneration() > attemptGeneration; + if (shouldRetry) { + UnauthenticatedException newEx = + new UnauthenticatedException( + unauthenticatedException.getMessage(), + unauthenticatedException.getCause(), + unauthenticatedException.getStatusCode(), + true, + unauthenticatedException.getErrorDetails()); + for (Throwable suppressed : unauthenticatedException.getSuppressed()) { + newEx.addSuppressed(suppressed); + } + if (t instanceof ServerStreamingAttemptException) { + ServerStreamingAttemptException attemptEx = (ServerStreamingAttemptException) t; + t = + new ServerStreamingAttemptException( + newEx, attemptEx.canResume(), attemptEx.hasSeenResponses()); + } else { + t = newEx; + } } } } diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/CertificateRotationTracker.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/CertificateRotationTracker.java new file mode 100644 index 000000000000..fc58d67ab1d6 --- /dev/null +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/CertificateRotationTracker.java @@ -0,0 +1,197 @@ +/* + * Copyright 2026 Google LLC + * + * Redistribution and use in source and binary forms, with or without + * modification, are permitted provided that the following conditions are + * met: + * + * * Redistributions of source code must retain the above copyright + * notice, this list of conditions and the following disclaimer. + * * Redistributions in binary form must reproduce the above + * copyright notice, this list of conditions and the following disclaimer + * in the documentation and/or other materials provided with the + * distribution. + * * Neither the name of Google LLC nor the names of its + * contributors may be used to endorse or promote products derived from + * this software without specific prior written permission. + * + * THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS + * "AS IS" AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT + * LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR + * A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT + * OWNER OR CONTRIBUTORS BE LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, + * SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT + * LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, + * DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY + * THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT + * (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE + * OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + */ +package com.google.api.gax.rpc.mtls; + +import com.google.api.core.InternalApi; +import com.google.common.base.Strings; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicLong; +import java.util.concurrent.locks.ReentrantLock; +import java.util.function.Function; +import java.util.function.Supplier; +import org.jspecify.annotations.Nullable; + +/** + * Thread-safe helper that tracks workload certificate fingerprints on disk to detect mTLS + * certificate rotations for transport channels ({@code ChannelPool} and {@code + * RefreshingHttpJsonChannel}). + * + *

For internal use only. + */ +@InternalApi +public class CertificateRotationTracker { + + /** + * Duration (1 second) for which a detected certificate rotation (where the disk fingerprint + * differs from {@code activeCertFingerprint}) is cached in memory. + * + *

When a certificate rotates on disk, many concurrent in-flight RPCs may fail with {@code + * UNAUTHENTICATED} simultaneously and call {@link #shouldRefresh()}. Caching positive rotation + * detections for 1 second coalesces disk reads ({@code Files.readAllBytes}) and SHA-256 hashing + * across concurrent threads to prevent a thundering herd of file I/O, while bounding maximum + * staleness to 1 second. + */ + private static final long POSITIVE_ROTATION_CACHE_TTL_NANOS = TimeUnit.SECONDS.toNanos(1); + + private static final class DiskCheckResult { + final String fingerprint; + final long sequence; + final long timestampNanos; + + DiskCheckResult(String fingerprint, long sequence, long timestampNanos) { + this.fingerprint = fingerprint; + this.sequence = sequence; + this.timestampNanos = timestampNanos; + } + } + + private final Supplier certPathSupplier; + private final Function fingerprintReader; + private volatile String activeCertFingerprint; + private volatile DiskCheckResult lastDiskCheck = null; + private final ReentrantLock diskCheckLock = new ReentrantLock(); + private final AtomicLong diskCheckSequence = new AtomicLong(0); + + /** + * Creates a tracker for a fixed workload certificate path using {@link + * WorkloadCertificateUtils#getCertificateFingerprint(String)}. + */ + public CertificateRotationTracker(@Nullable String workloadCertPath) { + this(() -> workloadCertPath, WorkloadCertificateUtils::getCertificateFingerprint); + } + + /** + * Creates a tracker with custom certificate path and fingerprint suppliers (used by transports + * that expose package-private overrides for testing). + */ + public CertificateRotationTracker( + Supplier certPathSupplier, Function fingerprintReader) { + this.certPathSupplier = certPathSupplier; + this.fingerprintReader = fingerprintReader; + String initialCertPath = certPathSupplier.get(); + this.activeCertFingerprint = + initialCertPath != null + ? Strings.nullToEmpty(fingerprintReader.apply(initialCertPath)) + : ""; + } + + /** + * Returns {@code true} if a workload certificate path is configured, readable, and its current + * SHA-256 fingerprint on disk differs from the active fingerprint. + */ + public boolean shouldRefresh() { + String certPath = certPathSupplier.get(); + if (certPath == null) { + return false; + } + String currentDiskFingerprint = getOrUpdateDiskFingerprint(certPath); + if (currentDiskFingerprint.isEmpty()) { + return false; + } + return !currentDiskFingerprint.equalsIgnoreCase(activeCertFingerprint); + } + + private String getOrUpdateDiskFingerprint(String certPath) { + long seqBeforeLock = diskCheckSequence.get(); + long now = System.nanoTime(); + DiskCheckResult cached = lastDiskCheck; + if (cached != null + && !cached.fingerprint.isEmpty() + && !cached.fingerprint.equalsIgnoreCase(this.activeCertFingerprint) + && (now - cached.timestampNanos < POSITIVE_ROTATION_CACHE_TTL_NANOS)) { + return cached.fingerprint; + } + + diskCheckLock.lock(); + try { + now = System.nanoTime(); + cached = lastDiskCheck; + if (cached != null + && !cached.fingerprint.isEmpty() + && (cached.sequence > seqBeforeLock + || (!cached.fingerprint.equalsIgnoreCase(this.activeCertFingerprint) + && (now - cached.timestampNanos < POSITIVE_ROTATION_CACHE_TTL_NANOS)))) { + return cached.fingerprint; + } + long newSeq = diskCheckSequence.incrementAndGet(); + String fingerprint = Strings.nullToEmpty(fingerprintReader.apply(certPath)); + if (!fingerprint.isEmpty()) { + lastDiskCheck = new DiskCheckResult(fingerprint, newSeq, System.nanoTime()); + } else { + lastDiskCheck = null; + } + return fingerprint; + } finally { + diskCheckLock.unlock(); + } + } + + /** + * Reads the current certificate fingerprint directly from disk (bypassing the 1-second cache), + * returning {@code ""} if no workload certificate path is configured or if the file is currently + * unreadable/empty. + */ + public String readDiskFingerprint() { + String certPath = certPathSupplier.get(); + if (certPath == null) { + return ""; + } + return Strings.nullToEmpty(fingerprintReader.apply(certPath)); + } + + /** + * Returns {@code true} if {@code diskFingerprint} matches the currently active certificate + * fingerprint. + */ + public boolean isAlreadyActive(String diskFingerprint) { + return diskFingerprint != null && diskFingerprint.equalsIgnoreCase(this.activeCertFingerprint); + } + + /** + * Updates the active certificate fingerprint after a successful channel refresh and clears any + * cached disk check result. + */ + public void markRefreshed(String newFingerprint) { + if (newFingerprint != null && !newFingerprint.isEmpty()) { + this.activeCertFingerprint = newFingerprint; + this.lastDiskCheck = null; + } + } + + /** Returns the currently active certificate fingerprint (or {@code ""} if none). */ + public String getActiveCertFingerprint() { + return activeCertFingerprint; + } + + /** Invalidates the cached disk check result. Visible for testing. */ + public void invalidateCache() { + this.lastDiskCheck = null; + } +} diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/WorkloadCertificateUtils.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/WorkloadCertificateUtils.java index b6dac90db34c..94fb36db1ac3 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/WorkloadCertificateUtils.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/mtls/WorkloadCertificateUtils.java @@ -31,15 +31,34 @@ import com.google.api.core.InternalApi; import com.google.auth.mtls.MtlsUtils; +import java.io.File; /** Internal utility class for managing dynamic workload certificates. */ @InternalApi public class WorkloadCertificateUtils { + private static final String EMPTY_FILE_SHA256 = + "e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855"; + private WorkloadCertificateUtils() {} + /** + * Computes the SHA-256 fingerprint of the certificate file at {@code certPath}, returning {@code + * ""} if the path is {@code null}, unreadable, empty (e.g., temporarily truncated to 0 bytes + * mid-write by an external certificate rotator), or hashes to the empty-byte digest. + * + *

Returning {@code ""} on unreadable or empty files ensures callers ({@code shouldRefresh()} + * and {@code refresh()}) safely skip refreshing during transient mid-write states rather than + * treating an empty digest as a certificate rotation mismatch. + */ public static String getCertificateFingerprint(String certPath) { + if (certPath == null || new File(certPath).length() == 0) { + return ""; + } String fingerprint = MtlsUtils.getCertificateFingerprint(certPath); - return fingerprint != null ? fingerprint : ""; + if (fingerprint == null || EMPTY_FILE_SHA256.equalsIgnoreCase(fingerprint)) { + return ""; + } + return fingerprint; } } diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ApiResultRetryAlgorithmTest.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ApiResultRetryAlgorithmTest.java index 300a0ad30130..41daa50745ea 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ApiResultRetryAlgorithmTest.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ApiResultRetryAlgorithmTest.java @@ -29,14 +29,22 @@ */ package com.google.api.gax.rpc; +import static org.junit.jupiter.api.Assertions.assertEquals; import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertNotNull; import static org.junit.jupiter.api.Assertions.assertTrue; import static org.mockito.Mockito.mock; import static org.mockito.Mockito.when; +import com.google.api.core.NanoClock; +import com.google.api.gax.retrying.ExponentialRetryAlgorithm; +import com.google.api.gax.retrying.RetryAlgorithm; +import com.google.api.gax.retrying.RetrySettings; +import com.google.api.gax.retrying.TimedAttemptSettings; import com.google.api.gax.rpc.StatusCode.Code; import com.google.api.gax.rpc.testing.FakeStatusCode; import com.google.common.collect.Sets; +import java.time.Duration; import java.util.Collections; import org.junit.jupiter.api.Test; import org.mockito.Mockito; @@ -111,4 +119,133 @@ void testShouldRetryWithContextWithEmptyRetryableCodes() { ApiResultRetryAlgorithm algorithm = new ApiResultRetryAlgorithm<>(); assertFalse(algorithm.shouldRetry(context, unavailableException, null)); } + + @Test + void testRotationRetryWithNonRetryableSettings_maxAttemptsOne() { + ApiCallContext context = + mock(ApiCallContext.class, Mockito.withSettings().withoutAnnotations()); + when(context.getRetryableCodes()).thenReturn(Collections.emptySet()); + + RetrySettings settings = + RetrySettings.newBuilder() + .setMaxAttempts(1) + .setInitialRpcTimeoutDuration(Duration.ofSeconds(10)) + .setMaxRpcTimeoutDuration(Duration.ofSeconds(10)) + .setTotalTimeoutDuration(Duration.ofSeconds(10)) + .build(); + + ApiResultRetryAlgorithm resultAlgorithm = new ApiResultRetryAlgorithm<>(); + ExponentialRetryAlgorithm timedAlgorithm = + new ExponentialRetryAlgorithm(settings, NanoClock.getDefaultClock()); + RetryAlgorithm retryAlgorithm = new RetryAlgorithm<>(resultAlgorithm, timedAlgorithm); + + TimedAttemptSettings firstAttempt = retryAlgorithm.createFirstAttempt(context); + UnauthenticatedException rotationEx = + new UnauthenticatedException( + "Expired cert", null, new FakeStatusCode(Code.UNAUTHENTICATED), /* retryable= */ true); + + // First rotation failure: grants immediate free retry without incrementing attemptCount + TimedAttemptSettings nextAttempt = + retryAlgorithm.createNextAttempt(context, rotationEx, null, firstAttempt); + assertNotNull(nextAttempt); + assertEquals(Duration.ZERO, nextAttempt.getRetryDelayDuration()); + assertEquals(Duration.ZERO, nextAttempt.getRandomizedRetryDelayDuration()); + assertEquals(0, nextAttempt.getAttemptCount()); + assertEquals(1, nextAttempt.getOverallAttemptCount()); + assertTrue(retryAlgorithm.shouldRetry(context, rotationEx, null, nextAttempt)); + + // Second consecutive failure: overallAttemptCount (1) != attemptCount (0), so no free retry + TimedAttemptSettings thirdAttempt = + retryAlgorithm.createNextAttempt(context, rotationEx, null, nextAttempt); + assertNotNull(thirdAttempt); + assertEquals(1, thirdAttempt.getAttemptCount()); + assertEquals(2, thirdAttempt.getOverallAttemptCount()); + assertFalse(retryAlgorithm.shouldRetry(context, rotationEx, null, thirdAttempt)); + } + + @Test + void testRotationRetryWithNonRetryableSettings_zeroMaxAttemptsZeroTotalTimeout() { + ApiCallContext context = + mock(ApiCallContext.class, Mockito.withSettings().withoutAnnotations()); + when(context.getRetryableCodes()).thenReturn(Collections.emptySet()); + + RetrySettings settings = + RetrySettings.newBuilder().setMaxAttempts(0).setTotalTimeoutDuration(Duration.ZERO).build(); + + ApiResultRetryAlgorithm resultAlgorithm = new ApiResultRetryAlgorithm<>(); + ExponentialRetryAlgorithm timedAlgorithm = + new ExponentialRetryAlgorithm(settings, NanoClock.getDefaultClock()); + RetryAlgorithm retryAlgorithm = new RetryAlgorithm<>(resultAlgorithm, timedAlgorithm); + + TimedAttemptSettings firstAttempt = retryAlgorithm.createFirstAttempt(context); + UnauthenticatedException rotationEx = + new UnauthenticatedException( + "Expired cert", null, new FakeStatusCode(Code.UNAUTHENTICATED), /* retryable= */ true); + + TimedAttemptSettings nextAttempt = + retryAlgorithm.createNextAttempt(context, rotationEx, null, firstAttempt); + assertNotNull(nextAttempt); + assertEquals(1, nextAttempt.getGlobalSettings().getMaxAttempts()); + assertEquals(0, nextAttempt.getAttemptCount()); + assertEquals(1, nextAttempt.getOverallAttemptCount()); + assertTrue(retryAlgorithm.shouldRetry(context, rotationEx, null, nextAttempt)); + + // Subsequent failure is rejected + TimedAttemptSettings thirdAttempt = + retryAlgorithm.createNextAttempt(context, rotationEx, null, nextAttempt); + assertNotNull(thirdAttempt); + assertFalse(retryAlgorithm.shouldRetry(context, rotationEx, null, thirdAttempt)); + } + + @Test + void testRotationRetryAfterTransientErrorPreservesRemainingBudget() { + ApiCallContext context = + mock(ApiCallContext.class, Mockito.withSettings().withoutAnnotations()); + when(context.getRetryableCodes()).thenReturn(Sets.newHashSet(Code.UNAVAILABLE)); + + RetrySettings settings = + RetrySettings.newBuilder() + .setMaxAttempts(3) + .setInitialRetryDelayDuration(Duration.ofMillis(100)) + .setRetryDelayMultiplier(2.0) + .setMaxRetryDelayDuration(Duration.ofSeconds(1)) + .setTotalTimeoutDuration(Duration.ofSeconds(30)) + .build(); + + ApiResultRetryAlgorithm resultAlgorithm = new ApiResultRetryAlgorithm<>(); + ExponentialRetryAlgorithm timedAlgorithm = + new ExponentialRetryAlgorithm(settings, NanoClock.getDefaultClock()); + RetryAlgorithm retryAlgorithm = new RetryAlgorithm<>(resultAlgorithm, timedAlgorithm); + + TimedAttemptSettings attempt0 = retryAlgorithm.createFirstAttempt(context); + ApiException unavailableEx = + new ApiException(null, new FakeStatusCode(Code.UNAVAILABLE), /* retryable= */ true); + UnauthenticatedException rotationEx = + new UnauthenticatedException( + "Expired cert", null, new FakeStatusCode(Code.UNAUTHENTICATED), /* retryable= */ true); + + // Attempt 0 fails with UNAVAILABLE -> normal retry (attemptCount = 1, overallAttemptCount = 1) + TimedAttemptSettings attempt1 = + retryAlgorithm.createNextAttempt(context, unavailableEx, null, attempt0); + assertEquals(1, attempt1.getAttemptCount()); + assertEquals(1, attempt1.getOverallAttemptCount()); + assertTrue(retryAlgorithm.shouldRetry(context, unavailableEx, null, attempt1)); + + // Attempt 1 fails with rotation 401 -> free zero-delay retry (attemptCount = 1, + // overallAttemptCount = 2) + TimedAttemptSettings attempt2 = + retryAlgorithm.createNextAttempt(context, rotationEx, null, attempt1); + assertEquals(Duration.ZERO, attempt2.getRetryDelayDuration()); + assertEquals(1, attempt2.getAttemptCount()); + assertEquals(2, attempt2.getOverallAttemptCount()); + assertTrue(retryAlgorithm.shouldRetry(context, rotationEx, null, attempt2)); + + // Attempt 2 fails with UNAVAILABLE -> normal retry still allowed (attemptCount = 2 < + // maxAttempts = 3) + TimedAttemptSettings attempt3 = + retryAlgorithm.createNextAttempt(context, unavailableEx, null, attempt2); + assertEquals(2, attempt3.getAttemptCount()); + assertEquals(3, attempt3.getOverallAttemptCount()); + assertTrue(retryAlgorithm.shouldRetry(context, unavailableEx, null, attempt3)); + } } diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/AttemptCallableTest.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/AttemptCallableTest.java index baeb35630014..84c51da90707 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/AttemptCallableTest.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/AttemptCallableTest.java @@ -275,7 +275,7 @@ void testPermanentUnauthenticatedFailure_sameGeneration_notMarkedRetryable() { } @Test - void testRefreshThrowsException_originalUnauthenticatedPropagated() { + void testRefreshThrowsException_notMarkedRetryableWhenGenerationUnchanged() { FakeChannel fakeChannel = new FakeChannel() { @Override @@ -317,6 +317,105 @@ public void refresh() { thrown = e.getCause(); } + assertThat(thrown).isInstanceOf(UnauthenticatedException.class); + UnauthenticatedException rethrown = (UnauthenticatedException) thrown; + assertThat(rethrown.isRetryable()).isFalse(); + } + + @Test + void testRefreshReturnsWithoutAdvancingGeneration_notMarkedRetryable() { + FakeChannel fakeChannel = + new FakeChannel() { + @Override + public boolean shouldRefresh() { + return true; + } + + @Override + public void refresh() { + // Simulates refresh() returning early without rotating any channel (e.g. unreadable + // cert) + } + }; + FakeTransportChannel transportChannel = FakeTransportChannel.create(fakeChannel); + + ApiCallContext callContext = + FakeCallContext.createDefault().withTransportChannel(transportChannel); + + UnauthenticatedException originalEx = + new UnauthenticatedException( + "Expired cert", null, FakeStatusCode.of(StatusCode.Code.UNAUTHENTICATED), false); + + SettableApiFuture failedFuture = SettableApiFuture.create(); + failedFuture.setException(originalEx); + when(mockInnerCallable.futureCall(Mockito.anyString(), Mockito.any())).thenReturn(failedFuture); + + AttemptCallable callable = + new AttemptCallable<>(mockInnerCallable, "fake-request", callContext); + callable.setExternalFuture(mockExternalFuture); + + callable.call(); + + ArgumentCaptor futureCaptor = ArgumentCaptor.forClass(ApiFuture.class); + Mockito.verify(mockExternalFuture, Mockito.times(2)).setAttemptFuture(futureCaptor.capture()); + + Throwable thrown = null; + try { + futureCaptor.getValue().get(); + } catch (Exception e) { + thrown = e.getCause(); + } + + assertThat(thrown).isInstanceOf(UnauthenticatedException.class); + UnauthenticatedException rethrown = (UnauthenticatedException) thrown; + assertThat(rethrown.isRetryable()).isFalse(); + } + + @Test + void testRefreshThrowsException_markedRetryableIfConcurrentThreadAdvancedGeneration() { + FakeChannel fakeChannel = + new FakeChannel() { + @Override + public boolean shouldRefresh() { + return true; + } + + @Override + public void refresh() { + // Concurrent thread advanced generation before/during refresh failure + setGeneration(getGeneration() + 1); + throw new RuntimeException("Refresh error on this thread"); + } + }; + FakeTransportChannel transportChannel = FakeTransportChannel.create(fakeChannel); + + ApiCallContext callContext = + FakeCallContext.createDefault().withTransportChannel(transportChannel); + + UnauthenticatedException originalEx = + new UnauthenticatedException( + "Expired cert", null, FakeStatusCode.of(StatusCode.Code.UNAUTHENTICATED), false); + + SettableApiFuture failedFuture = SettableApiFuture.create(); + failedFuture.setException(originalEx); + when(mockInnerCallable.futureCall(Mockito.anyString(), Mockito.any())).thenReturn(failedFuture); + + AttemptCallable callable = + new AttemptCallable<>(mockInnerCallable, "fake-request", callContext); + callable.setExternalFuture(mockExternalFuture); + + callable.call(); + + ArgumentCaptor futureCaptor = ArgumentCaptor.forClass(ApiFuture.class); + Mockito.verify(mockExternalFuture, Mockito.times(2)).setAttemptFuture(futureCaptor.capture()); + + Throwable thrown = null; + try { + futureCaptor.getValue().get(); + } catch (Exception e) { + thrown = e.getCause(); + } + assertThat(thrown).isInstanceOf(UnauthenticatedException.class); UnauthenticatedException rethrown = (UnauthenticatedException) thrown; assertThat(rethrown.isRetryable()).isTrue(); diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ServerStreamingAttemptCallableTest.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ServerStreamingAttemptCallableTest.java index ca6ea284d845..740dc37b7ffe 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ServerStreamingAttemptCallableTest.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ServerStreamingAttemptCallableTest.java @@ -42,6 +42,9 @@ import com.google.api.gax.rpc.StatusCode.Code; import com.google.api.gax.rpc.testing.FakeApiException; import com.google.api.gax.rpc.testing.FakeCallContext; +import com.google.api.gax.rpc.testing.FakeChannel; +import com.google.api.gax.rpc.testing.FakeStatusCode; +import com.google.api.gax.rpc.testing.FakeTransportChannel; import com.google.api.gax.rpc.testing.MockStreamingApi.MockServerStreamingCall; import com.google.api.gax.rpc.testing.MockStreamingApi.MockServerStreamingCallable; import com.google.api.gax.tracing.BaseApiTracer; @@ -289,6 +292,115 @@ void testUnauthenticatedRefresh() { Truth.assertThat(((UnauthenticatedException) outerError.getCause()).isRetryable()).isFalse(); } + @Test + @SuppressWarnings("ConstantConditions") + void testUnauthenticatedRefreshWithGenerationAdvanceRetries() { + TransportChannel transportChannel = Mockito.mock(TransportChannel.class); + java.util.concurrent.atomic.AtomicLong generation = + new java.util.concurrent.atomic.AtomicLong(0); + Mockito.when(transportChannel.getGeneration()).thenAnswer(inv -> generation.get()); + Mockito.when(transportChannel.shouldRefresh()).thenReturn(true); + Mockito.doAnswer( + inv -> { + generation.incrementAndGet(); + return null; + }) + .when(transportChannel) + .refresh(); + + ApiCallContext context = Mockito.mock(ApiCallContext.class); + Mockito.when(context.getTransportChannel()).thenReturn(transportChannel); + Mockito.when(context.getTracer()).thenReturn(BaseApiTracer.getInstance()); + Mockito.when(context.getTimeoutDuration()).thenReturn(java.time.Duration.ofHours(5)); + + resumptionStrategy = new MyStreamResumptionStrategy(); + ServerStreamingAttemptCallable callable = createCallable(context); + callable.start(); + + MockServerStreamingCall call = innerCallable.popLastCall(); + + UnauthenticatedException initialError = + new UnauthenticatedException( + "test", + null, + com.google.api.gax.rpc.testing.FakeStatusCode.of(Code.UNAUTHENTICATED), + false); + call.getController().getObserver().onError(initialError); + + ExecutionException ee = + assertThrows( + ExecutionException.class, + () -> fakeRetryingFuture.getAttemptResult().get(1, TimeUnit.SECONDS)); + Throwable outerError = ee.getCause(); + Mockito.verify(transportChannel).refresh(); + Truth.assertThat(outerError).isInstanceOf(ServerStreamingAttemptException.class); + Truth.assertThat(((ServerStreamingAttemptException) outerError).canResume()).isTrue(); + Truth.assertThat(outerError.getCause()).isInstanceOf(UnauthenticatedException.class); + Truth.assertThat(((UnauthenticatedException) outerError.getCause()).isRetryable()).isTrue(); + + // Verify retry call resumes stream + callable.call(); + call = innerCallable.popLastCall(); + Truth.assertThat(call.getRequest()).isEqualTo("request > 0"); + } + + @Test + @SuppressWarnings("ConstantConditions") + void testUnauthenticatedRefreshWithNonResumableStreamDoesNotRetry() { + TransportChannel transportChannel = Mockito.mock(TransportChannel.class); + java.util.concurrent.atomic.AtomicLong generation = + new java.util.concurrent.atomic.AtomicLong(0); + Mockito.when(transportChannel.getGeneration()).thenAnswer(inv -> generation.get()); + Mockito.when(transportChannel.shouldRefresh()).thenReturn(true); + Mockito.doAnswer( + inv -> { + generation.incrementAndGet(); + return null; + }) + .when(transportChannel) + .refresh(); + + ApiCallContext context = Mockito.mock(ApiCallContext.class); + Mockito.when(context.getTransportChannel()).thenReturn(transportChannel); + Mockito.when(context.getTracer()).thenReturn(BaseApiTracer.getInstance()); + Mockito.when(context.getTimeoutDuration()).thenReturn(java.time.Duration.ofHours(5)); + + // SimpleStreamResumptionStrategy cannot resume once a response has been received + resumptionStrategy = new com.google.api.gax.retrying.SimpleStreamResumptionStrategy<>(); + ServerStreamingAttemptCallable callable = createCallable(context); + callable.start(); + + MockServerStreamingCall call = innerCallable.popLastCall(); + call.getController().getObserver().onResponse("response1"); + + UnauthenticatedException initialError = + new UnauthenticatedException( + "test", + null, + com.google.api.gax.rpc.testing.FakeStatusCode.of(Code.UNAUTHENTICATED), + false); + call.getController().getObserver().onError(initialError); + + ExecutionException ee = + assertThrows( + ExecutionException.class, + () -> fakeRetryingFuture.getAttemptResult().get(1, TimeUnit.SECONDS)); + Throwable outerError = ee.getCause(); + Mockito.verify(transportChannel).refresh(); + Truth.assertThat(outerError).isInstanceOf(ServerStreamingAttemptException.class); + ServerStreamingAttemptException attemptEx = (ServerStreamingAttemptException) outerError; + Truth.assertThat(attemptEx.canResume()).isFalse(); + Truth.assertThat(((UnauthenticatedException) attemptEx.getCause()).isRetryable()).isTrue(); + Truth.assertThat( + new com.google.api.gax.retrying.StreamingRetryAlgorithm<>( + new ApiResultRetryAlgorithm<>(), + new com.google.api.gax.retrying.ExponentialRetryAlgorithm( + RetrySettings.newBuilder().build(), + com.google.api.core.NanoClock.getDefaultClock())) + .shouldRetry(attemptEx, null, fakeRetryingFuture.getAttemptSettings())) + .isFalse(); + } + @Test @SuppressWarnings("ConstantConditions") void testRefreshThrowsException_originalErrorNotLost() { @@ -480,6 +592,67 @@ public String processResponse(String response) { .containsExactly("first+suffix", "second+suffix", "third+suffix"); } + @Test + void testUnauthenticatedException_whenChannelRefreshes_setsRetryableTrue() throws Exception { + FakeChannel fakeChannel = new FakeChannel(); + fakeChannel.setShouldRefresh(true); + ApiCallContext context = + FakeCallContext.createDefault() + .withTransportChannel(FakeTransportChannel.create(fakeChannel)); + + ServerStreamingAttemptCallable callable = createCallable(context); + callable.start(); + + MockServerStreamingCall call = innerCallable.popLastCall(); + UnauthenticatedException unauthEx = + new UnauthenticatedException( + "cert expired", null, new FakeStatusCode(Code.UNAUTHENTICATED), false); + call.getController().getObserver().onError(unauthEx); + + ExecutionException ex = + assertThrows( + ExecutionException.class, + () -> fakeRetryingFuture.getAttemptResult().get(1, TimeUnit.SECONDS)); + Truth.assertThat(ex.getCause()).isInstanceOf(ServerStreamingAttemptException.class); + Throwable cause = ex.getCause().getCause(); + Truth.assertThat(cause).isInstanceOf(UnauthenticatedException.class); + Truth.assertThat(((UnauthenticatedException) cause).isRetryable()).isTrue(); + Truth.assertThat(fakeChannel.getRefreshCount()).isEqualTo(1); + } + + @Test + void testUnauthenticatedException_whenChannelRefreshFails_remainsNonRetryable() throws Exception { + FakeChannel fakeChannel = + new FakeChannel() { + @Override + public void refresh() { + throw new RuntimeException("Refresh failed"); + } + }; + fakeChannel.setShouldRefresh(true); + ApiCallContext context = + FakeCallContext.createDefault() + .withTransportChannel(FakeTransportChannel.create(fakeChannel)); + + ServerStreamingAttemptCallable callable = createCallable(context); + callable.start(); + + MockServerStreamingCall call = innerCallable.popLastCall(); + UnauthenticatedException unauthEx = + new UnauthenticatedException( + "cert expired", null, new FakeStatusCode(Code.UNAUTHENTICATED), false); + call.getController().getObserver().onError(unauthEx); + + ExecutionException ex = + assertThrows( + ExecutionException.class, + () -> fakeRetryingFuture.getAttemptResult().get(1, TimeUnit.SECONDS)); + Truth.assertThat(ex.getCause()).isInstanceOf(ServerStreamingAttemptException.class); + Throwable cause = ex.getCause().getCause(); + Truth.assertThat(cause).isInstanceOf(UnauthenticatedException.class); + Truth.assertThat(((UnauthenticatedException) cause).isRetryable()).isFalse(); + } + static class MyStreamResumptionStrategy implements StreamResumptionStrategy { private int responseCount; diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/CertificateBasedAccessTest.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/CertificateBasedAccessTest.java index 0e7055232298..92c0dfb96f3a 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/CertificateBasedAccessTest.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/mtls/CertificateBasedAccessTest.java @@ -34,6 +34,7 @@ import static org.junit.jupiter.api.Assertions.assertFalse; import static org.junit.jupiter.api.Assertions.assertNull; import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTrue; import java.util.HashMap; import java.util.Map; @@ -107,9 +108,10 @@ void testUseMtlsClientCertificateExplicitTrueNoCredentials() { TestEnv env = new TestEnv(); env.set("GOOGLE_API_USE_CLIENT_CERTIFICATE", "true"); CertificateBasedAccess cba = createCba(env); - // Explicit 'true' permits mTLS if certs exist, but if no certs are present, returns false/null - // cleanly (Row 3) - assertFalse(cba.useMtlsClientCertificate()); + // Explicit 'true' enables mTLS client certificate usage (for ECP / custom MtlsProvider) even + // when no workload cert files are present, while getWorkloadCertPath returns null so file + // rotation polling is not active. + assertTrue(cba.useMtlsClientCertificate()); assertNull(cba.getWorkloadCertPath()); } @@ -143,4 +145,18 @@ void testUseMtlsClientCertificateConfigMissingConfigFile_throwsIllegalStateExcep assertThrows(IllegalStateException.class, () -> cba.useMtlsClientCertificate()); assertThrows(IllegalStateException.class, () -> cba.getWorkloadCertPath()); } + + @Test + void testWorkloadCertificateUtilsEmptyFileReturnsEmptyString() throws Exception { + java.io.File tempFile = java.io.File.createTempFile("test-cert-empty", ".pem"); + tempFile.deleteOnExit(); + // 0-byte truncated file mid-write should return empty string rather than SHA-256 of empty bytes + assertEquals( + "", WorkloadCertificateUtils.getCertificateFingerprint(tempFile.getAbsolutePath())); + + java.nio.file.Files.write( + tempFile.toPath(), "test-cert-content".getBytes(java.nio.charset.StandardCharsets.UTF_8)); + String fp = WorkloadCertificateUtils.getCertificateFingerprint(tempFile.getAbsolutePath()); + assertFalse(fp.isEmpty()); + } } diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/testing/FakeChannel.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/testing/FakeChannel.java index 3f77f9affe93..9acfc703b9c3 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/testing/FakeChannel.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/testing/FakeChannel.java @@ -47,6 +47,7 @@ public boolean shouldRefresh() { public void refresh() { refreshCount++; + generation++; } public int getRefreshCount() { From 92135bfd1a336be034bd0d39c0e4b9529fd2f848 Mon Sep 17 00:00:00 2001 From: Matt Castelaz Date: Fri, 18 Sep 2026 16:18:16 +0000 Subject: [PATCH 15/29] fix(gax-httpjson,gax-grpc): fix Conscrypt mTLS KeyManagerFactory init and JDK 8 mock annotations (#13995) --- .../google/api/gax/grpc/ChannelPoolTest.java | 30 ++++++++++++------- .../InstantiatingHttpJsonChannelProvider.java | 17 ++++++++++- ...tantiatingHttpJsonChannelProviderTest.java | 24 +++++++++++++++ 3 files changed, 60 insertions(+), 11 deletions(-) diff --git a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java index 4dde16f9aedf..9c408f2d128c 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java @@ -511,7 +511,8 @@ void channelReactiveMTlsRefresh_failedCreationDoesNotMutateFingerprintAndAllowsR throws IOException { ManagedChannel channel1 = Mockito.mock(ManagedChannel.class); ManagedChannel channel2 = Mockito.mock(ManagedChannel.class); - ChannelFactory channelFactory = Mockito.mock(ChannelFactory.class); + ChannelFactory channelFactory = + Mockito.mock(ChannelFactory.class, Mockito.withSettings().withoutAnnotations()); // Initial creation returns channel1, refresh attempt 1 throws IOException, refresh attempt 2 // returns channel2 @@ -574,7 +575,8 @@ void channelReactiveMTlsRefresh_failedCreationDoesNotMutateFingerprintAndAllowsR ManagedChannel rotated1 = Mockito.mock(ManagedChannel.class); ManagedChannel rotated1SecondPass = Mockito.mock(ManagedChannel.class); ManagedChannel rotated2 = Mockito.mock(ManagedChannel.class); - ChannelFactory channelFactory = Mockito.mock(ChannelFactory.class); + ChannelFactory channelFactory = + Mockito.mock(ChannelFactory.class, Mockito.withSettings().withoutAnnotations()); // Initial creation: initial1, initial2 // Refresh pass 1: rotated1 succeeds, second throws IOException @@ -632,7 +634,8 @@ void refreshAll_runtimeExceptionOrError_doesNotLeakCreatedChannels() throws IOEx ManagedChannel initial2 = Mockito.mock(ManagedChannel.class); ManagedChannel createdBeforeRuntimeEx = Mockito.mock(ManagedChannel.class); ManagedChannel createdBeforeError = Mockito.mock(ManagedChannel.class); - ChannelFactory channelFactory = Mockito.mock(ChannelFactory.class); + ChannelFactory channelFactory = + Mockito.mock(ChannelFactory.class, Mockito.withSettings().withoutAnnotations()); Mockito.when(channelFactory.createSingleChannel()) .thenReturn(initial1, initial2) @@ -659,7 +662,8 @@ void refreshAll_runtimeExceptionOrError_doesNotLeakCreatedChannels() throws IOEx void refresh_onShutdownPool_noOpsAndCreatesNoChannels() throws IOException { ManagedChannel channel1 = mock(ManagedChannel.class); ManagedChannel channel2 = mock(ManagedChannel.class); - ChannelFactory channelFactory = Mockito.mock(ChannelFactory.class); + ChannelFactory channelFactory = + Mockito.mock(ChannelFactory.class, Mockito.withSettings().withoutAnnotations()); Mockito.when(channelFactory.createSingleChannel()).thenReturn(channel1, channel2); pool = ChannelPool.create(ChannelPoolSettings.staticallySized(1), channelFactory, null, null); @@ -680,7 +684,8 @@ void refresh_onShutdownPool_noOpsAndCreatesNoChannels() throws IOException { void generationCounterIncrementsOnRefresh() throws IOException { ManagedChannel channel1 = mock(ManagedChannel.class); ManagedChannel channel2 = mock(ManagedChannel.class); - ChannelFactory channelFactory = Mockito.mock(ChannelFactory.class); + ChannelFactory channelFactory = + Mockito.mock(ChannelFactory.class, Mockito.withSettings().withoutAnnotations()); Mockito.when(channelFactory.createSingleChannel()).thenReturn(channel1, channel2); pool = ChannelPool.create(ChannelPoolSettings.staticallySized(1), channelFactory, null, null); @@ -1234,7 +1239,8 @@ void shouldRefresh_doesNotCacheNegativeResultAndDetectsSubsequentRotationImmedia throws IOException { ManagedChannel initial = Mockito.mock(ManagedChannel.class); ManagedChannel rotated = Mockito.mock(ManagedChannel.class); - ChannelFactory channelFactory = Mockito.mock(ChannelFactory.class); + ChannelFactory channelFactory = + Mockito.mock(ChannelFactory.class, Mockito.withSettings().withoutAnnotations()); Mockito.when(channelFactory.createSingleChannel()).thenReturn(initial, rotated); tempCert = java.nio.file.Files.createTempFile("cert", ".pem"); @@ -1268,7 +1274,8 @@ void newCall_whenDelegateThrowsError_releasesEntryAndShutsDownRetiredChannel() throws IOException { ManagedChannel initial = Mockito.mock(ManagedChannel.class); ManagedChannel rotated = Mockito.mock(ManagedChannel.class); - ChannelFactory channelFactory = Mockito.mock(ChannelFactory.class); + ChannelFactory channelFactory = + Mockito.mock(ChannelFactory.class, Mockito.withSettings().withoutAnnotations()); Mockito.when(channelFactory.createSingleChannel()).thenReturn(initial, rotated); Mockito.when(initial.newCall(Mockito.any(), Mockito.any())) .thenThrow(new LinkageError("Simulated native/JNI linkage error")); @@ -1288,7 +1295,8 @@ void start_whenDelegateThrowsError_releasesEntryAndShutsDownRetiredChannel() thr ManagedChannel initial = Mockito.mock(ManagedChannel.class); ManagedChannel rotated = Mockito.mock(ManagedChannel.class); ClientCall mockCall = Mockito.mock(ClientCall.class); - ChannelFactory channelFactory = Mockito.mock(ChannelFactory.class); + ChannelFactory channelFactory = + Mockito.mock(ChannelFactory.class, Mockito.withSettings().withoutAnnotations()); Mockito.when(channelFactory.createSingleChannel()).thenReturn(initial, rotated); Mockito.when(initial.newCall(Mockito.any(), Mockito.any())).thenReturn((ClientCall) mockCall); Mockito.doThrow(new AssertionError("Simulated Error in start")) @@ -1316,7 +1324,8 @@ void cancel_whenDelegateThrowsException_releasesEntryAndShutsDownRetiredChannel( ManagedChannel initial = Mockito.mock(ManagedChannel.class); ManagedChannel rotated = Mockito.mock(ManagedChannel.class); ClientCall mockCall = Mockito.mock(ClientCall.class); - ChannelFactory channelFactory = Mockito.mock(ChannelFactory.class); + ChannelFactory channelFactory = + Mockito.mock(ChannelFactory.class, Mockito.withSettings().withoutAnnotations()); Mockito.when(channelFactory.createSingleChannel()).thenReturn(initial, rotated); Mockito.when(initial.newCall(Mockito.any(), Mockito.any())).thenReturn((ClientCall) mockCall); Mockito.doThrow(new RuntimeException("Simulated cancel exception")) @@ -1338,7 +1347,8 @@ void cancel_whenDelegateThrowsException_releasesEntryAndShutsDownRetiredChannel( void concurrentStartAndCancel_neverLeaksOrDoubleReleasesEntry() throws Exception { ManagedChannel initial = Mockito.mock(ManagedChannel.class); ManagedChannel rotated = Mockito.mock(ManagedChannel.class); - ChannelFactory channelFactory = Mockito.mock(ChannelFactory.class); + ChannelFactory channelFactory = + Mockito.mock(ChannelFactory.class, Mockito.withSettings().withoutAnnotations()); Mockito.when(channelFactory.createSingleChannel()).thenReturn(initial, rotated); Mockito.when(initial.newCall(Mockito.any(), Mockito.any())) diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProvider.java b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProvider.java index 730ae80bb8b9..11d1856b5d4b 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProvider.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProvider.java @@ -31,6 +31,7 @@ import com.google.api.client.http.HttpTransport; import com.google.api.client.http.javanet.NetHttpTransport; +import com.google.api.client.util.SslUtils; import com.google.api.core.InternalExtensionOnly; import com.google.api.gax.core.ExecutorProvider; import com.google.api.gax.rpc.FixedHeaderProvider; @@ -45,11 +46,13 @@ import java.io.IOException; import java.security.GeneralSecurityException; import java.security.KeyStore; +import java.security.Provider; import java.util.Map; import java.util.concurrent.Executor; import java.util.concurrent.ScheduledExecutorService; import java.util.logging.Level; import java.util.logging.Logger; +import javax.net.ssl.SSLContext; import org.jspecify.annotations.NullMarked; import org.jspecify.annotations.Nullable; @@ -200,8 +203,20 @@ public TransportChannelProvider withCredentials(Credentials credentials) { KeyStore mtlsKeyStore = mtlsProvider.getKeyStore(); if (mtlsKeyStore != null) { NetHttpTransport.Builder builder = new NetHttpTransport.Builder(); - HttpJsonConscryptUtils.configureConscryptSecurityProvider(builder); builder.trustCertificates(null, mtlsKeyStore, ""); + Provider conscryptProvider = HttpJsonConscryptUtils.getConscryptProvider(); + if (conscryptProvider != null) { + SSLContext sslContext = SSLContext.getInstance("TLS", conscryptProvider); + SslUtils.initSslContext( + sslContext, + null, + SslUtils.getPkixTrustManagerFactory(), + mtlsKeyStore, + "", + SslUtils.getDefaultKeyManagerFactory()); + builder.setSslSocketFactory(sslContext.getSocketFactory()); + } + HttpJsonConscryptUtils.configureConscryptSecurityProvider(builder); return builder.build(); } } diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProviderTest.java b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProviderTest.java index 833c5ff6f5e8..a84723fbdd64 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProviderTest.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProviderTest.java @@ -335,6 +335,30 @@ void createHttpTransport_withMtlsAndConscrypt_configuresSecurityProvider() assertThat(transport).isInstanceOf(com.google.api.client.http.javanet.NetHttpTransport.class); } + @Test + void createHttpTransport_whenMtlsProviderNullOrNotUsingClientCert_returnsNull() + throws IOException, GeneralSecurityException { + InstantiatingHttpJsonChannelProvider nullMtlsProviderChannelProvider = + InstantiatingHttpJsonChannelProvider.newBuilder() + .setEndpoint(DEFAULT_ENDPOINT) + .setMtlsProvider(null) + .setCertificateBasedAccess(certificateBasedAccess) + .build(); + assertThat(nullMtlsProviderChannelProvider.createHttpTransport()).isNull(); + + Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(false); + com.google.auth.mtls.MtlsProvider provider = + new com.google.api.gax.rpc.testing.FakeMtlsProvider( + com.google.api.gax.rpc.testing.FakeMtlsProvider.createTestMtlsKeyStore(), "", false); + InstantiatingHttpJsonChannelProvider disabledMtlsChannelProvider = + InstantiatingHttpJsonChannelProvider.newBuilder() + .setEndpoint(DEFAULT_ENDPOINT) + .setMtlsProvider(provider) + .setCertificateBasedAccess(certificateBasedAccess) + .build(); + assertThat(disabledMtlsChannelProvider.createHttpTransport()).isNull(); + } + @Override protected Object getMtlsObjectFromTransportChannel( MtlsProvider provider, CertificateBasedAccess certificateBasedAccess) From 6ec0bc917db7acc2befa48fbe3209b464e615935 Mon Sep 17 00:00:00 2001 From: Matt Castelaz Date: Fri, 18 Sep 2026 18:20:09 +0000 Subject: [PATCH 16/29] fix(gax-httpjson): add withoutAnnotations() to MtlsProvider mock for Java 8 (#13995) --- .../gax/httpjson/InstantiatingHttpJsonChannelProviderTest.java | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProviderTest.java b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProviderTest.java index a84723fbdd64..f1a542a7c75d 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProviderTest.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProviderTest.java @@ -274,7 +274,8 @@ void getTransportChannel_whenMtlsKeyStoreThrowsIOException_throwsCheckedIOExcept throws Exception { Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); Mockito.when(certificateBasedAccess.getWorkloadCertPath()).thenReturn("fake/cert/path.json"); - MtlsProvider failingMtlsProvider = Mockito.mock(MtlsProvider.class); + MtlsProvider failingMtlsProvider = + Mockito.mock(MtlsProvider.class, Mockito.withSettings().withoutAnnotations()); Mockito.when(failingMtlsProvider.getKeyStore()) .thenThrow(new IOException("Simulated keystore read failure")); From 2f549e2a4e0e2a93b35f24e039b6de790e696be8 Mon Sep 17 00:00:00 2001 From: Matt Castelaz Date: Fri, 18 Sep 2026 19:57:32 +0000 Subject: [PATCH 17/29] refactor(gax): make RetryingContext overload primary and simplify streaming exception wrap (#13995) --- .../api/gax/rpc/ApiResultRetryAlgorithm.java | 18 +++++++++--------- .../rpc/ServerStreamingAttemptCallable.java | 9 +-------- 2 files changed, 10 insertions(+), 17 deletions(-) diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ApiResultRetryAlgorithm.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ApiResultRetryAlgorithm.java index 1c80f1d90d59..010e6462d09d 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ApiResultRetryAlgorithm.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ApiResultRetryAlgorithm.java @@ -45,6 +45,15 @@ class ApiResultRetryAlgorithm extends BasicResultRetryAlgorithm extends BasicResultRetryAlgorithm Date: Wed, 30 Sep 2026 19:08:15 +0000 Subject: [PATCH 18/29] fix(gax,gax-grpc,gax-httpjson): restore timeout clearing and preserve stack trace - Restore the base-branch guard in GrpcCallContext/HttpJsonCallContext withTimeoutDuration so a null or zero timeout clears an existing one. - Copy the original stack trace onto the wrapped UnauthenticatedException in ServerStreamingAttemptCallable, matching AttemptCallable. --- .../com/google/api/gax/grpc/GrpcCallContext.java | 2 +- .../google/api/gax/grpc/GrpcCallContextTest.java | 16 ++++++++++++++++ .../api/gax/httpjson/HttpJsonCallContext.java | 2 +- .../gax/httpjson/HttpJsonCallContextTest.java | 14 ++++++++++++++ .../gax/rpc/ServerStreamingAttemptCallable.java | 1 + .../rpc/ServerStreamingAttemptCallableTest.java | 1 + 6 files changed, 34 insertions(+), 2 deletions(-) diff --git a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/GrpcCallContext.java b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/GrpcCallContext.java index c0288a1ae382..fab685328931 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/GrpcCallContext.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/GrpcCallContext.java @@ -284,7 +284,7 @@ public GrpcCallContext withTimeoutDuration(java.time.@Nullable Duration timeout) } // Prevent expanding timeouts - if (this.timeout != null && (timeout == null || this.timeout.compareTo(timeout) <= 0)) { + if (timeout != null && this.timeout != null && this.timeout.compareTo(timeout) <= 0) { return this; } diff --git a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/GrpcCallContextTest.java b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/GrpcCallContextTest.java index cbaa7af2475c..55f473407ec8 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/GrpcCallContextTest.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/GrpcCallContextTest.java @@ -210,6 +210,22 @@ void testWithLongerTimeout() { .isEqualTo(java.time.Duration.ofSeconds(5)); } + @Test + void testWithNullOrZeroTimeoutClearsExistingTimeout() { + GrpcCallContext ctxWithTimeout = + GrpcCallContext.createDefault().withTimeoutDuration(java.time.Duration.ofSeconds(5)); + + // Sanity check + Truth.assertThat(ctxWithTimeout.getTimeoutDuration()) + .isEqualTo(java.time.Duration.ofSeconds(5)); + + java.time.Duration nullTimeout = null; + Truth.assertThat(ctxWithTimeout.withTimeoutDuration(nullTimeout).getTimeoutDuration()).isNull(); + Truth.assertThat( + ctxWithTimeout.withTimeoutDuration(java.time.Duration.ZERO).getTimeoutDuration()) + .isNull(); + } + @Test void testMergeWithNullTimeout() { java.time.Duration timeout = java.time.Duration.ofSeconds(10); diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/HttpJsonCallContext.java b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/HttpJsonCallContext.java index 1944b6c9a6b3..81e970af4b06 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/HttpJsonCallContext.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/HttpJsonCallContext.java @@ -316,7 +316,7 @@ public HttpJsonCallContext withTimeoutDuration(java.time.Duration timeout) { } // Prevent expanding deadlines - if (this.timeout != null && (timeout == null || this.timeout.compareTo(timeout) <= 0)) { + if (timeout != null && this.timeout != null && this.timeout.compareTo(timeout) <= 0) { return this; } diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/HttpJsonCallContextTest.java b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/HttpJsonCallContextTest.java index 7ca2de0775c0..ce7b6682e586 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/HttpJsonCallContextTest.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/HttpJsonCallContextTest.java @@ -207,6 +207,20 @@ void testWithLongerTimeout() { .isEqualTo(java.time.Duration.ofSeconds(5)); } + @Test + void testWithNullOrZeroTimeoutClearsExistingTimeout() { + HttpJsonCallContext ctxWithTimeout = + HttpJsonCallContext.createDefault().withTimeoutDuration(java.time.Duration.ofSeconds(5)); + + // Sanity check + Truth.assertThat(ctxWithTimeout.getTimeoutDuration()) + .isEqualTo(java.time.Duration.ofSeconds(5)); + + java.time.Duration nullTimeout = null; + assertNull(ctxWithTimeout.withTimeoutDuration(nullTimeout).getTimeoutDuration()); + assertNull(ctxWithTimeout.withTimeoutDuration(java.time.Duration.ZERO).getTimeoutDuration()); + } + @Test void testMergeWithNullTimeout() { java.time.Duration timeout = java.time.Duration.ofSeconds(10); diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ServerStreamingAttemptCallable.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ServerStreamingAttemptCallable.java index bf1fdef29c05..b8ccd565ac5b 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ServerStreamingAttemptCallable.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ServerStreamingAttemptCallable.java @@ -271,6 +271,7 @@ public void onErrorImpl(Throwable t) { unauthenticatedException.getStatusCode(), true, unauthenticatedException.getErrorDetails()); + newEx.setStackTrace(unauthenticatedException.getStackTrace()); for (Throwable suppressed : unauthenticatedException.getSuppressed()) { newEx.addSuppressed(suppressed); } diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ServerStreamingAttemptCallableTest.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ServerStreamingAttemptCallableTest.java index 740dc37b7ffe..66fc73a7d2bf 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ServerStreamingAttemptCallableTest.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ServerStreamingAttemptCallableTest.java @@ -337,6 +337,7 @@ void testUnauthenticatedRefreshWithGenerationAdvanceRetries() { Truth.assertThat(((ServerStreamingAttemptException) outerError).canResume()).isTrue(); Truth.assertThat(outerError.getCause()).isInstanceOf(UnauthenticatedException.class); Truth.assertThat(((UnauthenticatedException) outerError.getCause()).isRetryable()).isTrue(); + Truth.assertThat(outerError.getCause().getStackTrace()).isEqualTo(initialError.getStackTrace()); // Verify retry call resumes stream callable.call(); From 82841ffd91bedad904ffc351a9f5476ac5648e2a Mon Sep 17 00:00:00 2001 From: Matt Castelaz Date: Wed, 30 Sep 2026 19:08:16 +0000 Subject: [PATCH 19/29] fix(gax-httpjson): record cert baseline before initial channel creation and fix refresh ordering - RefreshingHttpJsonChannel now initializes its CertificateRotationTracker before creating the initial channel, so a rotation during startup is detected. The provider lets RefreshingHttpJsonChannel create the initial channel, unwrapping checked IOException/GeneralSecurityException to keep the getTransportChannel() contract. - In refresh(), increment the generation before marking the tracker refreshed so a failing RPC that sees the new fingerprint also sees the new generation. - Catch and log channel factory failures in refresh(), keeping the old channel. --- .../InstantiatingHttpJsonChannelProvider.java | 30 ++++++++--- .../httpjson/RefreshingHttpJsonChannel.java | 33 ++++++------ ...tantiatingHttpJsonChannelProviderTest.java | 51 +++++++++++++++++++ .../RefreshingHttpJsonChannelTest.java | 28 +++++++++- 4 files changed, 115 insertions(+), 27 deletions(-) diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProvider.java b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProvider.java index 11d1856b5d4b..48e4d9ca5579 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProvider.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProvider.java @@ -249,22 +249,36 @@ private HttpJsonTransportChannel createChannel() throws IOException, GeneralSecu && certificateBasedAccess.useMtlsClientCertificate(); String workloadCertPath = isMtlsActive ? certificateBasedAccess.getWorkloadCertPath() : null; - ManagedHttpJsonChannel initialChannel = createSingleManagedChannel(); - try { + ManagedHttpJsonChannel baseChannel; + if (workloadCertPath != null) { java.util.function.Supplier channelFactory = () -> { try { return createSingleManagedChannel(); - } catch (Exception e) { + } catch (IOException | GeneralSecurityException e) { throw new java.lang.RuntimeException( "Failed to create fresh ManagedHttpJsonChannel", e); } }; + // RefreshingHttpJsonChannel records the baseline certificate fingerprint before creating the + // initial channel, so a rotation during startup is detected on the next auth failure. + try { + baseChannel = new RefreshingHttpJsonChannel(channelFactory, workloadCertPath); + } catch (RuntimeException e) { + if (e.getCause() instanceof IOException) { + throw (IOException) e.getCause(); + } + if (e.getCause() instanceof GeneralSecurityException) { + throw (GeneralSecurityException) e.getCause(); + } + throw e; + } + } else { + baseChannel = createSingleManagedChannel(); + } - ManagedHttpJsonChannel channel = - workloadCertPath != null - ? new RefreshingHttpJsonChannel(initialChannel, channelFactory, workloadCertPath) - : initialChannel; + try { + ManagedHttpJsonChannel channel = baseChannel; HttpJsonClientInterceptor headerInterceptor = new HttpJsonHeaderInterceptor(headerProvider.getHeaders()); @@ -279,7 +293,7 @@ private HttpJsonTransportChannel createChannel() throws IOException, GeneralSecu return HttpJsonTransportChannel.newBuilder().setManagedChannel(channel).build(); } catch (Throwable t) { - initialChannel.shutdownNow(); + baseChannel.shutdownNow(); throw t; } } diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java index 3d9557152845..8bc25d14d972 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java @@ -69,27 +69,16 @@ public class RefreshingHttpJsonChannel extends ManagedHttpJsonChannel { public RefreshingHttpJsonChannel( Supplier channelFactory, String workloadCertPath) { - this(channelFactory.get(), channelFactory, workloadCertPath); - } - - public RefreshingHttpJsonChannel( - ManagedHttpJsonChannel initialChannel, - Supplier channelFactory, - String workloadCertPath) { super(true); this.channelFactory = channelFactory; this.workloadCertPath = workloadCertPath; - ChannelEntry initial = new ChannelEntry(initialChannel); + // Record the baseline fingerprint before the initial channel loads the certificate from disk, + // so a rotation between the two steps is detected as a fingerprint change. + this.rotationTracker = + new CertificateRotationTracker(this::getWorkloadCertPath, this::getCertificateFingerprint); + ChannelEntry initial = new ChannelEntry(channelFactory.get()); this.activeEntry = new AtomicReference<>(initial); this.allEntries.add(initial); - try { - this.rotationTracker = - new CertificateRotationTracker( - this::getWorkloadCertPath, this::getCertificateFingerprint); - } catch (Throwable t) { - initialChannel.shutdownNow(); - throw t; - } } // Visible for testing @@ -128,14 +117,22 @@ public void refresh() { LOG.info("mTLS certificate rotation detected. Triggering HTTP/JSON channel pool refresh."); - ChannelEntry newEntry = new ChannelEntry(channelFactory.get()); + ChannelEntry newEntry; + try { + newEntry = new ChannelEntry(channelFactory.get()); + } catch (Exception e) { + LOG.log(Level.WARNING, "Failed to refresh HTTP/JSON channel, leaving old channel", e); + return; + } allEntries.add(newEntry); // Prune terminated entries after adding newEntry to ensure allEntries is never empty allEntries.removeIf(entry -> entry != newEntry && entry.channel.isTerminated()); ChannelEntry oldEntry = activeEntry.getAndSet(newEntry); - rotationTracker.markRefreshed(currentDiskFingerprint); + // Order matters: swap activeEntry, then bump generation, then mark the tracker refreshed, so + // any failing RPC that observes the new fingerprint also observes the new generation. generation.incrementAndGet(); + rotationTracker.markRefreshed(currentDiskFingerprint); if (oldEntry != null) { oldEntry.requestShutdown(); diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProviderTest.java b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProviderTest.java index f1a542a7c75d..1d967f4b108b 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProviderTest.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProviderTest.java @@ -316,6 +316,57 @@ void getTransportChannel_whenMtlsActiveAndKeyStoreNull_throwsIOException() { assertThat(thrown).hasMessageThat().contains("Failed to initialize mTLS HttpTransport"); } + @Test + void getTransportChannel_whenKeyStoreUninitialized_causeIsSecurityException() throws Exception { + Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); + Mockito.when(certificateBasedAccess.getWorkloadCertPath()).thenReturn("fake/cert/path.json"); + // A KeyStore that was never loaded makes transport creation fail with a KeyStoreException. + java.security.KeyStore uninitializedKeyStore = java.security.KeyStore.getInstance("PKCS12"); + com.google.auth.mtls.MtlsProvider providerWithBadKeyStore = + new com.google.api.gax.rpc.testing.FakeMtlsProvider(uninitializedKeyStore, "", false); + + InstantiatingHttpJsonChannelProvider provider = + InstantiatingHttpJsonChannelProvider.newBuilder() + .setEndpoint(DEFAULT_ENDPOINT) + .setMtlsProvider(providerWithBadKeyStore) + .setCertificateBasedAccess(certificateBasedAccess) + .build(); + InstantiatingHttpJsonChannelProvider finalProvider = + (InstantiatingHttpJsonChannelProvider) provider.withHeaders(DEFAULT_HEADER_MAP); + + // The GeneralSecurityException must be unwrapped from the factory's RuntimeException, so that + // getTransportChannel() surfaces it as the direct cause of its checked IOException. + IOException thrown = + org.junit.jupiter.api.Assertions.assertThrows( + IOException.class, finalProvider::getTransportChannel); + assertThat(thrown).hasCauseThat().isInstanceOf(GeneralSecurityException.class); + } + + @Test + void getTransportChannel_whenMtlsKeyStoreThrowsRuntimeException_propagatesUnwrapped() + throws Exception { + Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); + Mockito.when(certificateBasedAccess.getWorkloadCertPath()).thenReturn("fake/cert/path.json"); + IllegalStateException failure = new IllegalStateException("Simulated keystore failure"); + MtlsProvider failingMtlsProvider = + Mockito.mock(MtlsProvider.class, Mockito.withSettings().withoutAnnotations()); + Mockito.when(failingMtlsProvider.getKeyStore()).thenThrow(failure); + + InstantiatingHttpJsonChannelProvider provider = + InstantiatingHttpJsonChannelProvider.newBuilder() + .setEndpoint(DEFAULT_ENDPOINT) + .setMtlsProvider(failingMtlsProvider) + .setCertificateBasedAccess(certificateBasedAccess) + .build(); + InstantiatingHttpJsonChannelProvider finalProvider = + (InstantiatingHttpJsonChannelProvider) provider.withHeaders(DEFAULT_HEADER_MAP); + + IllegalStateException thrown = + org.junit.jupiter.api.Assertions.assertThrows( + IllegalStateException.class, finalProvider::getTransportChannel); + assertThat(thrown).isSameInstanceAs(failure); + } + @Test void createHttpTransport_withMtlsAndConscrypt_configuresSecurityProvider() throws IOException, GeneralSecurityException { diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java index aaaacd3c53db..54fcf80eea40 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java @@ -211,6 +211,29 @@ void shouldRefresh_doesNotCacheNegativeResultAndDetectsSubsequentRotationImmedia assertFalse(channel.shouldRefresh()); } + @Test + void rotationDuringInitialChannelCreation_isDetectedAndRefreshed() { + channelFactory = + () -> { + if (channelFactoryCount.incrementAndGet() == 1) { + // Simulate the certificate rotating on disk while the initial channel loads it. + testFingerprint = "fingerprint2"; + } + lastCreatedChannel = new FakeManagedHttpJsonChannel(); + return lastCreatedChannel; + }; + + RefreshingHttpJsonChannel channel = createTestChannel(); + + // The baseline was recorded before the initial channel was created, so the rotation is seen. + assertTrue(channel.shouldRefresh()); + + channel.refresh(); + assertEquals(2, channelFactoryCount.get()); + assertEquals(1, channel.getGeneration()); + assertFalse(channel.shouldRefresh()); + } + @Test void testRefreshSwapsChannel() throws InterruptedException { RefreshingHttpJsonChannel channel = createTestChannel(); @@ -322,7 +345,10 @@ void testRefreshFactoryExceptionDoesNotWedgeFingerprint() throws InterruptedExce channel.invalidateDiskFingerprintCache(); // Invalidate 1-second cache testFingerprint = "fingerprint2"; - assertThrows(RuntimeException.class, channel::refresh); + // Factory failure is logged and the existing channel is kept + channel.refresh(); + assertEquals(1, channelFactoryCount.get()); + assertEquals(0, channel.getGeneration()); // Because factory threw, activeCertFingerprint should NOT be updated to fingerprint2 // Therefore shouldRefresh() should still return true From baca8b14ec38e220e8c03e18ddb5222937f0b674 Mon Sep 17 00:00:00 2001 From: Matt Castelaz Date: Wed, 30 Sep 2026 19:08:16 +0000 Subject: [PATCH 20/29] fix(gax): stop retrying after the single rotation retry is used Return explicitly exhausted settings from ApiResultRetryAlgorithm createNextAttempt once the free certificate-rotation retry has been used, so a subsequent rotation-marked UNAUTHENTICATED failure stops instead of falling back to exponential backoff. --- .../api/gax/rpc/ApiResultRetryAlgorithm.java | 14 ++++ .../gax/rpc/ApiResultRetryAlgorithmTest.java | 74 +++++++++++++++++++ 2 files changed, 88 insertions(+) diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ApiResultRetryAlgorithm.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ApiResultRetryAlgorithm.java index 010e6462d09d..91b332d10660 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ApiResultRetryAlgorithm.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ApiResultRetryAlgorithm.java @@ -70,6 +70,20 @@ class ApiResultRetryAlgorithm extends BasicResultRetryAlgorithm resultAlgorithm = new ApiResultRetryAlgorithm<>(); + ExponentialRetryAlgorithm timedAlgorithm = + new ExponentialRetryAlgorithm(settings, NanoClock.getDefaultClock()); + RetryAlgorithm retryAlgorithm = new RetryAlgorithm<>(resultAlgorithm, timedAlgorithm); + + TimedAttemptSettings firstAttempt = retryAlgorithm.createFirstAttempt(context); + UnauthenticatedException rotationEx = + new UnauthenticatedException( + "Expired cert", null, new FakeStatusCode(Code.UNAUTHENTICATED), /* retryable= */ true); + + TimedAttemptSettings rotationRetry = + retryAlgorithm.createNextAttempt(context, rotationEx, null, firstAttempt); + assertTrue(retryAlgorithm.shouldRetry(context, rotationEx, null, rotationRetry)); + + // A second rotation-marked failure must stop immediately instead of retrying with backoff + // until the total timeout expires. + TimedAttemptSettings afterSecondFailure = + retryAlgorithm.createNextAttempt(context, rotationEx, null, rotationRetry); + assertNotNull(afterSecondFailure); + assertFalse(retryAlgorithm.shouldRetry(context, rotationEx, null, afterSecondFailure)); + } + + @Test + void testSecondRotationFailureDoesNotConsumeNormalRetryBudget() { + ApiCallContext context = + mock(ApiCallContext.class, Mockito.withSettings().withoutAnnotations()); + when(context.getRetryableCodes()).thenReturn(Sets.newHashSet(Code.UNAVAILABLE)); + + RetrySettings settings = + RetrySettings.newBuilder() + .setMaxAttempts(5) + .setInitialRetryDelayDuration(Duration.ofMillis(100)) + .setRetryDelayMultiplier(2.0) + .setMaxRetryDelayDuration(Duration.ofSeconds(1)) + .setTotalTimeoutDuration(Duration.ofMinutes(10)) + .build(); + + ApiResultRetryAlgorithm resultAlgorithm = new ApiResultRetryAlgorithm<>(); + ExponentialRetryAlgorithm timedAlgorithm = + new ExponentialRetryAlgorithm(settings, NanoClock.getDefaultClock()); + RetryAlgorithm retryAlgorithm = new RetryAlgorithm<>(resultAlgorithm, timedAlgorithm); + + TimedAttemptSettings firstAttempt = retryAlgorithm.createFirstAttempt(context); + UnauthenticatedException rotationEx = + new UnauthenticatedException( + "Expired cert", null, new FakeStatusCode(Code.UNAUTHENTICATED), /* retryable= */ true); + + TimedAttemptSettings rotationRetry = + retryAlgorithm.createNextAttempt(context, rotationEx, null, firstAttempt); + assertTrue(retryAlgorithm.shouldRetry(context, rotationEx, null, rotationRetry)); + + // UNAUTHENTICATED is not a retryable code, so a second rotation-marked failure must stop + // instead of consuming the method's remaining attempts. + TimedAttemptSettings afterSecondFailure = + retryAlgorithm.createNextAttempt(context, rotationEx, null, rotationRetry); + assertNotNull(afterSecondFailure); + assertFalse(retryAlgorithm.shouldRetry(context, rotationEx, null, afterSecondFailure)); + } } From f64c78000f1af78ff8c81c5d12b72e783dd4e8f5 Mon Sep 17 00:00:00 2001 From: Matt Castelaz Date: Wed, 30 Sep 2026 19:08:17 +0000 Subject: [PATCH 21/29] fix(gax-grpc): drop channels that fail to refresh on rotation On a rotation-triggered refresh, ChannelPool now keeps only successfully recreated channels so no traffic is routed to the old certificate, and marks the certificate refreshed once at least one channel was recreated. Statically sized pools schedule a one-time refill; dynamically sized pools are regrown by resize(). Non-rotation refreshes keep per-slot fallback. --- .../com/google/api/gax/grpc/ChannelPool.java | 68 ++++- .../google/api/gax/grpc/ChannelPoolTest.java | 271 +++++++++++++++--- 2 files changed, 298 insertions(+), 41 deletions(-) diff --git a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java index 8004593a2673..e81fc5867f82 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java @@ -50,6 +50,7 @@ import java.util.List; import java.util.concurrent.CancellationException; import java.util.concurrent.Executors; +import java.util.concurrent.RejectedExecutionException; import java.util.concurrent.ScheduledExecutorService; import java.util.concurrent.ScheduledFuture; import java.util.concurrent.TimeUnit; @@ -516,7 +517,8 @@ void refresh() { return; } - if (refreshAll()) { + // Drop any channel that fails to refresh so that no traffic is routed to the old certificate. + if (refreshAll(/* dropUnrefreshedChannels= */ true)) { rotationTracker.markRefreshed(currentDiskFingerprint); } } @@ -529,6 +531,23 @@ void refresh() { @InternalApi("Visible for testing") boolean refreshAll() { + return refreshAll(/* dropUnrefreshedChannels= */ false); + } + + /** + * Replaces the channels in the pool with freshly created ones. + * + * @param dropUnrefreshedChannels if {@code false}, a channel that fails to be recreated keeps its + * slot in the pool. If {@code true} (used for certificate rotation), only newly created + * channels are kept so no traffic is routed to a channel using the old certificate; a + * statically sized pool is then refilled asynchronously, while a dynamically sized pool is + * refilled by {@link #resize()}. + * @return if {@code dropUnrefreshedChannels} is {@code false}, whether every channel was + * recreated; otherwise, whether at least one channel was recreated (i.e. every channel left + * in the pool was newly created) + */ + @InternalApi("Visible for testing") + boolean refreshAll(boolean dropUnrefreshedChannels) { synchronized (entryWriteLock) { if (isShutdown) { return false; @@ -553,7 +572,12 @@ boolean refreshAll() { anyCreated = true; } catch (Exception e) { allCreated = false; - LOG.log(Level.WARNING, "Failed to refresh channel, leaving old channel", e); + LOG.log( + Level.WARNING, + dropUnrefreshedChannels + ? "Failed to refresh channel, dropping old channel" + : "Failed to refresh channel, leaving old channel", + e); } } @@ -561,17 +585,22 @@ boolean refreshAll() { return false; } - ImmutableList replacedEntries = entries.getAndSet(ImmutableList.copyOf(newEntries)); + ImmutableList finalEntries = + ImmutableList.copyOf(dropUnrefreshedChannels ? createdEntries : newEntries); + ImmutableList replacedEntries = entries.getAndSet(finalEntries); createdEntries.clear(); // Ownership transferred to pool // Shutdown the channels that were cycled out. for (Entry e : replacedEntries) { - if (!newEntries.contains(e)) { + if (!finalEntries.contains(e)) { e.requestShutdown(); } } generation.incrementAndGet(); - return allCreated; + if (dropUnrefreshedChannels && !allCreated && settings.isStaticSize()) { + scheduleRefill(); + } + return dropUnrefreshedChannels || allCreated; } finally { // If an Error aborted before getAndSet, shut down newly created channels so they don't leak for (Entry e : createdEntries) { @@ -581,6 +610,35 @@ boolean refreshAll() { } } + /** + * Schedules a one-shot task that restores a statically sized pool to its configured channel count + * after a certificate rotation refresh dropped channels that failed to refresh. + */ + private void scheduleRefill() { + try { + backgroundExecutorProvider.getExecutor().execute(this::refillSafely); + } catch (RejectedExecutionException e) { + LOG.log(Level.WARNING, "Failed to schedule channel pool refill", e); + } + } + + @VisibleForTesting + void refillSafely() { + try { + synchronized (entryWriteLock) { + if (isShutdown) { + return; + } + int targetSize = settings.getInitialChannelCount(); + if (entries.get().size() < targetSize) { + expand(targetSize); + } + } + } catch (Exception e) { + LOG.log(Level.WARNING, "Failed to refill channel pool", e); + } + } + /** * Returns the current channel pool generation counter. * diff --git a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java index 9c408f2d128c..5bc6f391065c 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java @@ -67,6 +67,7 @@ import java.util.concurrent.CancellationException; import java.util.concurrent.ExecutorService; import java.util.concurrent.Executors; +import java.util.concurrent.RejectedExecutionException; import java.util.concurrent.ScheduledExecutorService; import java.util.concurrent.ScheduledFuture; import java.util.concurrent.TimeUnit; @@ -566,66 +567,264 @@ void channelReactiveMTlsRefresh_failedCreationDoesNotMutateFingerprintAndAllowsR .newCall(Mockito.>any(), Mockito.any(CallOptions.class)); } + private void writeCert(String resourceName) throws IOException { + if (tempCert == null) { + tempCert = java.nio.file.Files.createTempFile("cert", ".pem"); + } + java.nio.file.Files.copy( + java.nio.file.Paths.get("src", "test", "resources", resourceName), + tempCert, + java.nio.file.StandardCopyOption.REPLACE_EXISTING); + } + + /** Creates an mTLS pool backed by {@code executor} and then rotates the certificate on disk. */ + private ChannelPool createMtlsPoolAndRotateCert( + ChannelPoolSettings settings, + ChannelFactory channelFactory, + ScheduledExecutorService executor) + throws IOException { + writeCert("client_cert.pem"); + pool = + new ChannelPool( + settings, channelFactory, FixedExecutorProvider.create(executor), tempCert.toString()); + pool.invalidateDiskFingerprintCache(); + writeCert("root_cert.pem"); + assertThat(pool.shouldRefresh()).isTrue(); + return pool; + } + + private static ScheduledExecutorService mockExecutor() { + return Mockito.mock( + ScheduledExecutorService.class, Mockito.withSettings().withoutAnnotations()); + } + + private static ChannelFactory mockChannelFactory() { + return Mockito.mock(ChannelFactory.class, Mockito.withSettings().withoutAnnotations()); + } + @Test - void - channelReactiveMTlsRefresh_partialFailureInMultiChannelPool_retainsShouldRefreshAndCompletesOnSubsequentRefresh() - throws IOException { + void channelReactiveMTlsRefresh_partialFailureInStaticPool_dropsUnrefreshedChannelsAndRefills() + throws IOException { + ScheduledExecutorService executor = mockExecutor(); ManagedChannel initial1 = Mockito.mock(ManagedChannel.class); ManagedChannel initial2 = Mockito.mock(ManagedChannel.class); ManagedChannel rotated1 = Mockito.mock(ManagedChannel.class); - ManagedChannel rotated1SecondPass = Mockito.mock(ManagedChannel.class); - ManagedChannel rotated2 = Mockito.mock(ManagedChannel.class); - ChannelFactory channelFactory = - Mockito.mock(ChannelFactory.class, Mockito.withSettings().withoutAnnotations()); - - // Initial creation: initial1, initial2 - // Refresh pass 1: rotated1 succeeds, second throws IOException - // Refresh pass 2: rotated1SecondPass, rotated2 both succeed + ManagedChannel refilled = Mockito.mock(ManagedChannel.class); + ChannelFactory channelFactory = mockChannelFactory(); Mockito.when(channelFactory.createSingleChannel()) .thenReturn(initial1, initial2) .thenReturn(rotated1) .thenThrow(new IOException("Transient failure on second sub-channel")) - .thenReturn(rotated1SecondPass, rotated2); + .thenReturn(refilled); - tempCert = java.nio.file.Files.createTempFile("cert", ".pem"); - java.nio.file.Path clientCert = - java.nio.file.Paths.get("src", "test", "resources", "client_cert.pem"); - java.nio.file.Files.copy( - clientCert, tempCert, java.nio.file.StandardCopyOption.REPLACE_EXISTING); + createMtlsPoolAndRotateCert(ChannelPoolSettings.staticallySized(2), channelFactory, executor); + long genBefore = pool.getGeneration(); - pool = - ChannelPool.create( - ChannelPoolSettings.staticallySized(2), channelFactory, null, tempCert.toString()); + pool.refresh(); - // Rotate cert on disk + // Only the newly created channel is kept; both channels on the old certificate are retired. + assertThat(pool.entries.get()).hasSize(1); + Mockito.verify(initial1).shutdown(); + Mockito.verify(initial2).shutdown(); + Mockito.verify(rotated1, Mockito.never()).shutdown(); + pool.newCall(FakeMethodDescriptor.create(), CallOptions.DEFAULT); + Mockito.verify(rotated1) + .newCall(Mockito.>any(), Mockito.any(CallOptions.class)); + assertThat(pool.getGeneration()).isEqualTo(genBefore + 1); + // Every channel left in the pool uses the new certificate, so it is recorded as active. pool.invalidateDiskFingerprintCache(); - java.nio.file.Path rootCert = - java.nio.file.Paths.get("src", "test", "resources", "root_cert.pem"); - java.nio.file.Files.copy(rootCert, tempCert, java.nio.file.StandardCopyOption.REPLACE_EXISTING); + assertThat(pool.shouldRefresh()).isFalse(); - assertThat(pool.shouldRefresh()).isTrue(); - long genBefore = pool.getGeneration(); + // A one-time refill restores the configured size. + ArgumentCaptor refillTask = ArgumentCaptor.forClass(Runnable.class); + Mockito.verify(executor).execute(refillTask.capture()); + refillTask.getValue().run(); + assertThat(pool.entries.get()).hasSize(2); + Mockito.verify(channelFactory, Mockito.times(5)).createSingleChannel(); + + // Running the refill again on a full pool creates no channels. + refillTask.getValue().run(); + assertThat(pool.entries.get()).hasSize(2); + Mockito.verify(channelFactory, Mockito.times(5)).createSingleChannel(); + Mockito.verify(refilled, Mockito.never()).shutdown(); + Mockito.verify(executor, Mockito.times(1)).execute(Mockito.any(Runnable.class)); + } + + @Test + void channelReactiveMTlsRefresh_fullSuccess_doesNotScheduleRefill() throws IOException { + ScheduledExecutorService executor = mockExecutor(); + ManagedChannel initial1 = Mockito.mock(ManagedChannel.class); + ManagedChannel initial2 = Mockito.mock(ManagedChannel.class); + ManagedChannel rotated1 = Mockito.mock(ManagedChannel.class); + ManagedChannel rotated2 = Mockito.mock(ManagedChannel.class); + ChannelFactory channelFactory = mockChannelFactory(); + Mockito.when(channelFactory.createSingleChannel()) + .thenReturn(initial1, initial2, rotated1, rotated2); + + createMtlsPoolAndRotateCert(ChannelPoolSettings.staticallySized(2), channelFactory, executor); - // First refresh: partial failure (channel 0 rotates to rotated1, channel 1 fails and keeps - // initial2) pool.refresh(); - // Generation should still increment since partial progress was committed - assertThat(pool.getGeneration()).isGreaterThan(genBefore); - // initial1 should have been shut down, initial2 should NOT be shut down yet + assertThat(pool.entries.get()).hasSize(2); + Mockito.verify(initial1).shutdown(); + Mockito.verify(initial2).shutdown(); + pool.invalidateDiskFingerprintCache(); + assertThat(pool.shouldRefresh()).isFalse(); + Mockito.verify(executor, Mockito.never()).execute(Mockito.any(Runnable.class)); + } + + @Test + void refreshAll_partialFailureWithoutRotation_keepsOldChannelAndDoesNotScheduleRefill() + throws IOException { + // Non-rotation refreshes (e.g. the preemptive refresh) keep per-slot fallback even for mTLS. + ScheduledExecutorService executor = mockExecutor(); + ManagedChannel initial1 = Mockito.mock(ManagedChannel.class); + ManagedChannel initial2 = Mockito.mock(ManagedChannel.class); + ManagedChannel refreshed1 = Mockito.mock(ManagedChannel.class); + ChannelFactory channelFactory = mockChannelFactory(); + Mockito.when(channelFactory.createSingleChannel()) + .thenReturn(initial1, initial2) + .thenReturn(refreshed1) + .thenThrow(new IOException("Transient failure on second sub-channel")); + + writeCert("client_cert.pem"); + pool = + new ChannelPool( + ChannelPoolSettings.staticallySized(2), + channelFactory, + FixedExecutorProvider.create(executor), + tempCert.toString()); + long genBefore = pool.getGeneration(); + + assertThat(pool.refreshAll()).isFalse(); + + assertThat(pool.entries.get()).hasSize(2); Mockito.verify(initial1).shutdown(); Mockito.verify(initial2, Mockito.never()).shutdown(); + assertThat(pool.getGeneration()).isEqualTo(genBefore + 1); + Mockito.verify(executor, Mockito.never()).execute(Mockito.any(Runnable.class)); + } - // Crucial assertion: shouldRefresh() MUST remain true so subsequent 401s on unrotated channel 1 - // trigger retry/refresh - pool.invalidateDiskFingerprintCache(); - assertThat(pool.shouldRefresh()).isTrue(); + @Test + void channelReactiveMTlsRefresh_partialFailureInDynamicPool_refilledByResize() + throws IOException { + ScheduledExecutorService executor = mockExecutor(); + ManagedChannel initial1 = Mockito.mock(ManagedChannel.class); + ManagedChannel initial2 = Mockito.mock(ManagedChannel.class); + ManagedChannel rotated1 = Mockito.mock(ManagedChannel.class); + ManagedChannel resized = Mockito.mock(ManagedChannel.class); + ChannelFactory channelFactory = mockChannelFactory(); + Mockito.when(channelFactory.createSingleChannel()) + .thenReturn(initial1, initial2) + .thenReturn(rotated1) + .thenThrow(new IOException("Transient failure on second sub-channel")) + .thenReturn(resized); + + createMtlsPoolAndRotateCert( + ChannelPoolSettings.builder() + .setInitialChannelCount(2) + .setMinChannelCount(2) + .setMaxChannelCount(4) + .setMinRpcsPerChannel(1) + .setMaxRpcsPerChannel(2) + .build(), + channelFactory, + executor); - // Second refresh: both channels succeed pool.refresh(); + assertThat(pool.entries.get()).hasSize(1); + Mockito.verify(initial1).shutdown(); + Mockito.verify(initial2).shutdown(); + pool.invalidateDiskFingerprintCache(); assertThat(pool.shouldRefresh()).isFalse(); + // Dynamic pools rely on the periodic resize instead of a one-time refill. + Mockito.verify(executor, Mockito.never()).execute(Mockito.any(Runnable.class)); + + pool.resize(); + assertThat(pool.entries.get()).hasSize(2); + } + + @Test + void refill_afterShutdown_createsNoChannels() throws IOException { + ScheduledExecutorService executor = mockExecutor(); + ChannelFactory channelFactory = mockChannelFactory(); + Mockito.when(channelFactory.createSingleChannel()) + .thenReturn( + Mockito.mock(ManagedChannel.class), + Mockito.mock(ManagedChannel.class), + Mockito.mock(ManagedChannel.class)) + .thenThrow(new IOException("Transient failure on second sub-channel")) + .thenReturn(Mockito.mock(ManagedChannel.class)); + + createMtlsPoolAndRotateCert(ChannelPoolSettings.staticallySized(2), channelFactory, executor); + pool.refresh(); + ArgumentCaptor refillTask = ArgumentCaptor.forClass(Runnable.class); + Mockito.verify(executor).execute(refillTask.capture()); + + pool.shutdown(); + refillTask.getValue().run(); + + assertThat(pool.entries.get()).hasSize(1); + Mockito.verify(channelFactory, Mockito.times(4)).createSingleChannel(); + } + + @Test + void refill_channelCreationFailure_isLoggedAndDoesNotThrow() throws IOException { + ScheduledExecutorService executor = mockExecutor(); + ChannelFactory channelFactory = mockChannelFactory(); + Mockito.when(channelFactory.createSingleChannel()) + .thenReturn( + Mockito.mock(ManagedChannel.class), + Mockito.mock(ManagedChannel.class), + Mockito.mock(ManagedChannel.class)) + .thenThrow(new IOException("Transient failure on second sub-channel")) + .thenThrow(new IOException("Checked failure during refill")) + .thenThrow(new RuntimeException("Unchecked failure during refill")) + .thenReturn(Mockito.mock(ManagedChannel.class)); + + createMtlsPoolAndRotateCert(ChannelPoolSettings.staticallySized(2), channelFactory, executor); + pool.refresh(); + ArgumentCaptor refillTask = ArgumentCaptor.forClass(Runnable.class); + Mockito.verify(executor).execute(refillTask.capture()); + + // IOException is handled by expand(); the pool stays usable at its reduced size. + refillTask.getValue().run(); + assertThat(pool.entries.get()).hasSize(1); + + // RuntimeException is caught by refillSafely(). + refillTask.getValue().run(); + assertThat(pool.entries.get()).hasSize(1); + + refillTask.getValue().run(); + assertThat(pool.entries.get()).hasSize(2); + } + + @Test + void channelReactiveMTlsRefresh_refillRejectedByExecutor_stillCompletesRefresh() + throws IOException { + ScheduledExecutorService executor = mockExecutor(); + Mockito.doThrow(new RejectedExecutionException("Executor shut down")) + .when(executor) + .execute(Mockito.any(Runnable.class)); + ManagedChannel initial1 = Mockito.mock(ManagedChannel.class); + ManagedChannel initial2 = Mockito.mock(ManagedChannel.class); + ChannelFactory channelFactory = mockChannelFactory(); + Mockito.when(channelFactory.createSingleChannel()) + .thenReturn(initial1, initial2, Mockito.mock(ManagedChannel.class)) + .thenThrow(new IOException("Transient failure on second sub-channel")); + + createMtlsPoolAndRotateCert(ChannelPoolSettings.staticallySized(2), channelFactory, executor); + long genBefore = pool.getGeneration(); + + pool.refresh(); + + assertThat(pool.entries.get()).hasSize(1); + Mockito.verify(initial1).shutdown(); Mockito.verify(initial2).shutdown(); + assertThat(pool.getGeneration()).isEqualTo(genBefore + 1); + pool.invalidateDiskFingerprintCache(); + assertThat(pool.shouldRefresh()).isFalse(); } @Test From 44dfcffe84f09ed63ec813127c0305485176d535 Mon Sep 17 00:00:00 2001 From: Matt Castelaz Date: Wed, 30 Sep 2026 19:08:18 +0000 Subject: [PATCH 22/29] chore(gax-grpc,gax-httpjson): address review nits - Use Guava @VisibleForTesting in RefreshingHttpJsonChannel. - Make the delegating-wrapper ManagedHttpJsonChannel constructor package-private and document it; note refresh() is an intentional no-op. - Add Javadocs for getGeneration/refresh/shouldRefresh with {@inheritDoc} on overrides. - Drop the redundant useMtlsClientCertificate() check in InstantiatingGrpcChannelProvider and document the rotation-tracking conditions. --- .../InstantiatingGrpcChannelProvider.java | 8 ++++--- .../gax/httpjson/ManagedHttpJsonChannel.java | 23 +++++++++++++++++-- .../ManagedHttpJsonInterceptorChannel.java | 3 +++ .../httpjson/RefreshingHttpJsonChannel.java | 7 ++++-- 4 files changed, 34 insertions(+), 7 deletions(-) diff --git a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/InstantiatingGrpcChannelProvider.java b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/InstantiatingGrpcChannelProvider.java index 8a1c1fc8b0bb..67fc57b075d4 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/InstantiatingGrpcChannelProvider.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/InstantiatingGrpcChannelProvider.java @@ -401,10 +401,12 @@ public TransportChannel getTransportChannel() throws IOException { } private TransportChannel createChannel() throws IOException { + // Only track the workload certificate for rotation when the pool's channels actually present + // it, mirroring the credential selection in createSingleChannel(): DirectPath channels use + // GoogleDefaultChannelCredentials (ALTS) rather than the client certificate, and without an + // mtlsProvider there is no client certificate KeyStore (S2A or plain TLS is used instead). String workloadCertPath = - !this.canUseDirectPath() - && mtlsProvider != null - && certificateBasedAccess.useMtlsClientCertificate() + !this.canUseDirectPath() && mtlsProvider != null ? certificateBasedAccess.getWorkloadCertPath() : null; return GrpcTransportChannel.newBuilder() diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/ManagedHttpJsonChannel.java b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/ManagedHttpJsonChannel.java index 86863e4ce51c..cd3afcfa92b9 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/ManagedHttpJsonChannel.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/ManagedHttpJsonChannel.java @@ -60,7 +60,12 @@ protected ManagedHttpJsonChannel() { this(null, true, null, null, true); } - protected ManagedHttpJsonChannel(boolean isDelegatingWrapper) { + /** + * Constructor for subclasses that delegate all calls to a wrapped channel. The argument is + * unused; it only distinguishes this overload from {@link #ManagedHttpJsonChannel()}, which would + * otherwise allocate a transport and executor that the wrapper never uses or shuts down. + */ + ManagedHttpJsonChannel(boolean isDelegatingWrapper) { this.executor = null; this.usingDefaultExecutor = false; this.endpoint = null; @@ -69,6 +74,10 @@ protected ManagedHttpJsonChannel(boolean isDelegatingWrapper) { this.deadlineScheduledExecutorService = null; } + /** + * Returns a monotonic generation counter tracking the number of successful refreshes or channel + * rotations performed by this channel. Always {@code 0} for channels that do not refresh. + */ public long getGeneration() { return 0; } @@ -114,8 +123,18 @@ public HttpJsonClientCall newCall( deadlineScheduledExecutorService); } - public void refresh() {} + /** + * Refreshes or recreates the underlying transport of this channel if a certificate rotation has + * been detected. By default, this is a no-op. + */ + public void refresh() { + // No-op: this channel has no certificate to rotate. Overridden by RefreshingHttpJsonChannel. + } + /** + * Returns true if a certificate rotation has been detected on disk and this channel should be + * refreshed, or false otherwise. Always {@code false} for channels that do not refresh. + */ public boolean shouldRefresh() { return false; } diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/ManagedHttpJsonInterceptorChannel.java b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/ManagedHttpJsonInterceptorChannel.java index 9ccbbc4ec1a2..9adbfdd5c329 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/ManagedHttpJsonInterceptorChannel.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/ManagedHttpJsonInterceptorChannel.java @@ -48,6 +48,7 @@ class ManagedHttpJsonInterceptorChannel extends ManagedHttpJsonChannel { this.interceptor = interceptor; } + /** {@inheritDoc} */ @Override public long getGeneration() { return channel.getGeneration(); @@ -81,11 +82,13 @@ public HttpJsonClientCall newCall( return interceptor.interceptCall(methodDescriptor, callOptions, channel); } + /** {@inheritDoc} */ @Override public void refresh() { channel.refresh(); } + /** {@inheritDoc} */ @Override public boolean shouldRefresh() { return channel.shouldRefresh(); diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java index 8bc25d14d972..d29b6a673513 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java @@ -81,21 +81,23 @@ public RefreshingHttpJsonChannel( this.allEntries.add(initial); } - // Visible for testing + @VisibleForTesting String getWorkloadCertPath() { return workloadCertPath; } - // Visible for testing + @VisibleForTesting String getCertificateFingerprint(String certPath) { return WorkloadCertificateUtils.getCertificateFingerprint(certPath); } + /** {@inheritDoc} */ @Override public boolean shouldRefresh() { return rotationTracker.shouldRefresh(); } + /** {@inheritDoc} */ @Override public void refresh() { synchronized (refreshLock) { @@ -140,6 +142,7 @@ public void refresh() { } } + /** {@inheritDoc} */ @Override public long getGeneration() { return generation.get(); From 4ee955cd5736cbdd36fd5f135f6d8d241c1d4f31 Mon Sep 17 00:00:00 2001 From: Matt Castelaz Date: Wed, 30 Sep 2026 19:40:31 +0000 Subject: [PATCH 23/29] refactor(gax-httpjson): swap HttpTransport on certificate rotation instead of replacing the channel RefreshingHttpJsonChannel now wraps a single ManagedHttpJsonChannel and replaces its HttpTransport when the workload certificate rotates. Calls capture their transport when created, so in-flight requests complete on the previous transport, and NetHttpTransport holds no pooled resources that need releasing. This removes the per-generation channel tracking and reference counting (ChannelEntry, allEntries, ReleasingHttpJsonClientCall) and keeps a single set of executors for the lifetime of the channel. --- .../InstantiatingHttpJsonChannelProvider.java | 21 +- .../gax/httpjson/ManagedHttpJsonChannel.java | 10 +- .../httpjson/RefreshingHttpJsonChannel.java | 296 ++--------- .../RefreshingHttpJsonChannelTest.java | 482 ++++++------------ 4 files changed, 229 insertions(+), 580 deletions(-) diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProvider.java b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProvider.java index 48e4d9ca5579..da4fbc6cd4b8 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProvider.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProvider.java @@ -234,6 +234,10 @@ private ManagedHttpJsonChannel createSingleManagedChannel() throw new IOException("Failed to initialize mTLS HttpTransport"); } } + return buildManagedChannel(httpTransportToUse); + } + + private ManagedHttpJsonChannel buildManagedChannel(@Nullable HttpTransport httpTransportToUse) { return ManagedHttpJsonChannel.newBuilder() .setEndpoint(endpoint) .setExecutor(executor) @@ -251,19 +255,24 @@ private HttpJsonTransportChannel createChannel() throws IOException, GeneralSecu ManagedHttpJsonChannel baseChannel; if (workloadCertPath != null) { - java.util.function.Supplier channelFactory = + java.util.function.Supplier transportFactory = () -> { try { - return createSingleManagedChannel(); + HttpTransport mtlsTransport = createHttpTransport(); + if (mtlsTransport == null) { + throw new IOException("Failed to initialize mTLS HttpTransport"); + } + return mtlsTransport; } catch (IOException | GeneralSecurityException e) { - throw new java.lang.RuntimeException( - "Failed to create fresh ManagedHttpJsonChannel", e); + throw new java.lang.RuntimeException("Failed to create mTLS HttpTransport", e); } }; // RefreshingHttpJsonChannel records the baseline certificate fingerprint before creating the - // initial channel, so a rotation during startup is detected on the next auth failure. + // initial transport, so a rotation during startup is detected on the next auth failure. try { - baseChannel = new RefreshingHttpJsonChannel(channelFactory, workloadCertPath); + baseChannel = + new RefreshingHttpJsonChannel( + transportFactory, this::buildManagedChannel, workloadCertPath); } catch (RuntimeException e) { if (e.getCause() instanceof IOException) { throw (IOException) e.getCause(); diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/ManagedHttpJsonChannel.java b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/ManagedHttpJsonChannel.java index cd3afcfa92b9..bf79e17d7d89 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/ManagedHttpJsonChannel.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/ManagedHttpJsonChannel.java @@ -51,7 +51,7 @@ public class ManagedHttpJsonChannel implements HttpJsonChannel, BackgroundResour private final Executor executor; private final boolean usingDefaultExecutor; private final String endpoint; - private final HttpTransport httpTransport; + private volatile HttpTransport httpTransport; private final boolean usingDefaultTransport; private final ScheduledExecutorService deadlineScheduledExecutorService; private boolean isTransportShutdown; @@ -91,6 +91,14 @@ HttpTransport getHttpTransport() { return httpTransport; } + /** + * Replaces the transport used by calls created after this method returns. Calls that were already + * created keep using the transport they were created with. + */ + void setHttpTransport(HttpTransport httpTransport) { + this.httpTransport = httpTransport; + } + private ManagedHttpJsonChannel( @Nullable Executor executor, boolean usingDefaultExecutor, diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java index d29b6a673513..fa00ba34ed36 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java @@ -31,27 +31,22 @@ import com.google.api.client.http.HttpTransport; import com.google.api.core.InternalApi; -import com.google.api.gax.httpjson.ForwardingHttpJsonClientCall.SimpleForwardingHttpJsonClientCall; -import com.google.api.gax.httpjson.ForwardingHttpJsonClientCallListener.SimpleForwardingHttpJsonClientCallListener; import com.google.api.gax.rpc.mtls.CertificateRotationTracker; import com.google.api.gax.rpc.mtls.WorkloadCertificateUtils; import com.google.common.annotations.VisibleForTesting; -import java.util.concurrent.CancellationException; -import java.util.concurrent.ConcurrentLinkedQueue; +import java.util.concurrent.Executor; import java.util.concurrent.TimeUnit; -import java.util.concurrent.atomic.AtomicBoolean; -import java.util.concurrent.atomic.AtomicInteger; import java.util.concurrent.atomic.AtomicLong; -import java.util.concurrent.atomic.AtomicReference; +import java.util.function.Function; import java.util.function.Supplier; import java.util.logging.Level; import java.util.logging.Logger; -import org.jspecify.annotations.Nullable; /** * An implementation of {@link ManagedHttpJsonChannel} that supports dynamic mTLS certificate - * rotation by thread-safely hot-swapping the underlying active HTTP/JSON channel while gracefully - * retiring older connections after all active in-flight requests complete. + * rotation. When the workload certificate on disk changes, {@link #refresh()} replaces the {@link + * HttpTransport} of the underlying channel. Calls that were already created keep using the + * transport they were created with, so in-flight requests are not interrupted. */ @InternalApi public class RefreshingHttpJsonChannel extends ManagedHttpJsonChannel { @@ -59,26 +54,29 @@ public class RefreshingHttpJsonChannel extends ManagedHttpJsonChannel { private static final Logger LOG = Logger.getLogger(RefreshingHttpJsonChannel.class.getName()); private final CertificateRotationTracker rotationTracker; - private final Supplier channelFactory; + private final Supplier transportFactory; private final String workloadCertPath; - private final AtomicReference activeEntry; - // Keep track of all entries to properly await their termination - private final ConcurrentLinkedQueue allEntries = new ConcurrentLinkedQueue<>(); + private final ManagedHttpJsonChannel delegate; private final Object refreshLock = new Object(); private final AtomicLong generation = new AtomicLong(0); + /** + * @param transportFactory creates a transport configured with the certificate currently on disk + * @param channelFactory creates the underlying channel from the initial transport + * @param workloadCertPath path of the workload certificate to monitor for rotation + */ public RefreshingHttpJsonChannel( - Supplier channelFactory, String workloadCertPath) { + Supplier transportFactory, + Function channelFactory, + String workloadCertPath) { super(true); - this.channelFactory = channelFactory; + this.transportFactory = transportFactory; this.workloadCertPath = workloadCertPath; - // Record the baseline fingerprint before the initial channel loads the certificate from disk, + // Record the baseline fingerprint before the initial transport loads the certificate from disk, // so a rotation between the two steps is detected as a fingerprint change. this.rotationTracker = new CertificateRotationTracker(this::getWorkloadCertPath, this::getCertificateFingerprint); - ChannelEntry initial = new ChannelEntry(channelFactory.get()); - this.activeEntry = new AtomicReference<>(initial); - this.allEntries.add(initial); + this.delegate = channelFactory.apply(transportFactory.get()); } @VisibleForTesting @@ -117,28 +115,22 @@ public void refresh() { return; } - LOG.info("mTLS certificate rotation detected. Triggering HTTP/JSON channel pool refresh."); + LOG.info("mTLS certificate rotation detected. Refreshing HTTP/JSON transport."); - ChannelEntry newEntry; + HttpTransport newTransport; try { - newEntry = new ChannelEntry(channelFactory.get()); + newTransport = transportFactory.get(); } catch (Exception e) { - LOG.log(Level.WARNING, "Failed to refresh HTTP/JSON channel, leaving old channel", e); + LOG.log(Level.WARNING, "Failed to refresh HTTP/JSON transport, keeping old transport", e); return; } - allEntries.add(newEntry); - // Prune terminated entries after adding newEntry to ensure allEntries is never empty - allEntries.removeIf(entry -> entry != newEntry && entry.channel.isTerminated()); - - ChannelEntry oldEntry = activeEntry.getAndSet(newEntry); - // Order matters: swap activeEntry, then bump generation, then mark the tracker refreshed, so - // any failing RPC that observes the new fingerprint also observes the new generation. + // The previous transport is not shut down: calls created before the swap may still be using + // it, and NetHttpTransport holds no pooled resources that need releasing. + delegate.setHttpTransport(newTransport); + // Order matters: swap the transport, then bump generation, then mark the tracker refreshed, + // so any failing RPC that observes the new fingerprint also observes the new generation. generation.incrementAndGet(); rotationTracker.markRefreshed(currentDiskFingerprint); - - if (oldEntry != null) { - oldEntry.requestShutdown(); - } } } @@ -148,252 +140,60 @@ public long getGeneration() { return generation.get(); } - private ChannelEntry getRetainedEntry() { - while (true) { - ChannelEntry entry = activeEntry.get(); - if (entry.retain()) { - return entry; - } - if (entry == activeEntry.get()) { - throw new IllegalStateException("Channel has been shut down"); - } - } - } - @Override public HttpJsonClientCall newCall( ApiMethodDescriptor methodDescriptor, HttpJsonCallOptions callOptions) { - ChannelEntry entry = getRetainedEntry(); - try { - HttpJsonClientCall delegateCall = - entry.channel.newCall(methodDescriptor, callOptions); - return new ReleasingHttpJsonClientCall<>(delegateCall, entry); - } catch (Throwable t) { - entry.release(); - throw t; - } + return delegate.newCall(methodDescriptor, callOptions); } @Override - java.util.concurrent.Executor getExecutor() { - return activeEntry.get().channel.getExecutor(); + Executor getExecutor() { + return delegate.getExecutor(); } + @Override + String getEndpoint() { + return delegate.getEndpoint(); + } + + @Override @VisibleForTesting - ManagedHttpJsonChannel getActiveChannel() { - return activeEntry.get().channel; + HttpTransport getHttpTransport() { + return delegate.getHttpTransport(); } - private volatile boolean isShuttingDown = false; + @VisibleForTesting + void invalidateDiskFingerprintCache() { + rotationTracker.invalidateCache(); + } @Override public void shutdown() { - synchronized (refreshLock) { - isShuttingDown = true; - for (ChannelEntry entry : allEntries) { - entry.requestShutdown(); - } - } + delegate.shutdown(); } @Override public boolean isShutdown() { - return isShuttingDown; + return delegate.isShutdown(); } @Override public boolean isTerminated() { - if (!isShuttingDown) { - return false; - } - for (ChannelEntry entry : allEntries) { - if (!entry.channel.isTerminated()) { - return false; - } - } - return true; + return delegate.isTerminated(); } @Override public void shutdownNow() { - synchronized (refreshLock) { - isShuttingDown = true; - for (ChannelEntry entry : allEntries) { - entry.shutdownRequested.set(true); - entry.shutdownInitiated.set(true); - entry.channel.shutdownNow(); - } - } - } - - @VisibleForTesting - void invalidateDiskFingerprintCache() { - rotationTracker.invalidateCache(); + delegate.shutdownNow(); } @Override public boolean awaitTermination(long duration, TimeUnit unit) throws InterruptedException { - long endNanos = System.nanoTime() + unit.toNanos(duration); - for (ChannelEntry entry : allEntries) { - if (entry.channel.isTerminated()) { - continue; - } - long remainingNanos = endNanos - System.nanoTime(); - if (remainingNanos <= 0) { - return false; - } - if (!entry.channel.awaitTermination(remainingNanos, TimeUnit.NANOSECONDS) - && !entry.channel.isTerminated()) { - return false; - } - } - return true; + return delegate.awaitTermination(duration, unit); } @Override public void close() { - shutdown(); - } - - @Override - String getEndpoint() { - return activeEntry.get().channel.getEndpoint(); - } - - @Override - @VisibleForTesting - HttpTransport getHttpTransport() { - return activeEntry.get().channel.getHttpTransport(); - } - - /** Internal container to manage request reference-counting and graceful shutdown. */ - private static class ChannelEntry { - private final ManagedHttpJsonChannel channel; - private final AtomicInteger outstandingCalls = new AtomicInteger(0); - private final AtomicBoolean shutdownRequested = new AtomicBoolean(false); - private final AtomicBoolean shutdownInitiated = new AtomicBoolean(false); - - ChannelEntry(ManagedHttpJsonChannel channel) { - this.channel = channel; - } - - boolean retain() { - outstandingCalls.incrementAndGet(); - if (shutdownRequested.get()) { - release(); - return false; - } - return true; - } - - void release() { - int count = outstandingCalls.decrementAndGet(); - if (count < 0) { - LOG.warning("Channel entry reference count dropped below 0"); - } - // Must check outstandingCalls after shutdownRequested (in reverse order of retain()) to - // ensure mutual exclusion. - if (shutdownRequested.get() && outstandingCalls.get() == 0) { - shutdown(); - } - } - - void requestShutdown() { - shutdownRequested.set(true); - if (outstandingCalls.get() == 0) { - shutdown(); - } - } - - private void shutdown() { - if (shutdownInitiated.compareAndSet(false, true)) { - try { - channel.shutdown(); - } catch (Exception e) { - LOG.log(Level.WARNING, "Error shutting down retired HTTP/JSON channel", e); - } - } - } - } - - /** A client call decorator that decrements the entry counter upon call completion. */ - private static class ReleasingHttpJsonClientCall - extends SimpleForwardingHttpJsonClientCall { - - private final Object callLock = new Object(); - private volatile @Nullable CancellationException cancellationException; - private final ChannelEntry entry; - private final AtomicBoolean wasClosed = new AtomicBoolean(false); - private final AtomicBoolean wasReleased = new AtomicBoolean(false); - private final AtomicBoolean wasStarted = new AtomicBoolean(false); - - ReleasingHttpJsonClientCall(HttpJsonClientCall delegate, ChannelEntry entry) { - super(delegate); - this.entry = entry; - } - - @Override - public void start(Listener responseListener, HttpJsonMetadata requestHeaders) { - synchronized (callLock) { - if (!wasStarted.compareAndSet(false, true)) { - throw new IllegalStateException("Call is already started"); - } - if (cancellationException != null) { - if (wasReleased.compareAndSet(false, true)) { - entry.release(); - } - throw new IllegalStateException("Call is already cancelled", cancellationException); - } - try { - super.start( - new SimpleForwardingHttpJsonClientCallListener(responseListener) { - @Override - public void onClose(int statusCode, HttpJsonMetadata trailers) { - if (!wasClosed.compareAndSet(false, true)) { - return; - } - try { - super.onClose(statusCode, trailers); - } finally { - if (wasReleased.compareAndSet(false, true)) { - entry.release(); - } - } - } - }, - requestHeaders); - } catch (Throwable t) { - if (wasReleased.compareAndSet(false, true)) { - entry.release(); - } - throw t; - } - } - } - - @Override - public void cancel(@Nullable String message, @Nullable Throwable cause) { - boolean releaseImmediately = false; - try { - synchronized (callLock) { - this.cancellationException = new CancellationException(message); - if (!wasStarted.get()) { - releaseImmediately = true; - } - if (delegate() != null) { - super.cancel(message, cause); - } - } - } catch (Throwable t) { - if (!wasStarted.get()) { - releaseImmediately = true; - } - throw t; - } finally { - if (releaseImmediately && wasReleased.compareAndSet(false, true)) { - entry.release(); - } - } - } + delegate.close(); } } diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java index 54fcf80eea40..bad0bd05ce32 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java @@ -31,14 +31,25 @@ import static org.junit.jupiter.api.Assertions.assertEquals; import static org.junit.jupiter.api.Assertions.assertFalse; -import static org.junit.jupiter.api.Assertions.assertNotNull; -import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertNotSame; +import static org.junit.jupiter.api.Assertions.assertSame; import static org.junit.jupiter.api.Assertions.assertTrue; +import com.google.api.client.http.HttpTransport; +import com.google.api.client.testing.http.MockHttpTransport; +import com.google.api.gax.httpjson.testing.MockHttpService; +import com.google.protobuf.Field; +import java.util.ArrayDeque; import java.util.ArrayList; +import java.util.Arrays; +import java.util.Collections; +import java.util.HashMap; import java.util.List; +import java.util.Map; +import java.util.Queue; import java.util.concurrent.TimeUnit; import java.util.concurrent.atomic.AtomicInteger; +import java.util.function.Function; import java.util.function.Supplier; import javax.annotation.Nullable; import org.junit.jupiter.api.AfterEach; @@ -46,14 +57,36 @@ import org.junit.jupiter.api.Test; class RefreshingHttpJsonChannelTest { + private static final ApiMethodDescriptor FAKE_METHOD_DESCRIPTOR = + ApiMethodDescriptor.newBuilder() + .setFullMethodName("google.cloud.v1.Fake/FakeMethod") + .setHttpMethod("POST") + .setRequestFormatter( + ProtoMessageRequestFormatter.newBuilder() + .setPath( + "/fake/v1/name/{name}", + request -> { + Map fields = new HashMap<>(); + ProtoRestSerializer serializer = ProtoRestSerializer.create(); + serializer.putPathParam(fields, "name", request.getName()); + return fields; + }) + .setQueryParamsExtractor(request -> new HashMap<>()) + .setRequestBodyExtractor( + request -> + ProtoRestSerializer.create() + .toBody("*", request.toBuilder().clearName().build(), false)) + .build()) + .setResponseParser( + ProtoMessageResponseParser.newBuilder() + .setDefaultInstance(Field.getDefaultInstance()) + .build()) + .build(); + private static class FakeHttpJsonClientCall extends HttpJsonClientCall { - protected Listener listener; - @Override - public void start(Listener responseListener, HttpJsonMetadata requestHeaders) { - this.listener = responseListener; - } + public void start(Listener responseListener, HttpJsonMetadata requestHeaders) {} @Override public void request(int numMessages) {} @@ -116,26 +149,34 @@ public HttpJsonClientCall newCall( } } - private AtomicInteger channelFactoryCount; + private AtomicInteger transportFactoryCount; + private HttpTransport lastCreatedTransport; private FakeManagedHttpJsonChannel lastCreatedChannel; - private String testCertPath = "/fake/path"; - private String testFingerprint = "fingerprint1"; - private boolean shouldThrowOnFactory = false; + private String testCertPath; + private String testFingerprint; + private boolean shouldThrowOnFactory; private List createdChannels; - private Supplier channelFactory = + private Supplier transportFactory = () -> { if (shouldThrowOnFactory) { throw new RuntimeException("Simulated factory failure"); } - channelFactoryCount.incrementAndGet(); + transportFactoryCount.incrementAndGet(); + lastCreatedTransport = new MockHttpTransport(); + return lastCreatedTransport; + }; + + private Function channelFactory = + transport -> { lastCreatedChannel = new FakeManagedHttpJsonChannel(); + lastCreatedChannel.setHttpTransport(transport); return lastCreatedChannel; }; @BeforeEach void setUp() { - channelFactoryCount = new AtomicInteger(0); + transportFactoryCount = new AtomicInteger(0); testCertPath = "/fake/path"; testFingerprint = "fingerprint1"; shouldThrowOnFactory = false; @@ -151,7 +192,10 @@ void tearDown() { private RefreshingHttpJsonChannel createTestChannel() { RefreshingHttpJsonChannel ch = - new RefreshingHttpJsonChannel(channelFactory, "fake/cert/path.json") { + new RefreshingHttpJsonChannel( + () -> transportFactory.get(), + transport -> channelFactory.apply(transport), + "fake/cert/path.json") { @Override String getWorkloadCertPath() { return testCertPath; @@ -166,6 +210,11 @@ String getCertificateFingerprint(String certPath) { return ch; } + private void rotateCertificate(RefreshingHttpJsonChannel channel) { + channel.invalidateDiskFingerprintCache(); + testFingerprint = "fingerprint2"; + } + @Test void testShouldRefreshNullCertPath() { testCertPath = null; @@ -174,7 +223,7 @@ void testShouldRefreshNullCertPath() { } @Test - void testShouldRefreshFalseWhenUnchanged() throws InterruptedException { + void testShouldRefreshFalseWhenUnchanged() { RefreshingHttpJsonChannel channel = createTestChannel(); channel.invalidateDiskFingerprintCache(); // Invalidate 1-second cache @@ -182,13 +231,10 @@ void testShouldRefreshFalseWhenUnchanged() throws InterruptedException { } @Test - void testShouldRefreshTrueWhenChanged() throws InterruptedException { + void testShouldRefreshTrueWhenChanged() { RefreshingHttpJsonChannel channel = createTestChannel(); - channel.invalidateDiskFingerprintCache(); // Invalidate 1-second cache - - // Simulate disk fingerprint changing - testFingerprint = "fingerprint2"; + rotateCertificate(channel); assertTrue(channel.shouldRefresh()); } @@ -212,151 +258,129 @@ void shouldRefresh_doesNotCacheNegativeResultAndDetectsSubsequentRotationImmedia } @Test - void rotationDuringInitialChannelCreation_isDetectedAndRefreshed() { - channelFactory = + void rotationDuringInitialTransportCreation_isDetectedAndRefreshed() { + transportFactory = () -> { - if (channelFactoryCount.incrementAndGet() == 1) { - // Simulate the certificate rotating on disk while the initial channel loads it. + if (transportFactoryCount.incrementAndGet() == 1) { + // Simulate the certificate rotating on disk while the initial transport loads it. testFingerprint = "fingerprint2"; } - lastCreatedChannel = new FakeManagedHttpJsonChannel(); - return lastCreatedChannel; + lastCreatedTransport = new MockHttpTransport(); + return lastCreatedTransport; }; RefreshingHttpJsonChannel channel = createTestChannel(); - // The baseline was recorded before the initial channel was created, so the rotation is seen. + // The baseline was recorded before the initial transport was created, so the rotation is seen. assertTrue(channel.shouldRefresh()); channel.refresh(); - assertEquals(2, channelFactoryCount.get()); + assertEquals(2, transportFactoryCount.get()); assertEquals(1, channel.getGeneration()); assertFalse(channel.shouldRefresh()); } @Test - void testRefreshSwapsChannel() throws InterruptedException { - RefreshingHttpJsonChannel channel = createTestChannel(); - FakeManagedHttpJsonChannel firstChannel = lastCreatedChannel; - assertEquals(1, channelFactoryCount.get()); - - channel.invalidateDiskFingerprintCache(); // Invalidate 1-second cache - - // Change fingerprint - testFingerprint = "fingerprint2"; - - // Act - channel.refresh(); - - // Verify a new channel was created and the old one retired - assertEquals(2, channelFactoryCount.get()); - FakeManagedHttpJsonChannel secondChannel = lastCreatedChannel; - - // The old channel should receive a shutdown request immediately since there are no active calls - assertTrue(firstChannel.isShutdown()); - assertFalse(secondChannel.isShutdown()); - } - - @Test - void testRefreshKeepsInFlightChannelsAlive() throws InterruptedException { + void refresh_swapsTransportAndKeepsChannel() { RefreshingHttpJsonChannel channel = createTestChannel(); - FakeManagedHttpJsonChannel firstChannel = lastCreatedChannel; - - // Simulate an in-flight API call - FakeHttpJsonClientCall fakeCall = new FakeHttpJsonClientCall<>(); - firstChannel.nextCall = fakeCall; - - HttpJsonClientCall activeCall = channel.newCall(null, null); - - channel.invalidateDiskFingerprintCache(); // Invalidate 1-second cache - - // Change fingerprint & refresh - testFingerprint = "fingerprint2"; + FakeManagedHttpJsonChannel underlyingChannel = lastCreatedChannel; + HttpTransport initialTransport = channel.getHttpTransport(); + assertSame(lastCreatedTransport, initialTransport); + rotateCertificate(channel); channel.refresh(); - // Verify a new channel was created - assertEquals(2, channelFactoryCount.get()); - - // IMPORTANT: The first channel should NOT be shut down yet because of the active call! - assertFalse(firstChannel.isShutdown()); - - // Now start the call - activeCall.start(new HttpJsonClientCall.Listener() {}, null); - - assertNotNull(fakeCall.listener); - - // Fire onClose - fakeCall.listener.onClose(0, null); - - // FIRST CHANNEL SHOULD BE SHUT DOWN NOW! - assertTrue(firstChannel.isShutdown()); + assertEquals(2, transportFactoryCount.get()); + assertNotSame(initialTransport, channel.getHttpTransport()); + assertSame(lastCreatedTransport, channel.getHttpTransport()); + // The underlying channel (and its executors) is reused rather than replaced or shut down. + assertSame(underlyingChannel, lastCreatedChannel); + assertFalse(underlyingChannel.isShutdown()); + assertEquals(1, channel.getGeneration()); + assertFalse(channel.shouldRefresh()); } @Test - void testCancelBeforeStartReleasesChannelEntry() { + void callCreatedBeforeRefresh_usesOriginalTransport() throws Exception { + MockHttpService originalService = + new MockHttpService(Collections.singletonList(FAKE_METHOD_DESCRIPTOR), "google.com:443"); + MockHttpService rotatedService = + new MockHttpService(Collections.singletonList(FAKE_METHOD_DESCRIPTOR), "google.com:443"); + Field message = Field.newBuilder().setName("bob").setNumber(1).build(); + originalService.addResponse(message); + rotatedService.addResponse(message); + Queue transports = + new ArrayDeque<>(Arrays.asList(originalService, rotatedService)); + transportFactory = transports::remove; + channelFactory = + transport -> + ManagedHttpJsonChannel.newBuilder() + .setEndpoint("google.com:443") + .setHttpTransport(transport) + .build(); RefreshingHttpJsonChannel channel = createTestChannel(); - FakeManagedHttpJsonChannel firstChannel = lastCreatedChannel; + HttpJsonCallOptions callOptions = HttpJsonCallOptions.newBuilder().build(); + HttpJsonCallContext callContext = HttpJsonCallContext.createDefault(); - HttpJsonClientCall activeCall = channel.newCall(null, null); + HttpJsonClientCall callBeforeRefresh = + channel.newCall(FAKE_METHOD_DESCRIPTOR, callOptions); - channel.invalidateDiskFingerprintCache(); - testFingerprint = "fingerprint2"; + rotateCertificate(channel); channel.refresh(); + assertEquals(1, channel.getGeneration()); - // Because activeCall was created, the old channel should NOT be shut down yet - assertFalse(firstChannel.isShutdown()); - - // Cancel before start() is called - activeCall.cancel("Cancelled early", null); - - // Because cancel() safely released the entry, the old channel should now be shut down! - assertTrue(firstChannel.isShutdown()); + assertEquals( + message, + HttpJsonClientCalls.futureUnaryCall(callBeforeRefresh, message, callContext) + .get(10, TimeUnit.SECONDS)); + assertEquals(1, originalService.getRequestPaths().size()); + assertEquals(0, rotatedService.getRequestPaths().size()); + + HttpJsonClientCall callAfterRefresh = + channel.newCall(FAKE_METHOD_DESCRIPTOR, callOptions); + assertEquals( + message, + HttpJsonClientCalls.futureUnaryCall(callAfterRefresh, message, callContext) + .get(10, TimeUnit.SECONDS)); + assertEquals(1, originalService.getRequestPaths().size()); + assertEquals(1, rotatedService.getRequestPaths().size()); } @Test - void testRefreshDoesNotSpawnChannelWhenShutdown() throws InterruptedException { + void testRefreshDoesNotCreateTransportWhenShutdown() { RefreshingHttpJsonChannel channel = createTestChannel(); - FakeManagedHttpJsonChannel firstChannel = lastCreatedChannel; - assertEquals(1, channelFactoryCount.get()); + assertEquals(1, transportFactoryCount.get()); - // Simulate that the channel pool is shut down. channel.shutdown(); - firstChannel.shutdown(); - - channel.invalidateDiskFingerprintCache(); // Invalidate 1-second cache - - // Change fingerprint - testFingerprint = "fingerprint2"; - - // Act + rotateCertificate(channel); channel.refresh(); - // Verify no new channel was spawned - assertEquals(1, channelFactoryCount.get()); + assertEquals(1, transportFactoryCount.get()); + assertEquals(0, channel.getGeneration()); } @Test - void testRefreshFactoryExceptionDoesNotWedgeFingerprint() throws InterruptedException { + void testRefreshFactoryExceptionDoesNotWedgeFingerprint() { RefreshingHttpJsonChannel channel = createTestChannel(); - assertEquals(1, channelFactoryCount.get()); + HttpTransport initialTransport = channel.getHttpTransport(); + assertEquals(1, transportFactoryCount.get()); shouldThrowOnFactory = true; - channel.invalidateDiskFingerprintCache(); // Invalidate 1-second cache - testFingerprint = "fingerprint2"; + rotateCertificate(channel); - // Factory failure is logged and the existing channel is kept + // Factory failure is logged and the existing transport is kept channel.refresh(); - assertEquals(1, channelFactoryCount.get()); + assertEquals(1, transportFactoryCount.get()); assertEquals(0, channel.getGeneration()); + assertSame(initialTransport, channel.getHttpTransport()); - // Because factory threw, activeCertFingerprint should NOT be updated to fingerprint2 - // Therefore shouldRefresh() should still return true + // Because the factory threw, the new fingerprint is not recorded as active, so the channel + // still reports that it should be refreshed. assertTrue(channel.shouldRefresh()); shouldThrowOnFactory = false; channel.refresh(); - assertEquals(2, channelFactoryCount.get()); + assertEquals(2, transportFactoryCount.get()); assertFalse(channel.shouldRefresh()); } @@ -373,8 +397,7 @@ void testShutdownNowSetsIsShutdown() { @Test void testAwaitTerminationZeroTimeoutOnTerminatedChannelReturnsTrue() throws InterruptedException { RefreshingHttpJsonChannel channel = createTestChannel(); - FakeManagedHttpJsonChannel firstChannel = lastCreatedChannel; - firstChannel.isTerminated = true; + lastCreatedChannel.isTerminated = true; channel.shutdown(); assertTrue(channel.awaitTermination(0, TimeUnit.MILLISECONDS)); @@ -383,22 +406,24 @@ void testAwaitTerminationZeroTimeoutOnTerminatedChannelReturnsTrue() throws Inte @Test void testChannelDelegationMethods() { RefreshingHttpJsonChannel channel = createTestChannel(); - FakeManagedHttpJsonChannel firstChannel = lastCreatedChannel; + FakeManagedHttpJsonChannel underlyingChannel = lastCreatedChannel; + FakeHttpJsonClientCall fakeCall = new FakeHttpJsonClientCall<>(); + underlyingChannel.nextCall = fakeCall; - assertEquals(firstChannel.getEndpoint(), channel.getEndpoint()); - assertEquals(firstChannel.getHttpTransport(), channel.getHttpTransport()); - assertEquals(firstChannel.getExecutor(), channel.getExecutor()); + assertEquals(underlyingChannel.getEndpoint(), channel.getEndpoint()); + assertEquals(underlyingChannel.getHttpTransport(), channel.getHttpTransport()); + assertEquals(underlyingChannel.getExecutor(), channel.getExecutor()); + assertSame(fakeCall, channel.newCall(null, null)); } @Test - void testNewCallAfterShutdownNowThrowsIllegalStateException() { + void close_shutsDownUnderlyingChannel() { RefreshingHttpJsonChannel channel = createTestChannel(); - channel.shutdownNow(); - assertThrows( - IllegalStateException.class, - () -> channel.newCall(null, null), - "Channel has been shut down"); + channel.close(); + + assertTrue(lastCreatedChannel.isShutdown()); + assertTrue(channel.isShutdown()); } @Test @@ -409,8 +434,7 @@ void testConcurrentNewCallDuringRefresh() throws InterruptedException { java.util.concurrent.Executors.newFixedThreadPool(threadCount); java.util.concurrent.CountDownLatch latch = new java.util.concurrent.CountDownLatch(threadCount); - java.util.concurrent.atomic.AtomicInteger successCount = - new java.util.concurrent.atomic.AtomicInteger(0); + AtomicInteger successCount = new AtomicInteger(0); for (int i = 0; i < threadCount; i++) { executorService.submit( @@ -424,8 +448,7 @@ void testConcurrentNewCallDuringRefresh() throws InterruptedException { }); } - channel.invalidateDiskFingerprintCache(); - testFingerprint = "fingerprint2"; + rotateCertificate(channel); channel.refresh(); latch.await(5, TimeUnit.SECONDS); @@ -439,8 +462,7 @@ void testGenerationIncrementAndLifecycleOnDelegatingWrapper() throws Exception { RefreshingHttpJsonChannel channel = createTestChannel(); assertEquals(0, channel.getGeneration()); - channel.invalidateDiskFingerprintCache(); - testFingerprint = "fingerprint2"; + rotateCertificate(channel); channel.refresh(); assertEquals(1, channel.getGeneration()); @@ -454,194 +476,4 @@ void testGenerationIncrementAndLifecycleOnDelegatingWrapper() throws Exception { assertTrue(channel.isTerminated()); assertTrue(channel.awaitTermination(1, TimeUnit.SECONDS)); } - - @Test - void testNewCall_whenDelegateThrowsError_releasesEntryAndShutsDownRetiredChannel() { - RefreshingHttpJsonChannel channel = createTestChannel(); - FakeManagedHttpJsonChannel firstChannel = lastCreatedChannel; - firstChannel.nextCall = null; - // Configure firstChannel to throw an Error on newCall - FakeManagedHttpJsonChannel throwingChannel = - new FakeManagedHttpJsonChannel() { - @Override - public HttpJsonClientCall newCall( - ApiMethodDescriptor methodDescriptor, - HttpJsonCallOptions callOptions) { - throw new LinkageError("Simulated native error in newCall"); - } - }; - channelFactory = - () -> { - channelFactoryCount.incrementAndGet(); - lastCreatedChannel = throwingChannel; - return throwingChannel; - }; - RefreshingHttpJsonChannel testChannel = createTestChannel(); - - assertThrows(LinkageError.class, () -> testChannel.newCall(null, null)); - - // Refresh should immediately shut down throwingChannel since ref count returned to 0 - testFingerprint = "fingerprint2"; - testChannel.refresh(); - assertTrue(throwingChannel.isShutdown()); - } - - @Test - void testStart_whenDelegateThrowsError_releasesEntryAndShutsDownRetiredChannel() { - RefreshingHttpJsonChannel channel = createTestChannel(); - FakeManagedHttpJsonChannel firstChannel = lastCreatedChannel; - firstChannel.nextCall = - new FakeHttpJsonClientCall() { - @Override - public void start(Listener responseListener, HttpJsonMetadata requestHeaders) { - throw new AssertionError("Simulated Error in start"); - } - }; - - HttpJsonClientCall call = channel.newCall(null, null); - testFingerprint = "fingerprint2"; - channel.refresh(); - assertFalse(firstChannel.isShutdown()); - - assertThrows( - AssertionError.class, () -> call.start(new HttpJsonClientCall.Listener() {}, null)); - assertTrue(firstChannel.isShutdown()); - } - - @Test - void testCancel_whenDelegateThrowsException_releasesEntryAndShutsDownRetiredChannel() { - RefreshingHttpJsonChannel channel = createTestChannel(); - FakeManagedHttpJsonChannel firstChannel = lastCreatedChannel; - firstChannel.nextCall = - new FakeHttpJsonClientCall() { - @Override - public void cancel(String message, Throwable cause) { - throw new RuntimeException("Simulated cancel failure"); - } - }; - - HttpJsonClientCall call = channel.newCall(null, null); - testFingerprint = "fingerprint2"; - channel.refresh(); - assertFalse(firstChannel.isShutdown()); - - assertThrows(RuntimeException.class, () -> call.cancel("cancel", null)); - assertTrue(firstChannel.isShutdown()); - } - - @Test - void testConcurrentStartAndCancel_neverLeaksOrDoubleReleasesEntry() throws Exception { - RefreshingHttpJsonChannel channel = createTestChannel(); - FakeManagedHttpJsonChannel firstChannel = lastCreatedChannel; - - int iterations = 100; - java.util.concurrent.ExecutorService executor = - java.util.concurrent.Executors.newFixedThreadPool(2); - try { - for (int i = 0; i < iterations; i++) { - firstChannel.nextCall = - new FakeHttpJsonClientCall() { - private boolean closed = false; - - @Override - public synchronized void start( - Listener responseListener, HttpJsonMetadata requestHeaders) { - if (closed) { - // Models HttpJsonClientCallImpl returning early when closed - return; - } - super.start(responseListener, requestHeaders); - } - - @Override - public synchronized void cancel(String message, Throwable cause) { - closed = true; - if (listener != null) { - listener.onClose(499, null); - } - } - }; - - HttpJsonClientCall call = channel.newCall(null, null); - java.util.concurrent.CyclicBarrier barrier = new java.util.concurrent.CyclicBarrier(2); - java.util.concurrent.Future f1 = - executor.submit( - () -> { - try { - barrier.await(); - call.start(new HttpJsonClientCall.Listener() {}, null); - } catch (Exception ignored) { - } - }); - java.util.concurrent.Future f2 = - executor.submit( - () -> { - try { - barrier.await(); - call.cancel("cancel", null); - } catch (Exception ignored) { - } - }); - f1.get(5, TimeUnit.SECONDS); - f2.get(5, TimeUnit.SECONDS); - } - } finally { - executor.shutdownNow(); - } - - testFingerprint = "fingerprint2"; - channel.refresh(); - assertTrue(firstChannel.isShutdown()); - } - - @Test - void cancel_whenStartedAndSuperCancelThrows_doesNotReleasePrematurelyUntilOnClose() { - RefreshingHttpJsonChannel channel = createTestChannel(); - FakeManagedHttpJsonChannel firstChannel = lastCreatedChannel; - FakeHttpJsonClientCall delegateCall = - new FakeHttpJsonClientCall() { - @Override - public void cancel(String message, Throwable cause) { - throw new RuntimeException("Simulated cancel failure"); - } - }; - firstChannel.nextCall = delegateCall; - - HttpJsonClientCall call = channel.newCall(null, null); - call.start(new HttpJsonClientCall.Listener() {}, null); - - assertThrows(RuntimeException.class, () -> call.cancel("abort", null)); - - // Rotate pool while call is still active (onClose hasn't fired yet): - // firstChannel must NOT be shut down yet because call is still active - testFingerprint = "fingerprint2"; - channel.refresh(); - assertFalse(firstChannel.isShutdown()); - - // Once onClose fires, entry is released and firstChannel shuts down - delegateCall.listener.onClose(200, null); - assertTrue(firstChannel.isShutdown()); - } - - @Test - void start_whenCalledTwice_throwsIllegalStateExceptionAndDoesNotReleaseFirstCallEntry() { - RefreshingHttpJsonChannel channel = createTestChannel(); - FakeManagedHttpJsonChannel firstChannel = lastCreatedChannel; - FakeHttpJsonClientCall delegateCall = new FakeHttpJsonClientCall<>(); - firstChannel.nextCall = delegateCall; - - HttpJsonClientCall call = channel.newCall(null, null); - call.start(new HttpJsonClientCall.Listener() {}, null); - - assertThrows( - IllegalStateException.class, - () -> call.start(new HttpJsonClientCall.Listener() {}, null)); - - testFingerprint = "fingerprint2"; - channel.refresh(); - assertFalse(firstChannel.isShutdown()); - - delegateCall.listener.onClose(200, null); - assertTrue(firstChannel.isShutdown()); - } } From 483e3960a29e6fd25b0516d805e22d934b857a10 Mon Sep 17 00:00:00 2001 From: Matt Castelaz Date: Thu, 1 Oct 2026 02:46:10 +0000 Subject: [PATCH 24/29] fix(gax): re-arm the rotation retry after a stream makes progress --- .../gax/retrying/StreamingRetryAlgorithm.java | 19 ++- .../gax/rpc/ApiResultRetryAlgorithmTest.java | 129 ++++++++++++++++++ 2 files changed, 145 insertions(+), 3 deletions(-) diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/retrying/StreamingRetryAlgorithm.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/retrying/StreamingRetryAlgorithm.java index e4d5461a6b4e..2ad0f6180e6c 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/retrying/StreamingRetryAlgorithm.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/retrying/StreamingRetryAlgorithm.java @@ -102,13 +102,26 @@ public StreamingRetryAlgorithm( (ServerStreamingAttemptException) previousThrowable; previousThrowable = previousThrowable.getCause(); - // If we have made progress in the last attempt, then reset the delays + // If we have made progress in the last attempt, then reset the delays and attempt counts. + // The next attempt is computed from a fresh baseline, so that result algorithms comparing + // the attempt count with the overall attempt count see the reset stream as a new sequence + // of attempts. The previous overall attempt count is then added back so that it keeps + // increasing across resets. if (attemptException.hasSeenResponses()) { - previousSettings = + int previousOverallAttemptCount = previousSettings.getOverallAttemptCount(); + TimedAttemptSettings resetSettings = createFirstAttempt(context).toBuilder() .setFirstAttemptStartTimeNanos(previousSettings.getFirstAttemptStartTimeNanos()) - .setOverallAttemptCount(previousSettings.getOverallAttemptCount()) .build(); + TimedAttemptSettings nextSettings = + super.createNextAttempt(context, previousThrowable, previousResponse, resetSettings); + if (nextSettings == null) { + return null; + } + return nextSettings.toBuilder() + .setOverallAttemptCount( + nextSettings.getOverallAttemptCount() + previousOverallAttemptCount) + .build(); } } diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ApiResultRetryAlgorithmTest.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ApiResultRetryAlgorithmTest.java index 373eafbb0640..1232be098e52 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ApiResultRetryAlgorithmTest.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ApiResultRetryAlgorithmTest.java @@ -32,6 +32,7 @@ import static org.junit.jupiter.api.Assertions.assertEquals; import static org.junit.jupiter.api.Assertions.assertFalse; import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertNull; import static org.junit.jupiter.api.Assertions.assertTrue; import static org.mockito.Mockito.mock; import static org.mockito.Mockito.when; @@ -40,6 +41,8 @@ import com.google.api.gax.retrying.ExponentialRetryAlgorithm; import com.google.api.gax.retrying.RetryAlgorithm; import com.google.api.gax.retrying.RetrySettings; +import com.google.api.gax.retrying.ServerStreamingAttemptException; +import com.google.api.gax.retrying.StreamingRetryAlgorithm; import com.google.api.gax.retrying.TimedAttemptSettings; import com.google.api.gax.rpc.StatusCode.Code; import com.google.api.gax.rpc.testing.FakeStatusCode; @@ -322,4 +325,130 @@ void testSecondRotationFailureDoesNotConsumeNormalRetryBudget() { assertNotNull(afterSecondFailure); assertFalse(retryAlgorithm.shouldRetry(context, rotationEx, null, afterSecondFailure)); } + + @Test + void testStreamRotationRetryIsAvailableAgainAfterProgress() { + ApiCallContext context = + mock(ApiCallContext.class, Mockito.withSettings().withoutAnnotations()); + when(context.getRetryableCodes()).thenReturn(Sets.newHashSet(Code.UNAVAILABLE)); + + RetrySettings settings = + RetrySettings.newBuilder() + .setMaxAttempts(5) + .setInitialRetryDelayDuration(Duration.ofMillis(100)) + .setRetryDelayMultiplier(2.0) + .setMaxRetryDelayDuration(Duration.ofSeconds(1)) + .setTotalTimeoutDuration(Duration.ofMinutes(10)) + .build(); + + StreamingRetryAlgorithm retryAlgorithm = + new StreamingRetryAlgorithm<>( + new ApiResultRetryAlgorithm(), + new ExponentialRetryAlgorithm(settings, NanoClock.getDefaultClock())); + ApiException unavailableEx = + new ApiException(null, new FakeStatusCode(Code.UNAVAILABLE), /* retryable= */ true); + UnauthenticatedException rotationEx = + new UnauthenticatedException( + "Expired cert", null, new FakeStatusCode(Code.UNAUTHENTICATED), /* retryable= */ true); + + // The first attempt fails with UNAVAILABLE before receiving any messages: normal retry. + TimedAttemptSettings attempt1 = + retryAlgorithm.createNextAttempt( + context, + new ServerStreamingAttemptException(unavailableEx, true, false), + null, + retryAlgorithm.createFirstAttempt(context)); + assertEquals(1, attempt1.getAttemptCount()); + assertEquals(1, attempt1.getOverallAttemptCount()); + + // The second attempt receives messages and then fails with a rotation error. The earlier + // retry must not prevent the free rotation retry. + ServerStreamingAttemptException rotationAfterProgress = + new ServerStreamingAttemptException(rotationEx, true, true); + TimedAttemptSettings attempt2 = + retryAlgorithm.createNextAttempt(context, rotationAfterProgress, null, attempt1); + assertNotNull(attempt2); + assertEquals(Duration.ZERO, attempt2.getRetryDelayDuration()); + assertEquals(Duration.ZERO, attempt2.getRandomizedRetryDelayDuration()); + assertEquals(0, attempt2.getAttemptCount()); + assertEquals(2, attempt2.getOverallAttemptCount()); + assertTrue(retryAlgorithm.shouldRetry(context, rotationAfterProgress, null, attempt2)); + + // A repeated rotation error without further progress stops the stream. + ServerStreamingAttemptException rotationWithoutProgress = + new ServerStreamingAttemptException(rotationEx, true, false); + TimedAttemptSettings attempt3 = + retryAlgorithm.createNextAttempt(context, rotationWithoutProgress, null, attempt2); + assertNotNull(attempt3); + assertEquals(3, attempt3.getOverallAttemptCount()); + assertFalse(retryAlgorithm.shouldRetry(context, rotationWithoutProgress, null, attempt3)); + } + + @Test + void testStreamProgressResetKeepsOverallAttemptCountIncreasing() { + ApiCallContext context = + mock(ApiCallContext.class, Mockito.withSettings().withoutAnnotations()); + when(context.getRetryableCodes()).thenReturn(Sets.newHashSet(Code.UNAVAILABLE)); + + RetrySettings settings = + RetrySettings.newBuilder() + .setMaxAttempts(5) + .setInitialRetryDelayDuration(Duration.ofMillis(100)) + .setRetryDelayMultiplier(2.0) + .setMaxRetryDelayDuration(Duration.ofSeconds(1)) + .setTotalTimeoutDuration(Duration.ofMinutes(10)) + .build(); + + StreamingRetryAlgorithm retryAlgorithm = + new StreamingRetryAlgorithm<>( + new ApiResultRetryAlgorithm(), + new ExponentialRetryAlgorithm(settings, NanoClock.getDefaultClock())); + ApiException unavailableEx = + new ApiException(null, new FakeStatusCode(Code.UNAVAILABLE), /* retryable= */ true); + ServerStreamingAttemptException unavailableWithoutProgress = + new ServerStreamingAttemptException(unavailableEx, true, false); + ServerStreamingAttemptException unavailableAfterProgress = + new ServerStreamingAttemptException(unavailableEx, true, true); + + TimedAttemptSettings attempt1 = + retryAlgorithm.createNextAttempt( + context, unavailableWithoutProgress, null, retryAlgorithm.createFirstAttempt(context)); + TimedAttemptSettings attempt2 = + retryAlgorithm.createNextAttempt(context, unavailableWithoutProgress, null, attempt1); + assertEquals(2, attempt2.getAttemptCount()); + assertEquals(2, attempt2.getOverallAttemptCount()); + + // Progress resets the attempt count, but the overall attempt count keeps increasing. + TimedAttemptSettings attempt3 = + retryAlgorithm.createNextAttempt(context, unavailableAfterProgress, null, attempt2); + assertEquals(1, attempt3.getAttemptCount()); + assertEquals(3, attempt3.getOverallAttemptCount()); + assertEquals(settings.getInitialRetryDelayDuration(), attempt3.getRetryDelayDuration()); + assertEquals( + attempt2.getFirstAttemptStartTimeNanos(), attempt3.getFirstAttemptStartTimeNanos()); + assertTrue(retryAlgorithm.shouldRetry(context, unavailableAfterProgress, null, attempt3)); + } + + @Test + void testStreamProgressResetReturnsNullWhenNotRetryable() { + ApiCallContext context = + mock(ApiCallContext.class, Mockito.withSettings().withoutAnnotations()); + when(context.getRetryableCodes()).thenReturn(Collections.emptySet()); + + StreamingRetryAlgorithm retryAlgorithm = + new StreamingRetryAlgorithm<>( + new ApiResultRetryAlgorithm(), + new ExponentialRetryAlgorithm( + RetrySettings.newBuilder().setMaxAttempts(5).build(), NanoClock.getDefaultClock())); + ApiException unavailableEx = + new ApiException(null, new FakeStatusCode(Code.UNAVAILABLE), /* retryable= */ true); + + TimedAttemptSettings next = + retryAlgorithm.createNextAttempt( + context, + new ServerStreamingAttemptException(unavailableEx, true, true), + null, + retryAlgorithm.createFirstAttempt(context)); + assertNull(next); + } } From fd46ba84f37b3babd52874bdfa0d84d6021bb0a1 Mon Sep 17 00:00:00 2001 From: Matt Castelaz Date: Thu, 1 Oct 2026 02:46:11 +0000 Subject: [PATCH 25/29] fix(gax-grpc): address audit findings --- .../com/google/api/gax/grpc/ChannelPool.java | 15 ++++---- .../google/api/gax/grpc/ChannelPoolTest.java | 35 +++++++++++++++++-- 2 files changed, 42 insertions(+), 8 deletions(-) diff --git a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java index e81fc5867f82..0e634fcd43c3 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java @@ -50,7 +50,6 @@ import java.util.List; import java.util.concurrent.CancellationException; import java.util.concurrent.Executors; -import java.util.concurrent.RejectedExecutionException; import java.util.concurrent.ScheduledExecutorService; import java.util.concurrent.ScheduledFuture; import java.util.concurrent.TimeUnit; @@ -441,7 +440,7 @@ private void expand(int desiredSize) { for (int i = 0; i < desiredSize - localEntries.size(); i++) { try { newEntries.add(new Entry(channelFactory.createSingleChannel())); - } catch (IOException e) { + } catch (Exception e) { LOG.log(Level.WARNING, "Failed to add channel", e); } } @@ -483,10 +482,14 @@ boolean shouldRefresh() { /** * Replace all of the channels in the channel pool with fresh ones. This is meant to mitigate the - * hourly GFE disconnects by giving clients the ability to prime the channel on reconnect. + * hourly GFE disconnects by giving clients the ability to prime the channel on reconnect, and to + * pick up a rotated mTLS workload certificate. * - *

This is done on a best effort basis. If the replacement channel fails to construct, the old - * channel will continue to be used. + *

This is done on a best effort basis. When no workload certificate is configured, a channel + * whose replacement fails to construct continues to be used. When a workload certificate is + * configured, a channel whose replacement fails to construct is dropped so that no traffic is + * routed to a channel using the old certificate; if no replacement can be constructed at all, the + * pool is left unchanged. */ @InternalApi("Visible for testing") void refresh() { @@ -617,7 +620,7 @@ boolean refreshAll(boolean dropUnrefreshedChannels) { private void scheduleRefill() { try { backgroundExecutorProvider.getExecutor().execute(this::refillSafely); - } catch (RejectedExecutionException e) { + } catch (RuntimeException e) { LOG.log(Level.WARNING, "Failed to schedule channel pool refill", e); } } diff --git a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java index 5bc6f391065c..0fe97f39c376 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java @@ -788,11 +788,11 @@ void refill_channelCreationFailure_isLoggedAndDoesNotThrow() throws IOException ArgumentCaptor refillTask = ArgumentCaptor.forClass(Runnable.class); Mockito.verify(executor).execute(refillTask.capture()); - // IOException is handled by expand(); the pool stays usable at its reduced size. + // Checked and unchecked failures are both handled by expand(); the pool stays usable at its + // reduced size. refillTask.getValue().run(); assertThat(pool.entries.get()).hasSize(1); - // RuntimeException is caught by refillSafely(). refillTask.getValue().run(); assertThat(pool.entries.get()).hasSize(1); @@ -800,6 +800,36 @@ void refill_channelCreationFailure_isLoggedAndDoesNotThrow() throws IOException assertThat(pool.entries.get()).hasSize(2); } + @Test + void refill_failureAfterPartialProgress_keepsCreatedChannels() throws IOException { + ScheduledExecutorService executor = mockExecutor(); + ManagedChannel refilled = Mockito.mock(ManagedChannel.class); + ChannelFactory channelFactory = mockChannelFactory(); + Mockito.when(channelFactory.createSingleChannel()) + .thenReturn( + Mockito.mock(ManagedChannel.class), + Mockito.mock(ManagedChannel.class), + Mockito.mock(ManagedChannel.class), + Mockito.mock(ManagedChannel.class)) + .thenThrow(new IOException("Transient failure on second sub-channel")) + .thenThrow(new IOException("Transient failure on third sub-channel")) + .thenReturn(refilled) + .thenThrow(new RuntimeException("Unchecked failure during refill")); + + createMtlsPoolAndRotateCert(ChannelPoolSettings.staticallySized(3), channelFactory, executor); + pool.refresh(); + assertThat(pool.entries.get()).hasSize(1); + ArgumentCaptor refillTask = ArgumentCaptor.forClass(Runnable.class); + Mockito.verify(executor).execute(refillTask.capture()); + + refillTask.getValue().run(); + + // The channel created before the failure is added to the pool rather than orphaned. + assertThat(pool.entries.get()).hasSize(2); + Mockito.verify(refilled, Mockito.never()).shutdown(); + Mockito.verify(channelFactory, Mockito.times(8)).createSingleChannel(); + } + @Test void channelReactiveMTlsRefresh_refillRejectedByExecutor_stillCompletesRefresh() throws IOException { @@ -819,6 +849,7 @@ void channelReactiveMTlsRefresh_refillRejectedByExecutor_stillCompletesRefresh() pool.refresh(); + Mockito.verify(executor).execute(Mockito.any(Runnable.class)); assertThat(pool.entries.get()).hasSize(1); Mockito.verify(initial1).shutdown(); Mockito.verify(initial2).shutdown(); From 0ed8e1c2c14330fece6185d6dcf5460d64aeae30 Mon Sep 17 00:00:00 2001 From: Matt Castelaz Date: Thu, 1 Oct 2026 02:46:11 +0000 Subject: [PATCH 26/29] fix(gax-httpjson): address audit findings --- .../httpjson/RefreshingHttpJsonChannel.java | 13 ++++-- ...tantiatingHttpJsonChannelProviderTest.java | 22 ++++++++++ .../RefreshingHttpJsonChannelTest.java | 42 +++++++++++++++++++ 3 files changed, 73 insertions(+), 4 deletions(-) diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java index fa00ba34ed36..f4201b4b30b9 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java @@ -125,7 +125,7 @@ public void refresh() { return; } // The previous transport is not shut down: calls created before the swap may still be using - // it, and NetHttpTransport holds no pooled resources that need releasing. + // it, and it needs no explicit shutdown because its idle keep-alive connections expire. delegate.setHttpTransport(newTransport); // Order matters: swap the transport, then bump generation, then mark the tracker refreshed, // so any failing RPC that observes the new fingerprint also observes the new generation. @@ -169,7 +169,10 @@ void invalidateDiskFingerprintCache() { @Override public void shutdown() { - delegate.shutdown(); + // Serialized with refresh() so that a transport swap cannot race with shutdown. + synchronized (refreshLock) { + delegate.shutdown(); + } } @Override @@ -184,7 +187,9 @@ public boolean isTerminated() { @Override public void shutdownNow() { - delegate.shutdownNow(); + synchronized (refreshLock) { + delegate.shutdownNow(); + } } @Override @@ -194,6 +199,6 @@ public boolean awaitTermination(long duration, TimeUnit unit) throws Interrupted @Override public void close() { - delegate.close(); + shutdown(); } } diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProviderTest.java b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProviderTest.java index 1d967f4b108b..b39af9581357 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProviderTest.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProviderTest.java @@ -316,6 +316,28 @@ void getTransportChannel_whenMtlsActiveAndKeyStoreNull_throwsIOException() { assertThat(thrown).hasMessageThat().contains("Failed to initialize mTLS HttpTransport"); } + @Test + void getTransportChannel_whenRotationEnabledAndKeyStoreNull_throwsIOException() { + Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); + Mockito.when(certificateBasedAccess.getWorkloadCertPath()).thenReturn("fake/cert/path.json"); + com.google.auth.mtls.MtlsProvider providerWithNullKeyStore = + new com.google.api.gax.rpc.testing.FakeMtlsProvider(null, "", false); + + InstantiatingHttpJsonChannelProvider provider = + InstantiatingHttpJsonChannelProvider.newBuilder() + .setEndpoint(DEFAULT_ENDPOINT) + .setMtlsProvider(providerWithNullKeyStore) + .setCertificateBasedAccess(certificateBasedAccess) + .build(); + InstantiatingHttpJsonChannelProvider finalProvider = + (InstantiatingHttpJsonChannelProvider) provider.withHeaders(DEFAULT_HEADER_MAP); + + IOException thrown = + org.junit.jupiter.api.Assertions.assertThrows( + IOException.class, finalProvider::getTransportChannel); + assertThat(thrown).hasMessageThat().contains("Failed to initialize mTLS HttpTransport"); + } + @Test void getTransportChannel_whenKeyStoreUninitialized_causeIsSecurityException() throws Exception { Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java index bad0bd05ce32..cc1dc72d5453 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java @@ -47,6 +47,7 @@ import java.util.List; import java.util.Map; import java.util.Queue; +import java.util.concurrent.CountDownLatch; import java.util.concurrent.TimeUnit; import java.util.concurrent.atomic.AtomicInteger; import java.util.function.Function; @@ -359,6 +360,47 @@ void testRefreshDoesNotCreateTransportWhenShutdown() { assertEquals(0, channel.getGeneration()); } + @Test + void shutdown_waitsForInProgressRefresh() throws Exception { + CountDownLatch refreshStarted = new CountDownLatch(1); + CountDownLatch releaseRefresh = new CountDownLatch(1); + transportFactory = + () -> { + if (transportFactoryCount.incrementAndGet() > 1) { + refreshStarted.countDown(); + try { + releaseRefresh.await(); + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + } + } + return new MockHttpTransport(); + }; + RefreshingHttpJsonChannel channel = createTestChannel(); + FakeManagedHttpJsonChannel underlyingChannel = lastCreatedChannel; + rotateCertificate(channel); + + Thread refreshThread = new Thread(channel::refresh); + Thread shutdownThread = new Thread(channel::shutdown); + try { + refreshThread.start(); + assertTrue(refreshStarted.await(5, TimeUnit.SECONDS)); + shutdownThread.start(); + + shutdownThread.join(200); + assertTrue(shutdownThread.isAlive()); + assertFalse(underlyingChannel.isShutdown()); + } finally { + releaseRefresh.countDown(); + } + refreshThread.join(5000); + shutdownThread.join(5000); + assertFalse(refreshThread.isAlive()); + assertFalse(shutdownThread.isAlive()); + assertEquals(1, channel.getGeneration()); + assertTrue(underlyingChannel.isShutdown()); + } + @Test void testRefreshFactoryExceptionDoesNotWedgeFingerprint() { RefreshingHttpJsonChannel channel = createTestChannel(); From 9a88c40a21f7d597fae1e73f4e120db1331765cb Mon Sep 17 00:00:00 2001 From: Matt Castelaz Date: Fri, 2 Oct 2026 03:45:17 +0000 Subject: [PATCH 27/29] fix(gax-httpjson,gax-grpc): restore base HttpJson transport creation and narrow surefire config - Restore createHttpTransport()/configureMtls() from the base (lost in a rebase) and the two tests that were dropped; keep failing closed when mTLS is enabled but no client certificate is available, in a separate helper used only for channel creation. - Add provider tests for the mTLS channel path and for keeping the current transport when the key store is unavailable during a rotation. - Add a test for refresh() when the certificate file is empty mid-rotation. - gax-grpc surefire: keep both existing exclusions and only drop the stray inclusion pattern that limited the module to a single test. --- sdk-platform-java/gax-java/gax-grpc/pom.xml | 3 +- .../InstantiatingHttpJsonChannelProvider.java | 92 +++++++------ ...tantiatingHttpJsonChannelProviderTest.java | 125 +++++++++++++++++- .../RefreshingHttpJsonChannelTest.java | 29 ++++ 4 files changed, 204 insertions(+), 45 deletions(-) diff --git a/sdk-platform-java/gax-java/gax-grpc/pom.xml b/sdk-platform-java/gax-java/gax-grpc/pom.xml index 0abd59208e3c..f5cc9b3798a3 100644 --- a/sdk-platform-java/gax-java/gax-grpc/pom.xml +++ b/sdk-platform-java/gax-java/gax-grpc/pom.xml @@ -162,7 +162,8 @@ maven-surefire-plugin - !InstantiatingGrpcChannelProviderTest#testLogDirectPathMisconfig_AttemptDirectPathNotSetAndAttemptDirectPathXdsSetViaEnv_warns + !InstantiatingGrpcChannelProviderTest#testLogDirectPathMisconfig_AttemptDirectPathNotSetAndAttemptDirectPathXdsSetViaEnv_warns,!InstantiatingGrpcChannelProviderTest#canUseDirectPath_directPathEnvVarNotSet_attemptDirectPathIsTrue + diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProvider.java b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProvider.java index da4fbc6cd4b8..4ba35b01a8bd 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProvider.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProvider.java @@ -195,49 +195,65 @@ public TransportChannelProvider withCredentials(Credentials credentials) { "InstantiatingHttpJsonChannelProvider doesn't need credentials"); } - @Nullable HttpTransport createHttpTransport() throws IOException, GeneralSecurityException { - if (mtlsProvider == null) { - return null; + HttpTransport createHttpTransport() throws IOException, GeneralSecurityException { + NetHttpTransport.Builder builder = new NetHttpTransport.Builder(); + configureMtls(builder); + HttpJsonConscryptUtils.configureConscryptSecurityProvider(builder); + return builder.build(); + } + + private NetHttpTransport.Builder configureMtls(NetHttpTransport.Builder builder) + throws IOException, GeneralSecurityException { + if (mtlsProvider == null || !certificateBasedAccess.useMtlsClientCertificate()) { + return builder; } - if (certificateBasedAccess.useMtlsClientCertificate()) { - KeyStore mtlsKeyStore = mtlsProvider.getKeyStore(); - if (mtlsKeyStore != null) { - NetHttpTransport.Builder builder = new NetHttpTransport.Builder(); - builder.trustCertificates(null, mtlsKeyStore, ""); - Provider conscryptProvider = HttpJsonConscryptUtils.getConscryptProvider(); - if (conscryptProvider != null) { - SSLContext sslContext = SSLContext.getInstance("TLS", conscryptProvider); - SslUtils.initSslContext( - sslContext, - null, - SslUtils.getPkixTrustManagerFactory(), - mtlsKeyStore, - "", - SslUtils.getDefaultKeyManagerFactory()); - builder.setSslSocketFactory(sslContext.getSocketFactory()); - } - HttpJsonConscryptUtils.configureConscryptSecurityProvider(builder); - return builder.build(); - } + KeyStore mtlsKeyStore = mtlsProvider.getKeyStore(); + if (mtlsKeyStore == null) { + return builder; + } + builder.trustCertificates(null, mtlsKeyStore, ""); + Provider conscryptProvider = HttpJsonConscryptUtils.getConscryptProvider(); + if (conscryptProvider == null) { + // Fall back to standard JDK JSSE if Conscrypt provider is unavailable + return builder; } - return null; + // Explicitly initialize SSLContext with the Conscrypt provider so that the client certificate + // key managers + // and trust manager factory (TMF) are bound to Conscrypt's TLS implementation (supporting PQC + // key exchange). + SSLContext sslContext = SSLContext.getInstance("TLS", conscryptProvider); + SslUtils.initSslContext( + sslContext, + null, + SslUtils.getPkixTrustManagerFactory(), + mtlsKeyStore, + "", + SslUtils.getDefaultKeyManagerFactory()); + builder.setSslSocketFactory(sslContext.getSocketFactory()); + return builder; + } + + private HttpTransport createChannelHttpTransport() throws IOException, GeneralSecurityException { + HttpTransport transport = createHttpTransport(); + if (mtlsProvider != null + && certificateBasedAccess.useMtlsClientCertificate() + && !((NetHttpTransport) transport).isMtls()) { + // mTLS is enabled but the provider returned no client certificate. Fail instead of silently + // using a transport without the certificate, matching InstantiatingGrpcChannelProvider. + // During certificate rotation, this makes RefreshingHttpJsonChannel keep the current + // authenticated transport instead of swapping in one without a client certificate. + throw new IOException("Failed to initialize mTLS HttpTransport"); + } + return transport; } private ManagedHttpJsonChannel createSingleManagedChannel() throws IOException, GeneralSecurityException { - HttpTransport httpTransportToUse = httpTransport; - if (httpTransportToUse == null) { - httpTransportToUse = createHttpTransport(); - if (httpTransportToUse == null - && mtlsProvider != null - && certificateBasedAccess.useMtlsClientCertificate()) { - throw new IOException("Failed to initialize mTLS HttpTransport"); - } - } - return buildManagedChannel(httpTransportToUse); + return buildManagedChannel( + httpTransport != null ? httpTransport : createChannelHttpTransport()); } - private ManagedHttpJsonChannel buildManagedChannel(@Nullable HttpTransport httpTransportToUse) { + private ManagedHttpJsonChannel buildManagedChannel(HttpTransport httpTransportToUse) { return ManagedHttpJsonChannel.newBuilder() .setEndpoint(endpoint) .setExecutor(executor) @@ -258,11 +274,7 @@ private HttpJsonTransportChannel createChannel() throws IOException, GeneralSecu java.util.function.Supplier transportFactory = () -> { try { - HttpTransport mtlsTransport = createHttpTransport(); - if (mtlsTransport == null) { - throw new IOException("Failed to initialize mTLS HttpTransport"); - } - return mtlsTransport; + return createChannelHttpTransport(); } catch (IOException | GeneralSecurityException e) { throw new java.lang.RuntimeException("Failed to create mTLS HttpTransport", e); } diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProviderTest.java b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProviderTest.java index b39af9581357..b63ce6ba04be 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProviderTest.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProviderTest.java @@ -33,12 +33,14 @@ import static org.junit.jupiter.api.Assertions.assertEquals; import static org.mockito.Mockito.mock; +import com.google.api.client.http.javanet.NetHttpTransport; import com.google.api.gax.rpc.HeaderProvider; import com.google.api.gax.rpc.TransportChannelProvider; import com.google.api.gax.rpc.mtls.AbstractMtlsTransportChannelTest; import com.google.api.gax.rpc.mtls.CertificateBasedAccess; import com.google.auth.mtls.MtlsProvider; import java.io.IOException; +import java.nio.charset.StandardCharsets; import java.security.GeneralSecurityException; import java.util.Collections; import java.util.Map; @@ -269,6 +271,94 @@ void channelCreation_withCustomHttpTransport_ignoresWorkloadCertPathAndDoesNotWr httpJsonTransportChannel.shutdownNow(); } + @Test + void getTransportChannel_withMtlsKeyStore_usesMtlsTransport() throws IOException { + Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); + com.google.auth.mtls.MtlsProvider mtlsProvider = + new com.google.api.gax.rpc.testing.FakeMtlsProvider( + com.google.api.gax.rpc.testing.FakeMtlsProvider.createTestMtlsKeyStore(), "", false); + + InstantiatingHttpJsonChannelProvider provider = + InstantiatingHttpJsonChannelProvider.newBuilder() + .setEndpoint(DEFAULT_ENDPOINT) + .setMtlsProvider(mtlsProvider) + .setCertificateBasedAccess(certificateBasedAccess) + .build(); + provider = (InstantiatingHttpJsonChannelProvider) provider.withHeaders(DEFAULT_HEADER_MAP); + + HttpJsonTransportChannel httpJsonTransportChannel = provider.getTransportChannel(); + ManagedHttpJsonInterceptorChannel interceptorChannel = + (ManagedHttpJsonInterceptorChannel) httpJsonTransportChannel.getManagedChannel(); + ManagedHttpJsonChannel channel = + ((ManagedHttpJsonInterceptorChannel) interceptorChannel.getChannel()).getChannel(); + + assertThat(channel).isNotInstanceOf(RefreshingHttpJsonChannel.class); + assertThat(((NetHttpTransport) channel.getHttpTransport()).isMtls()).isTrue(); + + httpJsonTransportChannel.shutdownNow(); + } + + @Test + void refresh_whenKeyStoreUnavailableDuringRotation_keepsCurrentTransport( + @org.junit.jupiter.api.io.TempDir java.nio.file.Path tempDir) throws IOException { + java.nio.file.Path certPath = tempDir.resolve("cert.pem"); + java.nio.file.Files.write(certPath, "certificate-v1".getBytes(StandardCharsets.UTF_8)); + Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); + Mockito.when(certificateBasedAccess.getWorkloadCertPath()).thenReturn(certPath.toString()); + java.security.KeyStore keyStore = + com.google.api.gax.rpc.testing.FakeMtlsProvider.createTestMtlsKeyStore(); + java.util.concurrent.atomic.AtomicReference currentKeyStore = + new java.util.concurrent.atomic.AtomicReference<>(keyStore); + com.google.auth.mtls.MtlsProvider mtlsProvider = + new com.google.auth.mtls.MtlsProvider() { + @Override + public java.security.KeyStore getKeyStore() { + return currentKeyStore.get(); + } + + @Override + public boolean isAvailable() { + return true; + } + }; + + InstantiatingHttpJsonChannelProvider provider = + InstantiatingHttpJsonChannelProvider.newBuilder() + .setEndpoint(DEFAULT_ENDPOINT) + .setMtlsProvider(mtlsProvider) + .setCertificateBasedAccess(certificateBasedAccess) + .build(); + provider = (InstantiatingHttpJsonChannelProvider) provider.withHeaders(DEFAULT_HEADER_MAP); + HttpJsonTransportChannel httpJsonTransportChannel = provider.getTransportChannel(); + ManagedHttpJsonInterceptorChannel interceptorChannel = + (ManagedHttpJsonInterceptorChannel) httpJsonTransportChannel.getManagedChannel(); + RefreshingHttpJsonChannel channel = + (RefreshingHttpJsonChannel) + ((ManagedHttpJsonInterceptorChannel) interceptorChannel.getChannel()).getChannel(); + com.google.api.client.http.HttpTransport initialTransport = channel.getHttpTransport(); + assertThat(((NetHttpTransport) initialTransport).isMtls()).isTrue(); + + // The certificate rotates on disk, but the provider cannot return the new key store yet. + java.nio.file.Files.write(certPath, "certificate-v2".getBytes(StandardCharsets.UTF_8)); + currentKeyStore.set(null); + assertThat(channel.shouldRefresh()).isTrue(); + channel.refresh(); + + // The current authenticated transport is kept instead of one without a client certificate. + assertThat(channel.getHttpTransport()).isSameInstanceAs(initialTransport); + assertThat(channel.getGeneration()).isEqualTo(0); + assertThat(channel.shouldRefresh()).isTrue(); + + // Once the key store is available again, the next refresh swaps in a new mTLS transport. + currentKeyStore.set(keyStore); + channel.refresh(); + assertThat(channel.getHttpTransport()).isNotSameInstanceAs(initialTransport); + assertThat(((NetHttpTransport) channel.getHttpTransport()).isMtls()).isTrue(); + assertThat(channel.getGeneration()).isEqualTo(1); + + httpJsonTransportChannel.shutdownNow(); + } + @Test void getTransportChannel_whenMtlsKeyStoreThrowsIOException_throwsCheckedIOException() throws Exception { @@ -410,7 +500,7 @@ void createHttpTransport_withMtlsAndConscrypt_configuresSecurityProvider() } @Test - void createHttpTransport_whenMtlsProviderNullOrNotUsingClientCert_returnsNull() + void createHttpTransport_whenMtlsProviderNullOrNotUsingClientCert_returnsNonMtlsTransport() throws IOException, GeneralSecurityException { InstantiatingHttpJsonChannelProvider nullMtlsProviderChannelProvider = InstantiatingHttpJsonChannelProvider.newBuilder() @@ -418,7 +508,10 @@ void createHttpTransport_whenMtlsProviderNullOrNotUsingClientCert_returnsNull() .setMtlsProvider(null) .setCertificateBasedAccess(certificateBasedAccess) .build(); - assertThat(nullMtlsProviderChannelProvider.createHttpTransport()).isNull(); + NetHttpTransport nullMtlsProviderTransport = + (NetHttpTransport) nullMtlsProviderChannelProvider.createHttpTransport(); + assertThat(nullMtlsProviderTransport).isNotNull(); + assertThat(nullMtlsProviderTransport.isMtls()).isFalse(); Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(false); com.google.auth.mtls.MtlsProvider provider = @@ -430,7 +523,30 @@ void createHttpTransport_whenMtlsProviderNullOrNotUsingClientCert_returnsNull() .setMtlsProvider(provider) .setCertificateBasedAccess(certificateBasedAccess) .build(); - assertThat(disabledMtlsChannelProvider.createHttpTransport()).isNull(); + NetHttpTransport disabledMtlsTransport = + (NetHttpTransport) disabledMtlsChannelProvider.createHttpTransport(); + assertThat(disabledMtlsTransport).isNotNull(); + assertThat(disabledMtlsTransport.isMtls()).isFalse(); + } + + @Test + void testCreateHttpTransport_returnsValidTransport() throws Exception { + InstantiatingHttpJsonChannelProvider channelProvider = + InstantiatingHttpJsonChannelProvider.newBuilder() + .setEndpoint("localhost:8080") + .setHeaderProvider(Collections::emptyMap) + .setExecutor(Runnable::run) + .build(); + NetHttpTransport transport = (NetHttpTransport) channelProvider.createHttpTransport(); + assertThat(transport).isNotNull(); + } + + @Test + void testConfigureConscryptSecurityProvider_returnsConfiguredBuilder() { + NetHttpTransport.Builder builder = new NetHttpTransport.Builder(); + NetHttpTransport.Builder result = + HttpJsonConscryptUtils.configureConscryptSecurityProvider(builder); + assertThat(result).isSameInstanceAs(builder); } @Override @@ -446,6 +562,7 @@ protected Object getMtlsObjectFromTransportChannel( mock(HeaderProvider.class, Mockito.withSettings().withoutAnnotations())) .setExecutor(mock(Executor.class)) .build(); - return channelProvider.createHttpTransport(); + NetHttpTransport transport = (NetHttpTransport) channelProvider.createHttpTransport(); + return (transport != null && transport.isMtls()) ? transport : null; } } diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java index cc1dc72d5453..bf7b33da84ba 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java @@ -301,6 +301,35 @@ void refresh_swapsTransportAndKeepsChannel() { assertFalse(channel.shouldRefresh()); } + @Test + void refresh_whenCertificateFileEmpty_keepsTransportAndGeneration() { + RefreshingHttpJsonChannel channel = createTestChannel(); + HttpTransport initialTransport = channel.getHttpTransport(); + + // A rotation is detected, but when refresh() re-reads the file the rotator has truncated it + // and not yet written the new certificate. + rotateCertificate(channel); + assertTrue(channel.shouldRefresh()); + testFingerprint = ""; + channel.refresh(); + + // Nothing is swapped and the generation is unchanged, so the failed call is not retried. + assertEquals(1, transportFactoryCount.get()); + assertSame(initialTransport, channel.getHttpTransport()); + assertEquals(0, channel.getGeneration()); + // While the file stays empty, no refresh is requested. + channel.invalidateDiskFingerprintCache(); + assertFalse(channel.shouldRefresh()); + + // Once the new certificate is written, the next check refreshes normally. + testFingerprint = "fingerprint2"; + assertTrue(channel.shouldRefresh()); + channel.refresh(); + assertEquals(2, transportFactoryCount.get()); + assertNotSame(initialTransport, channel.getHttpTransport()); + assertEquals(1, channel.getGeneration()); + } + @Test void callCreatedBeforeRefresh_usesOriginalTransport() throws Exception { MockHttpService originalService = From 65bf74fcd600058523b39049c7864cc1b310735a Mon Sep 17 00:00:00 2001 From: Matt Castelaz Date: Fri, 2 Oct 2026 21:44:49 +0000 Subject: [PATCH 28/29] fix(gax): address round-4 review feedback on rotation retries - Check the channel generation before reading the certificate from disk, and keep shouldRefresh() inside the try block. - Only advance the ChannelPool generation when every channel switched to a new certificate, so periodic refreshes don't enable rotation retries. - Mark rotation retries with a package-private channelRefreshed flag on UnauthenticatedException instead of reusing isRetryable(), so configured UNAUTHENTICATED retries and per-call retry codes keep their behaviour. Pin serialVersionUID to the released value; the flag is transient. - Log the new certificate fingerprint after a successful switch. - Cancel resize/refresh futures before taking entryWriteLock on shutdown so an in-progress refresh is interrupted. --- .../com/google/api/gax/grpc/ChannelPool.java | 87 ++++--- .../google/api/gax/grpc/ChannelPoolTest.java | 218 +++++++++++++++++- .../api/gax/rpc/ApiResultRetryAlgorithm.java | 22 +- .../google/api/gax/rpc/AttemptCallable.java | 24 +- .../rpc/ServerStreamingAttemptCallable.java | 24 +- .../api/gax/rpc/UnauthenticatedException.java | 49 ++++ .../gax/rpc/ApiResultRetryAlgorithmTest.java | 180 +++++++++++++-- .../api/gax/rpc/AttemptCallableTest.java | 117 +++++++++- .../ServerStreamingAttemptCallableTest.java | 84 ++++++- .../gax/rpc/UnauthenticatedExceptionTest.java | 183 +++++++++++++++ 10 files changed, 884 insertions(+), 104 deletions(-) create mode 100644 sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/UnauthenticatedExceptionTest.java diff --git a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java index 0e634fcd43c3..f7b2d76d9d8b 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/main/java/com/google/api/gax/grpc/ChannelPool.java @@ -34,7 +34,6 @@ import com.google.api.gax.rpc.mtls.CertificateRotationTracker; import com.google.common.annotations.VisibleForTesting; import com.google.common.base.Preconditions; -import com.google.common.base.Strings; import com.google.common.collect.ImmutableList; import io.grpc.CallOptions; import io.grpc.Channel; @@ -203,18 +202,19 @@ Channel getChannel(int affinity) { public ManagedChannel shutdown() { LOG.fine("Initiating graceful shutdown due to explicit request"); + // Resize and refresh tasks can block on channel priming. We don't need + // to wait for the channels to be ready since we're shutting down the + // pool. Allowing interrupt to speed it up. This is done before acquiring + // entryWriteLock, which a running refresh or resize holds. + if (resizeFuture != null) { + resizeFuture.cancel(true); + } + if (refreshFuture != null) { + refreshFuture.cancel(true); + } + synchronized (entryWriteLock) { isShutdown = true; - // Resize and refresh tasks can block on channel priming. We don't need - // to wait for the channels to be ready since we're shutting down the - // pool. Allowing interrupt to speed it up. - if (resizeFuture != null) { - resizeFuture.cancel(true); - } - if (refreshFuture != null) { - refreshFuture.cancel(true); - } - List localEntries = entries.get(); for (Entry entry : localEntries) { entry.channel.shutdown(); @@ -262,15 +262,17 @@ public boolean isTerminated() { public ManagedChannel shutdownNow() { LOG.fine("Initiating immediate shutdown due to explicit request"); + // Cancel before acquiring entryWriteLock, which a running refresh or resize holds, so that + // they are interrupted instead of delaying shutdown. + if (resizeFuture != null) { + resizeFuture.cancel(true); + } + if (refreshFuture != null) { + refreshFuture.cancel(true); + } + synchronized (entryWriteLock) { isShutdown = true; - if (resizeFuture != null) { - resizeFuture.cancel(true); - } - if (refreshFuture != null) { - refreshFuture.cancel(true); - } - List localEntries = entries.get(); for (Entry entry : localEntries) { entry.channel.shutdownNow(); @@ -454,6 +456,10 @@ private void expand(int desiredSize) { * disconnects). This applies to all channels even when {@code workloadCertPath == null}. If * {@code workloadCertPath} is configured, also updates the tracked certificate fingerprint on * success (or skips if the certificate file is currently unreadable or mid-write on disk). + * + *

The generation is only advanced if this refresh switched every channel to a certificate + * different from the active one, so that a periodic refresh without a rotation does not make + * in-flight {@code UNAUTHENTICATED} failures eligible for a rotation retry. */ private void refreshSafely() { try { @@ -462,8 +468,15 @@ private void refreshSafely() { if (workloadCertPath != null && currentDiskFingerprint.isEmpty()) { return; } + boolean rotated = + !currentDiskFingerprint.isEmpty() + && !rotationTracker.isAlreadyActive(currentDiskFingerprint); if (refreshAll() && !currentDiskFingerprint.isEmpty()) { - rotationTracker.markRefreshed(currentDiskFingerprint); + if (rotated) { + completeCertificateSwitch(currentDiskFingerprint); + } else { + rotationTracker.markRefreshed(currentDiskFingerprint); + } } } } catch (Exception e) { @@ -522,11 +535,25 @@ void refresh() { // Drop any channel that fails to refresh so that no traffic is routed to the old certificate. if (refreshAll(/* dropUnrefreshedChannels= */ true)) { - rotationTracker.markRefreshed(currentDiskFingerprint); + completeCertificateSwitch(currentDiskFingerprint); } } } + /** + * Records that every channel in the pool now uses the certificate with the given fingerprint. + * Must be called while holding {@code entryWriteLock}, after the channels have been swapped. + * + *

The generation is incremented before the fingerprint is marked active, so that a concurrent + * failing RPC that no longer sees a pending rotation ({@link #shouldRefresh()} is {@code false}) + * is guaranteed to see the new generation and be retried on the new channels. + */ + private void completeCertificateSwitch(String newFingerprint) { + generation.incrementAndGet(); + rotationTracker.markRefreshed(newFingerprint); + LOG.fine("Channel pool switched to certificate with fingerprint: " + newFingerprint); + } + @InternalApi("Visible for testing") @Nullable String getWorkloadCertPath() { return workloadCertPath; @@ -555,12 +582,7 @@ boolean refreshAll(boolean dropUnrefreshedChannels) { if (isShutdown) { return false; } - String activeFingerprint = rotationTracker.getActiveCertFingerprint(); - LOG.fine( - "Refreshing all channels" - + (Strings.isNullOrEmpty(activeFingerprint) - ? "" - : " with certificate fingerprint: " + activeFingerprint)); + LOG.fine("Refreshing all channels"); ArrayList newEntries = new ArrayList<>(entries.get()); boolean anyCreated = false; boolean allCreated = !newEntries.isEmpty(); @@ -599,7 +621,6 @@ boolean refreshAll(boolean dropUnrefreshedChannels) { e.requestShutdown(); } } - generation.incrementAndGet(); if (dropUnrefreshedChannels && !allCreated && settings.isStaticSize()) { scheduleRefill(); } @@ -645,11 +666,13 @@ void refillSafely() { /** * Returns the current channel pool generation counter. * - *

The generation is a monotonically increasing counter incremented each time {@link - * #refreshAll()} replaces the channels in the pool. Retry loops ({@code AttemptCallable} and - * {@code ServerStreamingAttemptCallable}) snapshot the generation before starting an RPC attempt - * and compare it after an {@code UNAUTHENTICATED} failure to determine whether the pool rotated - * to a new certificate generation during or after the attempt. + *

The generation is a monotonically increasing counter incremented each time the pool switches + * every channel to a new mTLS certificate (a reactive refresh after a rotation, or a periodic + * refresh that picks up a rotated certificate). Periodic refreshes that do not change the + * certificate do not increment it. Retry loops ({@code AttemptCallable} and {@code + * ServerStreamingAttemptCallable}) snapshot the generation before starting an RPC attempt and + * compare it after an {@code UNAUTHENTICATED} failure to determine whether the pool rotated to a + * new certificate during or after the attempt. */ long getGeneration() { return generation.get(); diff --git a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java index 0fe97f39c376..e7e386e2304c 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java @@ -47,6 +47,7 @@ import com.google.api.gax.rpc.StreamController; import com.google.api.gax.rpc.UnaryCallSettings; import com.google.api.gax.rpc.UnaryCallable; +import com.google.api.gax.rpc.mtls.CertificateRotationTracker; import com.google.api.gax.util.FakeLogHandler; import com.google.auth.Credentials; import com.google.common.collect.ImmutableList; @@ -65,13 +66,16 @@ import java.util.Arrays; import java.util.List; import java.util.concurrent.CancellationException; +import java.util.concurrent.CountDownLatch; import java.util.concurrent.ExecutorService; import java.util.concurrent.Executors; import java.util.concurrent.RejectedExecutionException; import java.util.concurrent.ScheduledExecutorService; import java.util.concurrent.ScheduledFuture; import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicBoolean; import java.util.concurrent.atomic.AtomicInteger; +import java.util.logging.Level; import org.junit.jupiter.api.AfterEach; import org.junit.jupiter.api.Assertions; import org.junit.jupiter.api.Test; @@ -701,7 +705,7 @@ void refreshAll_partialFailureWithoutRotation_keepsOldChannelAndDoesNotScheduleR assertThat(pool.entries.get()).hasSize(2); Mockito.verify(initial1).shutdown(); Mockito.verify(initial2, Mockito.never()).shutdown(); - assertThat(pool.getGeneration()).isEqualTo(genBefore + 1); + assertThat(pool.getGeneration()).isEqualTo(genBefore); Mockito.verify(executor, Mockito.never()).execute(Mockito.any(Runnable.class)); } @@ -911,7 +915,7 @@ void refresh_onShutdownPool_noOpsAndCreatesNoChannels() throws IOException { } @Test - void generationCounterIncrementsOnRefresh() throws IOException { + void refreshAll_doesNotIncrementGeneration() throws IOException { ManagedChannel channel1 = mock(ManagedChannel.class); ManagedChannel channel2 = mock(ManagedChannel.class); ChannelFactory channelFactory = @@ -921,8 +925,216 @@ void generationCounterIncrementsOnRefresh() throws IOException { pool = ChannelPool.create(ChannelPoolSettings.staticallySized(1), channelFactory, null, null); assertThat(pool.getGeneration()).isEqualTo(0); - pool.refreshAll(); + // Only a switch to a new certificate advances the generation. + assertThat(pool.refreshAll()).isTrue(); + Mockito.verify(channel1).shutdown(); + assertThat(pool.getGeneration()).isEqualTo(0); + } + + /** + * Creates an mTLS pool with preemptive refresh enabled and returns the scheduled periodic refresh + * task. + */ + private Runnable createPreemptiveRefreshMtlsPool(int size, ChannelFactory channelFactory) + throws IOException { + List refreshTasks = new ArrayList<>(); + ScheduledExecutorService executor = mockExecutor(); + Mockito.doAnswer( + invocation -> { + refreshTasks.add(invocation.getArgument(0)); + return Mockito.mock( + ScheduledFuture.class, Mockito.withSettings().withoutAnnotations()); + }) + .when(executor) + .scheduleAtFixedRate( + Mockito.any(Runnable.class), Mockito.anyLong(), Mockito.anyLong(), Mockito.any()); + writeCert("client_cert.pem"); + pool = + new ChannelPool( + ChannelPoolSettings.staticallySized(size).toBuilder() + .setPreemptiveRefreshEnabled(true) + .build(), + channelFactory, + FixedExecutorProvider.create(executor), + tempCert.toString()); + assertThat(refreshTasks).hasSize(1); + return refreshTasks.get(0); + } + + private String readCertFingerprint() { + return new CertificateRotationTracker(tempCert.toString()).readDiskFingerprint(); + } + + @Test + void preemptiveRefresh_withoutRotation_doesNotIncrementGeneration() throws IOException { + ManagedChannel initial = Mockito.mock(ManagedChannel.class); + ManagedChannel refreshed = Mockito.mock(ManagedChannel.class); + ChannelFactory channelFactory = mockChannelFactory(); + Mockito.when(channelFactory.createSingleChannel()).thenReturn(initial, refreshed); + Runnable preemptiveRefresh = createPreemptiveRefreshMtlsPool(1, channelFactory); + + FakeLogHandler logHandler = new FakeLogHandler(); + Level originalLevel = ChannelPool.LOG.getLevel(); + ChannelPool.LOG.setLevel(Level.FINE); + ChannelPool.LOG.addHandler(logHandler); + try { + pool.invalidateDiskFingerprintCache(); + preemptiveRefresh.run(); + } finally { + ChannelPool.LOG.removeHandler(logHandler); + ChannelPool.LOG.setLevel(originalLevel); + } + + // The channels are replaced, but the certificate did not change. + Mockito.verify(initial).shutdown(); + assertThat(pool.getGeneration()).isEqualTo(0); + assertThat(logHandler.getAllMessages()).contains("Refreshing all channels"); + assertThat(String.join("\n", logHandler.getAllMessages())) + .doesNotContain("Channel pool switched to certificate"); + } + + @Test + void preemptiveRefresh_pickingUpRotatedCert_incrementsGeneration() throws IOException { + ManagedChannel initial = Mockito.mock(ManagedChannel.class); + ManagedChannel rotated = Mockito.mock(ManagedChannel.class); + ChannelFactory channelFactory = mockChannelFactory(); + Mockito.when(channelFactory.createSingleChannel()).thenReturn(initial, rotated); + Runnable preemptiveRefresh = createPreemptiveRefreshMtlsPool(1, channelFactory); + + pool.invalidateDiskFingerprintCache(); + writeCert("root_cert.pem"); + preemptiveRefresh.run(); + + Mockito.verify(initial).shutdown(); assertThat(pool.getGeneration()).isEqualTo(1); + pool.invalidateDiskFingerprintCache(); + assertThat(pool.shouldRefresh()).isFalse(); + } + + @Test + void preemptiveRefresh_partialFailureDuringRotation_doesNotIncrementOrMarkRefreshed() + throws IOException { + ManagedChannel initial1 = Mockito.mock(ManagedChannel.class); + ManagedChannel initial2 = Mockito.mock(ManagedChannel.class); + ManagedChannel rotated1 = Mockito.mock(ManagedChannel.class); + ChannelFactory channelFactory = mockChannelFactory(); + Mockito.when(channelFactory.createSingleChannel()) + .thenReturn(initial1, initial2) + .thenReturn(rotated1) + .thenThrow(new IOException("Transient failure on second sub-channel")); + Runnable preemptiveRefresh = createPreemptiveRefreshMtlsPool(2, channelFactory); + + pool.invalidateDiskFingerprintCache(); + writeCert("root_cert.pem"); + preemptiveRefresh.run(); + + // One channel still uses the old certificate, so the switch is not complete. + Mockito.verify(initial1).shutdown(); + Mockito.verify(initial2, Mockito.never()).shutdown(); + assertThat(pool.getGeneration()).isEqualTo(0); + pool.invalidateDiskFingerprintCache(); + assertThat(pool.shouldRefresh()).isTrue(); + } + + @Test + void refresh_whenAlreadyActive_doesNotIncrementGeneration() throws IOException { + ChannelFactory channelFactory = mockChannelFactory(); + Mockito.when(channelFactory.createSingleChannel()) + .thenReturn(Mockito.mock(ManagedChannel.class), Mockito.mock(ManagedChannel.class)); + writeCert("client_cert.pem"); + pool = + new ChannelPool( + ChannelPoolSettings.staticallySized(1), + channelFactory, + FixedExecutorProvider.create(mockExecutor()), + tempCert.toString()); + + pool.invalidateDiskFingerprintCache(); + pool.refresh(); + + assertThat(pool.getGeneration()).isEqualTo(0); + Mockito.verify(channelFactory, Mockito.times(1)).createSingleChannel(); + } + + @Test + void refresh_onRotation_logsNewCertificateFingerprint() throws IOException { + ChannelFactory channelFactory = mockChannelFactory(); + Mockito.when(channelFactory.createSingleChannel()) + .thenReturn(Mockito.mock(ManagedChannel.class), Mockito.mock(ManagedChannel.class)); + writeCert("client_cert.pem"); + String oldFingerprint = readCertFingerprint(); + createMtlsPoolAndRotateCert( + ChannelPoolSettings.staticallySized(1), channelFactory, mockExecutor()); + String newFingerprint = readCertFingerprint(); + assertThat(newFingerprint).isNotEqualTo(oldFingerprint); + + FakeLogHandler logHandler = new FakeLogHandler(); + Level originalLevel = ChannelPool.LOG.getLevel(); + ChannelPool.LOG.setLevel(Level.FINE); + ChannelPool.LOG.addHandler(logHandler); + try { + pool.refresh(); + } finally { + ChannelPool.LOG.removeHandler(logHandler); + ChannelPool.LOG.setLevel(originalLevel); + } + + assertThat(pool.getGeneration()).isEqualTo(1); + assertThat(logHandler.getAllMessages()) + .contains("Channel pool switched to certificate with fingerprint: " + newFingerprint); + assertThat(String.join("\n", logHandler.getAllMessages())).doesNotContain(oldFingerprint); + } + + @Test + void shutdown_interruptsInProgressRefresh() throws Exception { + CountDownLatch refreshStarted = new CountDownLatch(1); + CountDownLatch release = new CountDownLatch(1); + AtomicBoolean refreshInterrupted = new AtomicBoolean(); + AtomicInteger createdChannels = new AtomicInteger(); + ChannelFactory channelFactory = + () -> { + if (createdChannels.getAndIncrement() > 0) { + // The preemptive refresh blocks while creating its replacement channel. + refreshStarted.countDown(); + try { + release.await(); + } catch (InterruptedException e) { + refreshInterrupted.set(true); + Thread.currentThread().interrupt(); + throw new IOException("Interrupted while creating channel", e); + } + } + return Mockito.mock(ManagedChannel.class); + }; + ScheduledExecutorService realExecutor = Executors.newSingleThreadScheduledExecutor(); + ScheduledExecutorService executor = mockExecutor(); + Mockito.doAnswer( + invocation -> + realExecutor.schedule( + (Runnable) invocation.getArgument(0), 0, TimeUnit.MILLISECONDS)) + .when(executor) + .scheduleAtFixedRate( + Mockito.any(Runnable.class), Mockito.anyLong(), Mockito.anyLong(), Mockito.any()); + try { + pool = + new ChannelPool( + ChannelPoolSettings.staticallySized(1).toBuilder() + .setPreemptiveRefreshEnabled(true) + .build(), + channelFactory, + FixedExecutorProvider.create(executor), + null); + assertThat(refreshStarted.await(5, TimeUnit.SECONDS)).isTrue(); + + // The refresh holds the pool's write lock; shutdown must interrupt it rather than wait. + Assertions.assertTimeoutPreemptively(java.time.Duration.ofSeconds(5), () -> pool.shutdown()); + + assertThat(refreshInterrupted.get()).isTrue(); + assertThat(pool.isShutdown()).isTrue(); + } finally { + release.countDown(); + realExecutor.shutdownNow(); + } } @Test diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ApiResultRetryAlgorithm.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ApiResultRetryAlgorithm.java index 91b332d10660..c8fdbe15f9ee 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ApiResultRetryAlgorithm.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ApiResultRetryAlgorithm.java @@ -54,8 +54,7 @@ class ApiResultRetryAlgorithm extends BasicResultRetryAlgorithm extends BasicResultRetryAlgorithm extends BasicResultRetryAlgorithm { TransportChannel channel = finalContext.getTransportChannel(); if (channel != null) { - if (channel.shouldRefresh()) { + // If another request already refreshed the channel, retry without checking the + // certificate on disk again. Otherwise, check for a rotation and refresh. + boolean shouldRetry = channel.getGeneration() > attemptGeneration; + if (!shouldRetry) { try { - channel.refresh(); + if (channel.shouldRefresh()) { + channel.refresh(); + } } catch (Exception e) { LOG.log( Level.WARNING, "Failed to refresh transport channel after authentication error", e); } + shouldRetry = channel.getGeneration() > attemptGeneration; } - boolean shouldRetry = channel.getGeneration() > attemptGeneration; if (shouldRetry) { - UnauthenticatedException newEx = - new UnauthenticatedException( - unauthenticatedException.getMessage(), - unauthenticatedException.getCause(), - unauthenticatedException.getStatusCode(), - true, // isRetryable = true - unauthenticatedException.getErrorDetails()); - newEx.setStackTrace(unauthenticatedException.getStackTrace()); - for (Throwable suppressed : unauthenticatedException.getSuppressed()) { - newEx.addSuppressed(suppressed); - } - throw newEx; + throw unauthenticatedException.withChannelRefreshed(); } } throw unauthenticatedException; diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ServerStreamingAttemptCallable.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ServerStreamingAttemptCallable.java index b8ccd565ac5b..efc374a50ec3 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ServerStreamingAttemptCallable.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/ServerStreamingAttemptCallable.java @@ -252,30 +252,24 @@ public void onErrorImpl(Throwable t) { UnauthenticatedException unauthenticatedException = (UnauthenticatedException) cause; TransportChannel transportChannel = finalContext.getTransportChannel(); if (transportChannel != null) { - if (transportChannel.shouldRefresh()) { + // If another request already refreshed the channel, retry without checking the + // certificate on disk again. Otherwise, check for a rotation and refresh. + boolean shouldRetry = transportChannel.getGeneration() > attemptGeneration; + if (!shouldRetry) { try { - transportChannel.refresh(); + if (transportChannel.shouldRefresh()) { + transportChannel.refresh(); + } } catch (Exception e) { LOG.log( Level.WARNING, "Failed to refresh transport channel after authentication error", e); } + shouldRetry = transportChannel.getGeneration() > attemptGeneration; } - boolean shouldRetry = transportChannel.getGeneration() > attemptGeneration; if (shouldRetry) { - UnauthenticatedException newEx = - new UnauthenticatedException( - unauthenticatedException.getMessage(), - unauthenticatedException.getCause(), - unauthenticatedException.getStatusCode(), - true, - unauthenticatedException.getErrorDetails()); - newEx.setStackTrace(unauthenticatedException.getStackTrace()); - for (Throwable suppressed : unauthenticatedException.getSuppressed()) { - newEx.addSuppressed(suppressed); - } - t = newEx; + t = unauthenticatedException.withChannelRefreshed(); } } } diff --git a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/UnauthenticatedException.java b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/UnauthenticatedException.java index 0c93f07e6b51..d18989264572 100644 --- a/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/UnauthenticatedException.java +++ b/sdk-platform-java/gax-java/gax/src/main/java/com/google/api/gax/rpc/UnauthenticatedException.java @@ -30,6 +30,7 @@ package com.google.api.gax.rpc; import org.jspecify.annotations.NullMarked; +import org.jspecify.annotations.Nullable; /** * Exception thrown when the request does not have valid authentication credentials for the @@ -37,18 +38,33 @@ */ @NullMarked public class UnauthenticatedException extends ApiException { + // Pinned to the value computed for previous releases (gax 2.83.0 to 2.87.0) so that adding + // members does not break Java serialization compatibility with them. + private static final long serialVersionUID = 6971115068105015909L; + + /** + * Whether this failure happened on a transport channel that has since been refreshed (for + * example, after an mTLS certificate rotation), making the request eligible for a single + * immediate retry on the refreshed channel. Only meaningful within the process that observed the + * failure, so it is not serialized. + */ + private final transient boolean channelRefreshed; + public UnauthenticatedException(Throwable cause, StatusCode statusCode, boolean retryable) { super(cause, statusCode, retryable); + this.channelRefreshed = false; } public UnauthenticatedException( String message, Throwable cause, StatusCode statusCode, boolean retryable) { super(message, cause, statusCode, retryable); + this.channelRefreshed = false; } public UnauthenticatedException( Throwable cause, StatusCode statusCode, boolean retryable, ErrorDetails errorDetails) { super(cause, statusCode, retryable, errorDetails); + this.channelRefreshed = false; } public UnauthenticatedException( @@ -58,5 +74,38 @@ public UnauthenticatedException( boolean retryable, ErrorDetails errorDetails) { super(message, cause, statusCode, retryable, errorDetails); + this.channelRefreshed = false; + } + + private UnauthenticatedException( + @Nullable String message, + @Nullable Throwable cause, + StatusCode statusCode, + boolean retryable, + @Nullable ErrorDetails errorDetails, + boolean channelRefreshed) { + super(message, cause, statusCode, retryable, errorDetails); + this.channelRefreshed = channelRefreshed; + } + + /** Returns whether this failure happened on a transport channel that has since been refreshed. */ + boolean isChannelRefreshed() { + return channelRefreshed; + } + + /** + * Returns a copy of this exception marked as having happened on a channel that has since been + * refreshed. The copy keeps the message, cause, status code, {@link #isRetryable()} value, error + * details, stack trace and suppressed exceptions of this exception. + */ + UnauthenticatedException withChannelRefreshed() { + UnauthenticatedException newEx = + new UnauthenticatedException( + getMessage(), getCause(), getStatusCode(), isRetryable(), getErrorDetails(), true); + newEx.setStackTrace(getStackTrace()); + for (Throwable suppressed : getSuppressed()) { + newEx.addSuppressed(suppressed); + } + return newEx; } } diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ApiResultRetryAlgorithmTest.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ApiResultRetryAlgorithmTest.java index 1232be098e52..200aaddf7730 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ApiResultRetryAlgorithmTest.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ApiResultRetryAlgorithmTest.java @@ -143,9 +143,7 @@ void testRotationRetryWithNonRetryableSettings_maxAttemptsOne() { RetryAlgorithm retryAlgorithm = new RetryAlgorithm<>(resultAlgorithm, timedAlgorithm); TimedAttemptSettings firstAttempt = retryAlgorithm.createFirstAttempt(context); - UnauthenticatedException rotationEx = - new UnauthenticatedException( - "Expired cert", null, new FakeStatusCode(Code.UNAUTHENTICATED), /* retryable= */ true); + UnauthenticatedException rotationEx = rotationException(); // First rotation failure: grants immediate free retry without incrementing attemptCount TimedAttemptSettings nextAttempt = @@ -181,9 +179,7 @@ void testRotationRetryWithNonRetryableSettings_zeroMaxAttemptsZeroTotalTimeout() RetryAlgorithm retryAlgorithm = new RetryAlgorithm<>(resultAlgorithm, timedAlgorithm); TimedAttemptSettings firstAttempt = retryAlgorithm.createFirstAttempt(context); - UnauthenticatedException rotationEx = - new UnauthenticatedException( - "Expired cert", null, new FakeStatusCode(Code.UNAUTHENTICATED), /* retryable= */ true); + UnauthenticatedException rotationEx = rotationException(); TimedAttemptSettings nextAttempt = retryAlgorithm.createNextAttempt(context, rotationEx, null, firstAttempt); @@ -223,9 +219,7 @@ void testRotationRetryAfterTransientErrorPreservesRemainingBudget() { TimedAttemptSettings attempt0 = retryAlgorithm.createFirstAttempt(context); ApiException unavailableEx = new ApiException(null, new FakeStatusCode(Code.UNAVAILABLE), /* retryable= */ true); - UnauthenticatedException rotationEx = - new UnauthenticatedException( - "Expired cert", null, new FakeStatusCode(Code.UNAUTHENTICATED), /* retryable= */ true); + UnauthenticatedException rotationEx = rotationException(); // Attempt 0 fails with UNAVAILABLE -> normal retry (attemptCount = 1, overallAttemptCount = 1) TimedAttemptSettings attempt1 = @@ -273,9 +267,7 @@ void testSecondRotationFailureStopsWithTotalTimeoutAndNoMaxAttempts() { RetryAlgorithm retryAlgorithm = new RetryAlgorithm<>(resultAlgorithm, timedAlgorithm); TimedAttemptSettings firstAttempt = retryAlgorithm.createFirstAttempt(context); - UnauthenticatedException rotationEx = - new UnauthenticatedException( - "Expired cert", null, new FakeStatusCode(Code.UNAUTHENTICATED), /* retryable= */ true); + UnauthenticatedException rotationEx = rotationException(); TimedAttemptSettings rotationRetry = retryAlgorithm.createNextAttempt(context, rotationEx, null, firstAttempt); @@ -310,9 +302,7 @@ void testSecondRotationFailureDoesNotConsumeNormalRetryBudget() { RetryAlgorithm retryAlgorithm = new RetryAlgorithm<>(resultAlgorithm, timedAlgorithm); TimedAttemptSettings firstAttempt = retryAlgorithm.createFirstAttempt(context); - UnauthenticatedException rotationEx = - new UnauthenticatedException( - "Expired cert", null, new FakeStatusCode(Code.UNAUTHENTICATED), /* retryable= */ true); + UnauthenticatedException rotationEx = rotationException(); TimedAttemptSettings rotationRetry = retryAlgorithm.createNextAttempt(context, rotationEx, null, firstAttempt); @@ -347,9 +337,7 @@ void testStreamRotationRetryIsAvailableAgainAfterProgress() { new ExponentialRetryAlgorithm(settings, NanoClock.getDefaultClock())); ApiException unavailableEx = new ApiException(null, new FakeStatusCode(Code.UNAVAILABLE), /* retryable= */ true); - UnauthenticatedException rotationEx = - new UnauthenticatedException( - "Expired cert", null, new FakeStatusCode(Code.UNAUTHENTICATED), /* retryable= */ true); + UnauthenticatedException rotationEx = rotationException(); // The first attempt fails with UNAVAILABLE before receiving any messages: normal retry. TimedAttemptSettings attempt1 = @@ -451,4 +439,160 @@ void testStreamProgressResetReturnsNullWhenNotRetryable() { retryAlgorithm.createFirstAttempt(context)); assertNull(next); } + + @Test + void testConfiguredUnauthenticated_notFlagged_usesNormalBackoff() { + // No per-call retryable codes: the method's configuration is carried by isRetryable(). + ApiCallContext context = + mock(ApiCallContext.class, Mockito.withSettings().withoutAnnotations()); + when(context.getRetryableCodes()).thenReturn(null); + RetrySettings settings = + RetrySettings.newBuilder() + .setMaxAttempts(5) + .setInitialRetryDelayDuration(Duration.ofMillis(100)) + .setRetryDelayMultiplier(2.0) + .setMaxRetryDelayDuration(Duration.ofSeconds(1)) + .setTotalTimeoutDuration(Duration.ofMinutes(10)) + .build(); + RetryAlgorithm retryAlgorithm = + new RetryAlgorithm<>( + new ApiResultRetryAlgorithm(), + new ExponentialRetryAlgorithm(settings, NanoClock.getDefaultClock())); + // UNAUTHENTICATED configured as retryable, not caused by a channel refresh. + UnauthenticatedException configuredEx = + new UnauthenticatedException( + "Invalid token", null, new FakeStatusCode(Code.UNAUTHENTICATED), /* retryable= */ true); + + TimedAttemptSettings attempt = retryAlgorithm.createFirstAttempt(context); + for (int i = 1; i < 5; i++) { + attempt = retryAlgorithm.createNextAttempt(context, configuredEx, null, attempt); + assertNotNull(attempt); + assertEquals(i, attempt.getAttemptCount()); + assertEquals(i, attempt.getOverallAttemptCount()); + assertTrue(attempt.getRetryDelayDuration().compareTo(Duration.ZERO) > 0); + assertTrue(retryAlgorithm.shouldRetry(context, configuredEx, null, attempt)); + } + // The fifth failure exhausts maxAttempts = 5. + attempt = retryAlgorithm.createNextAttempt(context, configuredEx, null, attempt); + assertFalse(retryAlgorithm.shouldRetry(context, configuredEx, null, attempt)); + } + + @Test + void testContextRetryableCodesExcludeUnauthenticated_notFlagged_notRetried() { + ApiCallContext context = + mock(ApiCallContext.class, Mockito.withSettings().withoutAnnotations()); + when(context.getRetryableCodes()).thenReturn(Sets.newHashSet(Code.UNAVAILABLE)); + // Retryable according to the method's configuration, but the per-call codes exclude it. + UnauthenticatedException configuredEx = + new UnauthenticatedException( + "Invalid token", null, new FakeStatusCode(Code.UNAUTHENTICATED), /* retryable= */ true); + + ApiResultRetryAlgorithm algorithm = new ApiResultRetryAlgorithm<>(); + assertFalse(algorithm.shouldRetry(context, configuredEx, null)); + assertNull( + algorithm.createNextAttempt( + context, + configuredEx, + null, + new ExponentialRetryAlgorithm( + RetrySettings.newBuilder().setMaxAttempts(5).build(), + NanoClock.getDefaultClock()) + .createFirstAttempt())); + } + + @Test + void testFlagged_nullContext_retriesOnce() { + RetrySettings settings = RetrySettings.newBuilder().setMaxAttempts(1).build(); + RetryAlgorithm retryAlgorithm = + new RetryAlgorithm<>( + new ApiResultRetryAlgorithm(), + new ExponentialRetryAlgorithm(settings, NanoClock.getDefaultClock())); + UnauthenticatedException rotationEx = rotationException(); + + TimedAttemptSettings rotationRetry = + retryAlgorithm.createNextAttempt( + null, rotationEx, null, retryAlgorithm.createFirstAttempt(null)); + assertNotNull(rotationRetry); + assertEquals(Duration.ZERO, rotationRetry.getRetryDelayDuration()); + assertEquals(0, rotationRetry.getAttemptCount()); + assertEquals(1, rotationRetry.getOverallAttemptCount()); + assertTrue(retryAlgorithm.shouldRetry(null, rotationEx, null, rotationRetry)); + + TimedAttemptSettings afterSecondFailure = + retryAlgorithm.createNextAttempt(null, rotationEx, null, rotationRetry); + assertFalse(retryAlgorithm.shouldRetry(null, rotationEx, null, afterSecondFailure)); + } + + @Test + void testFlagged_contextWithoutRetryableCodes_retriesOnce() { + ApiCallContext context = + mock(ApiCallContext.class, Mockito.withSettings().withoutAnnotations()); + when(context.getRetryableCodes()).thenReturn(null); + RetrySettings settings = RetrySettings.newBuilder().setMaxAttempts(1).build(); + RetryAlgorithm retryAlgorithm = + new RetryAlgorithm<>( + new ApiResultRetryAlgorithm(), + new ExponentialRetryAlgorithm(settings, NanoClock.getDefaultClock())); + // Not retryable by configuration (isRetryable() == false), but caused by a channel refresh. + UnauthenticatedException rotationEx = rotationException(); + + TimedAttemptSettings rotationRetry = + retryAlgorithm.createNextAttempt( + context, rotationEx, null, retryAlgorithm.createFirstAttempt(context)); + assertNotNull(rotationRetry); + assertEquals(Duration.ZERO, rotationRetry.getRetryDelayDuration()); + assertTrue(retryAlgorithm.shouldRetry(context, rotationEx, null, rotationRetry)); + + TimedAttemptSettings afterSecondFailure = + retryAlgorithm.createNextAttempt(context, rotationEx, null, rotationRetry); + assertFalse(retryAlgorithm.shouldRetry(context, rotationEx, null, afterSecondFailure)); + } + + @Test + void testConfiguredUnauthenticated_afterRotationRetry_continuesWithNormalPolicy() { + ApiCallContext context = + mock(ApiCallContext.class, Mockito.withSettings().withoutAnnotations()); + when(context.getRetryableCodes()).thenReturn(null); + RetrySettings settings = + RetrySettings.newBuilder() + .setMaxAttempts(3) + .setInitialRetryDelayDuration(Duration.ofMillis(100)) + .setRetryDelayMultiplier(2.0) + .setMaxRetryDelayDuration(Duration.ofSeconds(1)) + .setTotalTimeoutDuration(Duration.ofMinutes(10)) + .build(); + RetryAlgorithm retryAlgorithm = + new RetryAlgorithm<>( + new ApiResultRetryAlgorithm(), + new ExponentialRetryAlgorithm(settings, NanoClock.getDefaultClock())); + UnauthenticatedException configuredEx = + new UnauthenticatedException( + "Expired cert", null, new FakeStatusCode(Code.UNAUTHENTICATED), /* retryable= */ true); + UnauthenticatedException rotationEx = configuredEx.withChannelRefreshed(); + + // A rotation failure gets the free zero-delay retry without consuming an attempt. + TimedAttemptSettings rotationRetry = + retryAlgorithm.createNextAttempt( + context, rotationEx, null, retryAlgorithm.createFirstAttempt(context)); + assertEquals(Duration.ZERO, rotationRetry.getRetryDelayDuration()); + assertEquals(0, rotationRetry.getAttemptCount()); + + // A later, unrelated failure on the same channel follows the configured policy with backoff. + TimedAttemptSettings normalRetry = + retryAlgorithm.createNextAttempt(context, configuredEx, null, rotationRetry); + assertNotNull(normalRetry); + assertEquals(1, normalRetry.getAttemptCount()); + assertTrue(normalRetry.getRetryDelayDuration().compareTo(Duration.ZERO) > 0); + assertTrue(retryAlgorithm.shouldRetry(context, configuredEx, null, normalRetry)); + } + + /** + * An UNAUTHENTICATED failure caused by a channel refresh, where UNAUTHENTICATED is not configured + * as retryable (the common case). + */ + private static UnauthenticatedException rotationException() { + return new UnauthenticatedException( + "Expired cert", null, new FakeStatusCode(Code.UNAUTHENTICATED), /* retryable= */ false) + .withChannelRefreshed(); + } } diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/AttemptCallableTest.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/AttemptCallableTest.java index 84c51da90707..6014f4e4ff7e 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/AttemptCallableTest.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/AttemptCallableTest.java @@ -30,6 +30,7 @@ package com.google.api.gax.rpc; import static com.google.common.truth.Truth.assertThat; +import static org.junit.jupiter.api.Assertions.assertThrows; import static org.mockito.ArgumentMatchers.any; import static org.mockito.Mockito.mock; import static org.mockito.Mockito.when; @@ -44,6 +45,8 @@ import com.google.api.gax.rpc.testing.FakeStatusCode; import com.google.api.gax.rpc.testing.FakeTransportChannel; import com.google.api.gax.tracing.ApiTracer; +import java.util.concurrent.ExecutionException; +import java.util.concurrent.atomic.AtomicInteger; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.extension.ExtendWith; @@ -177,7 +180,11 @@ void testUnauthenticatedExceptionReThrowPreservesContext() { assertThat(thrown).isInstanceOf(UnauthenticatedException.class); UnauthenticatedException rethrown = (UnauthenticatedException) thrown; - assertThat(rethrown.isRetryable()).isTrue(); + assertThat(rethrown.isChannelRefreshed()).isTrue(); + // The original retryable setting is preserved. + assertThat(rethrown.isRetryable()).isFalse(); + assertThat(rethrown.getMessage()).isEqualTo(originalEx.getMessage()); + assertThat(rethrown.getStatusCode()).isEqualTo(originalEx.getStatusCode()); assertThat(rethrown.getCause()).isEqualTo(originalEx.getCause()); assertThat(rethrown.getStackTrace()).isEqualTo(originalEx.getStackTrace()); assertThat(rethrown.getSuppressed().length).isEqualTo(1); @@ -185,7 +192,7 @@ void testUnauthenticatedExceptionReThrowPreservesContext() { } @Test - void testSiblingInFlightRequest_channelRotatedInFlight_markedRetryableWithoutDuplicateRefresh() { + void testSiblingInFlightRequest_channelRotatedInFlight_flaggedWithoutDuplicateRefresh() { FakeChannel innerChannel = new FakeChannel(); innerChannel.setGeneration(1); // Initially shouldRefresh is false because sibling request already completed the refresh @@ -227,14 +234,14 @@ void testSiblingInFlightRequest_channelRotatedInFlight_markedRetryableWithoutDup assertThat(thrown).isInstanceOf(UnauthenticatedException.class); UnauthenticatedException rethrown = (UnauthenticatedException) thrown; - // Sibling request should be marked retryable to run on the new channel - assertThat(rethrown.isRetryable()).isTrue(); + // Sibling request should be flagged for a retry on the new channel + assertThat(rethrown.isChannelRefreshed()).isTrue(); // But should NOT have triggered a second refresh call assertThat(innerChannel.getRefreshCount()).isEqualTo(0); } @Test - void testPermanentUnauthenticatedFailure_sameGeneration_notMarkedRetryable() { + void testPermanentUnauthenticatedFailure_sameGeneration_notFlagged() { FakeChannel innerChannel = new FakeChannel(); innerChannel.setGeneration(1); innerChannel.setShouldRefresh(false); @@ -271,11 +278,12 @@ void testPermanentUnauthenticatedFailure_sameGeneration_notMarkedRetryable() { UnauthenticatedException rethrown = (UnauthenticatedException) thrown; // Genuine permanent error on same generation is NOT retryable assertThat(rethrown.isRetryable()).isFalse(); + assertThat(rethrown.isChannelRefreshed()).isFalse(); assertThat(innerChannel.getRefreshCount()).isEqualTo(0); } @Test - void testRefreshThrowsException_notMarkedRetryableWhenGenerationUnchanged() { + void testRefreshThrowsException_notFlaggedWhenGenerationUnchanged() { FakeChannel fakeChannel = new FakeChannel() { @Override @@ -320,10 +328,11 @@ public void refresh() { assertThat(thrown).isInstanceOf(UnauthenticatedException.class); UnauthenticatedException rethrown = (UnauthenticatedException) thrown; assertThat(rethrown.isRetryable()).isFalse(); + assertThat(rethrown.isChannelRefreshed()).isFalse(); } @Test - void testRefreshReturnsWithoutAdvancingGeneration_notMarkedRetryable() { + void testRefreshReturnsWithoutAdvancingGeneration_notFlagged() { FakeChannel fakeChannel = new FakeChannel() { @Override @@ -369,10 +378,11 @@ public void refresh() { assertThat(thrown).isInstanceOf(UnauthenticatedException.class); UnauthenticatedException rethrown = (UnauthenticatedException) thrown; assertThat(rethrown.isRetryable()).isFalse(); + assertThat(rethrown.isChannelRefreshed()).isFalse(); } @Test - void testRefreshThrowsException_markedRetryableIfConcurrentThreadAdvancedGeneration() { + void testRefreshThrowsException_flaggedIfConcurrentThreadAdvancedGeneration() { FakeChannel fakeChannel = new FakeChannel() { @Override @@ -418,6 +428,97 @@ public void refresh() { assertThat(thrown).isInstanceOf(UnauthenticatedException.class); UnauthenticatedException rethrown = (UnauthenticatedException) thrown; + assertThat(rethrown.isChannelRefreshed()).isTrue(); + } + + @Test + void testGenerationAlreadyAdvanced_skipsShouldRefresh() { + AtomicInteger shouldRefreshCalls = new AtomicInteger(); + FakeChannel fakeChannel = + new FakeChannel() { + @Override + public boolean shouldRefresh() { + shouldRefreshCalls.incrementAndGet(); + return true; + } + }; + UnauthenticatedException originalEx = + new UnauthenticatedException( + "Expired cert", null, FakeStatusCode.of(StatusCode.Code.UNAUTHENTICATED), false); + SettableApiFuture failedFuture = SettableApiFuture.create(); + failedFuture.setException(originalEx); + when(mockInnerCallable.futureCall(Mockito.anyString(), Mockito.any())) + .thenAnswer( + invocation -> { + // Another request rotated the channel while this one was in flight. + fakeChannel.setGeneration(fakeChannel.getGeneration() + 1); + return failedFuture; + }); + + UnauthenticatedException rethrown = callAndGetUnauthenticated(fakeChannel); + + assertThat(rethrown.isChannelRefreshed()).isTrue(); + // The certificate on disk is not checked again, and no second refresh happens. + assertThat(shouldRefreshCalls.get()).isEqualTo(0); + assertThat(fakeChannel.getRefreshCount()).isEqualTo(0); + } + + @Test + void testShouldRefreshThrows_originalUnauthenticatedPreserved() { + FakeChannel fakeChannel = + new FakeChannel() { + @Override + public boolean shouldRefresh() { + throw new IllegalStateException("Unable to read certificate"); + } + }; + UnauthenticatedException originalEx = + new UnauthenticatedException( + "Expired cert", null, FakeStatusCode.of(StatusCode.Code.UNAUTHENTICATED), false); + SettableApiFuture failedFuture = SettableApiFuture.create(); + failedFuture.setException(originalEx); + when(mockInnerCallable.futureCall(Mockito.anyString(), Mockito.any())).thenReturn(failedFuture); + + UnauthenticatedException rethrown = callAndGetUnauthenticated(fakeChannel); + + assertThat(rethrown).isSameInstanceAs(originalEx); + assertThat(rethrown.isChannelRefreshed()).isFalse(); + assertThat(fakeChannel.getRefreshCount()).isEqualTo(0); + } + + @Test + void testConfiguredRetryableUnauthenticated_sameGeneration_notFlagged() { + FakeChannel fakeChannel = new FakeChannel().setShouldRefresh(false); + UnauthenticatedException originalEx = + new UnauthenticatedException( + "Token expired", null, FakeStatusCode.of(StatusCode.Code.UNAUTHENTICATED), true); + SettableApiFuture failedFuture = SettableApiFuture.create(); + failedFuture.setException(originalEx); + when(mockInnerCallable.futureCall(Mockito.anyString(), Mockito.any())).thenReturn(failedFuture); + + UnauthenticatedException rethrown = callAndGetUnauthenticated(fakeChannel); + + // Left to the normal retry policy: still retryable, but not a rotation retry. + assertThat(rethrown).isSameInstanceAs(originalEx); assertThat(rethrown.isRetryable()).isTrue(); + assertThat(rethrown.isChannelRefreshed()).isFalse(); + } + + private UnauthenticatedException callAndGetUnauthenticated(FakeChannel fakeChannel) { + ApiCallContext callContext = + FakeCallContext.createDefault() + .withTransportChannel(FakeTransportChannel.create(fakeChannel)); + AttemptCallable callable = + new AttemptCallable<>(mockInnerCallable, "fake-request", callContext); + callable.setExternalFuture(mockExternalFuture); + + callable.call(); + + ArgumentCaptor futureCaptor = ArgumentCaptor.forClass(ApiFuture.class); + Mockito.verify(mockExternalFuture, Mockito.times(2)).setAttemptFuture(futureCaptor.capture()); + ExecutionException e = + assertThrows(ExecutionException.class, () -> futureCaptor.getValue().get()); + assertThat(e.getCause()).isInstanceOf(UnauthenticatedException.class); + return (UnauthenticatedException) e.getCause(); } } diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ServerStreamingAttemptCallableTest.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ServerStreamingAttemptCallableTest.java index 66fc73a7d2bf..fd9eca2a4c05 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ServerStreamingAttemptCallableTest.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ServerStreamingAttemptCallableTest.java @@ -290,6 +290,8 @@ void testUnauthenticatedRefresh() { Truth.assertThat(((ServerStreamingAttemptException) outerError).canResume()).isTrue(); Truth.assertThat(outerError.getCause()).isInstanceOf(UnauthenticatedException.class); Truth.assertThat(((UnauthenticatedException) outerError.getCause()).isRetryable()).isFalse(); + Truth.assertThat(((UnauthenticatedException) outerError.getCause()).isChannelRefreshed()) + .isFalse(); } @Test @@ -336,7 +338,9 @@ void testUnauthenticatedRefreshWithGenerationAdvanceRetries() { Truth.assertThat(outerError).isInstanceOf(ServerStreamingAttemptException.class); Truth.assertThat(((ServerStreamingAttemptException) outerError).canResume()).isTrue(); Truth.assertThat(outerError.getCause()).isInstanceOf(UnauthenticatedException.class); - Truth.assertThat(((UnauthenticatedException) outerError.getCause()).isRetryable()).isTrue(); + Truth.assertThat(((UnauthenticatedException) outerError.getCause()).isChannelRefreshed()) + .isTrue(); + Truth.assertThat(((UnauthenticatedException) outerError.getCause()).isRetryable()).isFalse(); Truth.assertThat(outerError.getCause().getStackTrace()).isEqualTo(initialError.getStackTrace()); // Verify retry call resumes stream @@ -391,7 +395,8 @@ void testUnauthenticatedRefreshWithNonResumableStreamDoesNotRetry() { Truth.assertThat(outerError).isInstanceOf(ServerStreamingAttemptException.class); ServerStreamingAttemptException attemptEx = (ServerStreamingAttemptException) outerError; Truth.assertThat(attemptEx.canResume()).isFalse(); - Truth.assertThat(((UnauthenticatedException) attemptEx.getCause()).isRetryable()).isTrue(); + Truth.assertThat(((UnauthenticatedException) attemptEx.getCause()).isChannelRefreshed()) + .isTrue(); Truth.assertThat( new com.google.api.gax.retrying.StreamingRetryAlgorithm<>( new ApiResultRetryAlgorithm<>(), @@ -594,7 +599,7 @@ public String processResponse(String response) { } @Test - void testUnauthenticatedException_whenChannelRefreshes_setsRetryableTrue() throws Exception { + void testUnauthenticatedException_whenChannelRefreshes_flagsChannelRefreshed() throws Exception { FakeChannel fakeChannel = new FakeChannel(); fakeChannel.setShouldRefresh(true); ApiCallContext context = @@ -617,12 +622,12 @@ void testUnauthenticatedException_whenChannelRefreshes_setsRetryableTrue() throw Truth.assertThat(ex.getCause()).isInstanceOf(ServerStreamingAttemptException.class); Throwable cause = ex.getCause().getCause(); Truth.assertThat(cause).isInstanceOf(UnauthenticatedException.class); - Truth.assertThat(((UnauthenticatedException) cause).isRetryable()).isTrue(); + Truth.assertThat(((UnauthenticatedException) cause).isChannelRefreshed()).isTrue(); Truth.assertThat(fakeChannel.getRefreshCount()).isEqualTo(1); } @Test - void testUnauthenticatedException_whenChannelRefreshFails_remainsNonRetryable() throws Exception { + void testUnauthenticatedException_whenChannelRefreshFails_notFlagged() throws Exception { FakeChannel fakeChannel = new FakeChannel() { @Override @@ -652,6 +657,75 @@ public void refresh() { Throwable cause = ex.getCause().getCause(); Truth.assertThat(cause).isInstanceOf(UnauthenticatedException.class); Truth.assertThat(((UnauthenticatedException) cause).isRetryable()).isFalse(); + Truth.assertThat(((UnauthenticatedException) cause).isChannelRefreshed()).isFalse(); + } + + @Test + void testUnauthenticated_generationAlreadyAdvanced_skipsShouldRefresh() throws Exception { + TransportChannel transportChannel = Mockito.mock(TransportChannel.class); + java.util.concurrent.atomic.AtomicLong generation = + new java.util.concurrent.atomic.AtomicLong(0); + Mockito.when(transportChannel.getGeneration()).thenAnswer(inv -> generation.get()); + + ApiCallContext context = Mockito.mock(ApiCallContext.class); + Mockito.when(context.getTransportChannel()).thenReturn(transportChannel); + Mockito.when(context.getTracer()).thenReturn(BaseApiTracer.getInstance()); + Mockito.when(context.getTimeoutDuration()).thenReturn(java.time.Duration.ofHours(5)); + + resumptionStrategy = new MyStreamResumptionStrategy(); + ServerStreamingAttemptCallable callable = createCallable(context); + callable.start(); + + MockServerStreamingCall call = innerCallable.popLastCall(); + // Another request rotated the channel while this stream was open. + generation.incrementAndGet(); + call.getController() + .getObserver() + .onError( + new UnauthenticatedException( + "cert expired", null, new FakeStatusCode(Code.UNAUTHENTICATED), false)); + + ExecutionException ex = + assertThrows( + ExecutionException.class, + () -> fakeRetryingFuture.getAttemptResult().get(1, TimeUnit.SECONDS)); + Throwable cause = ex.getCause().getCause(); + Truth.assertThat(cause).isInstanceOf(UnauthenticatedException.class); + Truth.assertThat(((UnauthenticatedException) cause).isChannelRefreshed()).isTrue(); + Mockito.verify(transportChannel, Mockito.never()).shouldRefresh(); + Mockito.verify(transportChannel, Mockito.never()).refresh(); + } + + @Test + void testUnauthenticated_shouldRefreshThrows_originalErrorPreserved() throws Exception { + FakeChannel fakeChannel = + new FakeChannel() { + @Override + public boolean shouldRefresh() { + throw new IllegalStateException("Unable to read certificate"); + } + }; + ApiCallContext context = + FakeCallContext.createDefault() + .withTransportChannel(FakeTransportChannel.create(fakeChannel)); + + ServerStreamingAttemptCallable callable = createCallable(context); + callable.start(); + + MockServerStreamingCall call = innerCallable.popLastCall(); + UnauthenticatedException unauthEx = + new UnauthenticatedException( + "cert expired", null, new FakeStatusCode(Code.UNAUTHENTICATED), false); + call.getController().getObserver().onError(unauthEx); + + ExecutionException ex = + assertThrows( + ExecutionException.class, + () -> fakeRetryingFuture.getAttemptResult().get(1, TimeUnit.SECONDS)); + Truth.assertThat(ex.getCause()).isInstanceOf(ServerStreamingAttemptException.class); + Truth.assertThat(ex.getCause().getCause()).isSameInstanceAs(unauthEx); + Truth.assertThat(unauthEx.isChannelRefreshed()).isFalse(); + Truth.assertThat(fakeChannel.getRefreshCount()).isEqualTo(0); } static class MyStreamResumptionStrategy implements StreamResumptionStrategy { diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/UnauthenticatedExceptionTest.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/UnauthenticatedExceptionTest.java new file mode 100644 index 000000000000..083de5286277 --- /dev/null +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/UnauthenticatedExceptionTest.java @@ -0,0 +1,183 @@ +/* + * Copyright 2026 Google LLC + * + * Redistribution and use in source and binary forms, with or without + * modification, are permitted provided that the following conditions are + * met: + * + * * Redistributions of source code must retain the above copyright + * notice, this list of conditions and the following disclaimer. + * * Redistributions in binary form must reproduce the above + * copyright notice, this list of conditions and the following disclaimer + * in the documentation and/or other materials provided with the + * distribution. + * * Neither the name of Google LLC nor the names of its + * contributors may be used to endorse or promote products derived from + * this software without specific prior written permission. + * + * THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS + * "AS IS" AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT + * LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR + * A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT + * OWNER OR CONTRIBUTORS BE LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, + * SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT + * LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, + * DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY + * THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT + * (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE + * OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + */ +package com.google.api.gax.rpc; + +import static com.google.common.truth.Truth.assertThat; + +import com.google.api.gax.rpc.testing.FakeStatusCode; +import java.io.ByteArrayInputStream; +import java.io.ByteArrayOutputStream; +import java.io.ObjectInputStream; +import java.io.ObjectOutputStream; +import java.io.ObjectStreamClass; +import java.io.Serializable; +import java.util.Base64; +import java.util.Collections; +import org.junit.jupiter.api.Test; + +class UnauthenticatedExceptionTest { + + /** + * An {@link UnauthenticatedException} with message "serialized by released gax", no cause, an + * empty stack trace, retryable {@code true} and a {@link SerializableStatusCode} of + * UNAUTHENTICATED, Java-serialized with the released gax 2.87.0 jar. + */ + private static final String SERIALIZED_BY_GAX_2_87_0 = + "rO0ABXNyAC9jb20uZ29vZ2xlLmFwaS5nYXgucnBjLlVuYXV0aGVudGljYXRlZEV4Y2VwdGlvbmC+YDhKxcJlAgAAeHIAI2NvbS5nb29nbGUuYXBpLmdheC5ycGMuQXBpRXhjZXB0aW9uw0h4sCzSVFQCAANaAAlyZXRyeWFibGVMAAxlcnJvckRldGFpbHN0ACVMY29tL2dvb2dsZS9hcGkvZ2F4L3JwYy9FcnJvckRldGFpbHM7TAAKc3RhdHVzQ29kZXQAI0xjb20vZ29vZ2xlL2FwaS9nYXgvcnBjL1N0YXR1c0NvZGU7eHIAGmphdmEubGFuZy5SdW50aW1lRXhjZXB0aW9unl8GRwo0g+UCAAB4cgATamF2YS5sYW5nLkV4Y2VwdGlvbtD9Hz4aOxzEAgAAeHIAE2phdmEubGFuZy5UaHJvd2FibGXVxjUnOXe4ywMABEwABWNhdXNldAAVTGphdmEvbGFuZy9UaHJvd2FibGU7TAANZGV0YWlsTWVzc2FnZXQAEkxqYXZhL2xhbmcvU3RyaW5nO1sACnN0YWNrVHJhY2V0AB5bTGphdmEvbGFuZy9TdGFja1RyYWNlRWxlbWVudDtMABRzdXBwcmVzc2VkRXhjZXB0aW9uc3QAEExqYXZhL3V0aWwvTGlzdDt4cHB0ABpzZXJpYWxpemVkIGJ5IHJlbGVhc2VkIGdheHVyAB5bTGphdmEubGFuZy5TdGFja1RyYWNlRWxlbWVudDsCRio8PP0iOQIAAHhwAAAAAHNyAB9qYXZhLnV0aWwuQ29sbGVjdGlvbnMkRW1wdHlMaXN0ergXtDynnt4CAAB4cHgBcHNyAEpjb20uZ29vZ2xlLmFwaS5nYXgucnBjLlVuYXV0aGVudGljYXRlZEV4Y2VwdGlvblRlc3QkU2VyaWFsaXphYmxlU3RhdHVzQ29kZQAAAAAAAAABAgABTAAEY29kZXQAKExjb20vZ29vZ2xlL2FwaS9nYXgvcnBjL1N0YXR1c0NvZGUkQ29kZTt4cH5yACZjb20uZ29vZ2xlLmFwaS5nYXgucnBjLlN0YXR1c0NvZGUkQ29kZQAAAAAAAAAAEgAAeHIADmphdmEubGFuZy5FbnVtAAAAAAAAAAASAAB4cHQAD1VOQVVUSEVOVElDQVRFRA=="; + + /** A serializable {@link StatusCode}, since the transport implementations are not. */ + static final class SerializableStatusCode implements StatusCode, Serializable { + private static final long serialVersionUID = 1L; + private final Code code; + + SerializableStatusCode(Code code) { + this.code = code; + } + + @Override + public Code getCode() { + return code; + } + + @Override + public Object getTransportCode() { + return code.name(); + } + } + + @Test + void publicConstructors_areNotChannelRefreshed() { + StatusCode statusCode = FakeStatusCode.of(StatusCode.Code.UNAUTHENTICATED); + ErrorDetails errorDetails = + ErrorDetails.builder().setRawErrorMessages(Collections.emptyList()).build(); + + assertThat(new UnauthenticatedException(null, statusCode, true).isChannelRefreshed()).isFalse(); + assertThat(new UnauthenticatedException("msg", null, statusCode, true).isChannelRefreshed()) + .isFalse(); + assertThat( + new UnauthenticatedException(null, statusCode, true, errorDetails).isChannelRefreshed()) + .isFalse(); + assertThat( + new UnauthenticatedException("msg", null, statusCode, true, errorDetails) + .isChannelRefreshed()) + .isFalse(); + } + + @Test + void withChannelRefreshed_preservesFields() { + ErrorDetails errorDetails = + ErrorDetails.builder().setRawErrorMessages(Collections.emptyList()).build(); + IllegalStateException cause = new IllegalStateException("root cause"); + StatusCode statusCode = FakeStatusCode.of(StatusCode.Code.UNAUTHENTICATED); + UnauthenticatedException original = + new UnauthenticatedException("Expired cert", cause, statusCode, false, errorDetails); + original.setStackTrace( + new StackTraceElement[] {new StackTraceElement("foo", "bar", "Baz.java", 123)}); + RuntimeException suppressed = new RuntimeException("suppressed"); + original.addSuppressed(suppressed); + + UnauthenticatedException refreshed = original.withChannelRefreshed(); + + assertThat(refreshed).isNotSameInstanceAs(original); + assertThat(refreshed.isChannelRefreshed()).isTrue(); + assertThat(original.isChannelRefreshed()).isFalse(); + assertThat(refreshed.getMessage()).isEqualTo(original.getMessage()); + assertThat(refreshed.getCause()).isSameInstanceAs(cause); + assertThat(refreshed.getStatusCode()).isSameInstanceAs(statusCode); + assertThat(refreshed.getErrorDetails()).isSameInstanceAs(errorDetails); + assertThat(refreshed.getStackTrace()).isEqualTo(original.getStackTrace()); + assertThat(refreshed.getSuppressed()).asList().containsExactly(suppressed); + } + + @Test + void withChannelRefreshed_keepsIsRetryable() { + StatusCode statusCode = FakeStatusCode.of(StatusCode.Code.UNAUTHENTICATED); + + assertThat( + new UnauthenticatedException("msg", null, statusCode, false) + .withChannelRefreshed() + .isRetryable()) + .isFalse(); + assertThat( + new UnauthenticatedException("msg", null, statusCode, true) + .withChannelRefreshed() + .isRetryable()) + .isTrue(); + } + + @Test + void serialVersionUID_matchesReleasedValue() { + // gax 2.83.0 to 2.87.0 did not declare a serialVersionUID; this is the value the JVM computed + // for them. Keeping it avoids InvalidClassException across versions. + assertThat(ObjectStreamClass.lookup(UnauthenticatedException.class).getSerialVersionUID()) + .isEqualTo(6971115068105015909L); + } + + @Test + void deserializesExceptionSerializedByReleasedGax() throws Exception { + Object deserialized = deserialize(Base64.getDecoder().decode(SERIALIZED_BY_GAX_2_87_0)); + + assertThat(deserialized).isInstanceOf(UnauthenticatedException.class); + UnauthenticatedException ex = (UnauthenticatedException) deserialized; + assertThat(ex.getMessage()).isEqualTo("serialized by released gax"); + assertThat(ex.getStatusCode().getCode()).isEqualTo(StatusCode.Code.UNAUTHENTICATED); + assertThat(ex.isRetryable()).isTrue(); + assertThat(ex.isChannelRefreshed()).isFalse(); + } + + @Test + void withChannelRefreshed_flagNotSerialized() throws Exception { + UnauthenticatedException refreshed = + new UnauthenticatedException( + "msg", null, new SerializableStatusCode(StatusCode.Code.UNAUTHENTICATED), false) + .withChannelRefreshed(); + + UnauthenticatedException roundTripped = + (UnauthenticatedException) deserialize(serialize(refreshed)); + + assertThat(roundTripped.getMessage()).isEqualTo("msg"); + assertThat(roundTripped.isRetryable()).isFalse(); + assertThat(roundTripped.isChannelRefreshed()).isFalse(); + } + + private static byte[] serialize(Object object) throws Exception { + ByteArrayOutputStream bytes = new ByteArrayOutputStream(); + try (ObjectOutputStream out = new ObjectOutputStream(bytes)) { + out.writeObject(object); + } + return bytes.toByteArray(); + } + + private static Object deserialize(byte[] bytes) throws Exception { + try (ObjectInputStream in = new ObjectInputStream(new ByteArrayInputStream(bytes))) { + return in.readObject(); + } + } +} From ff9f3357fafa97081244e29c74133b0a7787196e Mon Sep 17 00:00:00 2001 From: Matt Castelaz Date: Sat, 3 Oct 2026 04:42:25 +0000 Subject: [PATCH 29/29] test(gax): name tests after their assertions and tighten weak assertions --- .../google/api/gax/grpc/ChannelPoolTest.java | 29 ++++++++++++++----- .../InstantiatingGrpcChannelProviderTest.java | 2 +- .../gax/httpjson/HttpJsonCallContextTest.java | 3 +- ...tantiatingHttpJsonChannelProviderTest.java | 18 ++++++++++-- .../RefreshingHttpJsonChannelTest.java | 12 ++++++-- .../gax/rpc/ApiResultRetryAlgorithmTest.java | 4 +-- .../api/gax/rpc/AttemptCallableTest.java | 9 ++++-- .../ServerStreamingAttemptCallableTest.java | 22 ++++++++++++-- 8 files changed, 76 insertions(+), 23 deletions(-) diff --git a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java index e7e386e2304c..7821dccb2b7d 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/ChannelPoolTest.java @@ -461,7 +461,7 @@ void testCancelBeforeStartReleasesChannelEntry() throws IOException { } @Test - void channelReactiveMTlsRefreshShouldConditionallySwapChannels() + void channelReactiveMTlsRefresh_swapsChannelsOnlyWhenCertChanges() throws IOException, InterruptedException { ManagedChannel underlyingChannel1 = Mockito.mock(ManagedChannel.class); ManagedChannel underlyingChannel2 = Mockito.mock(ManagedChannel.class); @@ -794,11 +794,19 @@ void refill_channelCreationFailure_isLoggedAndDoesNotThrow() throws IOException // Checked and unchecked failures are both handled by expand(); the pool stays usable at its // reduced size. - refillTask.getValue().run(); - assertThat(pool.entries.get()).hasSize(1); + FakeLogHandler logHandler = new FakeLogHandler(); + ChannelPool.LOG.addHandler(logHandler); + try { + refillTask.getValue().run(); + assertThat(pool.entries.get()).hasSize(1); - refillTask.getValue().run(); - assertThat(pool.entries.get()).hasSize(1); + refillTask.getValue().run(); + assertThat(pool.entries.get()).hasSize(1); + } finally { + ChannelPool.LOG.removeHandler(logHandler); + } + assertThat(logHandler.getAllMessages()) + .containsExactly("Failed to add channel", "Failed to add channel"); refillTask.getValue().run(); assertThat(pool.entries.get()).hasSize(2); @@ -1012,8 +1020,7 @@ void preemptiveRefresh_pickingUpRotatedCert_incrementsGeneration() throws IOExce } @Test - void preemptiveRefresh_partialFailureDuringRotation_doesNotIncrementOrMarkRefreshed() - throws IOException { + void preemptiveRefresh_partialFailureDuringRotation_leavesRotationPending() throws IOException { ManagedChannel initial1 = Mockito.mock(ManagedChannel.class); ManagedChannel initial2 = Mockito.mock(ManagedChannel.class); ManagedChannel rotated1 = Mockito.mock(ManagedChannel.class); @@ -1032,12 +1039,14 @@ void preemptiveRefresh_partialFailureDuringRotation_doesNotIncrementOrMarkRefres Mockito.verify(initial1).shutdown(); Mockito.verify(initial2, Mockito.never()).shutdown(); assertThat(pool.getGeneration()).isEqualTo(0); + // The new certificate is not recorded as active, so the rotation stays pending and the next + // UNAUTHENTICATED failure triggers a reactive refresh that completes the switch. pool.invalidateDiskFingerprintCache(); assertThat(pool.shouldRefresh()).isTrue(); } @Test - void refresh_whenAlreadyActive_doesNotIncrementGeneration() throws IOException { + void refresh_whenCertUnchanged_noOpsAndDoesNotIncrementGeneration() throws IOException { ChannelFactory channelFactory = mockChannelFactory(); Mockito.when(channelFactory.createSingleChannel()) .thenReturn(Mockito.mock(ManagedChannel.class), Mockito.mock(ManagedChannel.class)); @@ -1861,6 +1870,10 @@ public void sendMessage(Color message) {} executor.shutdownNow(); } + // Every start/cancel race must leave the entry with no outstanding RPCs: a leak would leave it + // positive and a double release would make it negative. + assertThat(pool.entries.get().get(0).outstandingRpcs.get()).isEqualTo(0); + // Rotate pool: initial channel must shut down cleanly, proving outstandingRpcs == 0 (no leaks // or negative counts) pool.refresh(); diff --git a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/InstantiatingGrpcChannelProviderTest.java b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/InstantiatingGrpcChannelProviderTest.java index d083c01ad845..d766b288d77d 100644 --- a/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/InstantiatingGrpcChannelProviderTest.java +++ b/sdk-platform-java/gax-java/gax-grpc/src/test/java/com/google/api/gax/grpc/InstantiatingGrpcChannelProviderTest.java @@ -1424,7 +1424,7 @@ void createChannel_whenMtlsActive_passesWorkloadCertPathToChannelPool() throws E } @Test - void createChannelBuilder_whenMtlsActiveAndCredentialsNull_throwsIOException() { + void createChannelBuilder_whenMtlsActiveAndKeyStoreNull_throwsIOException() { CertificateBasedAccess mtlsCertificateBasedAccess = mock(CertificateBasedAccess.class, Mockito.withSettings().withoutAnnotations()); Mockito.when(mtlsCertificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/HttpJsonCallContextTest.java b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/HttpJsonCallContextTest.java index ce7b6682e586..39c82205c294 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/HttpJsonCallContextTest.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/HttpJsonCallContextTest.java @@ -373,7 +373,8 @@ void testWithChannelClearsStaleTransportChannel() { Truth.assertThat(nullChannelContext.getTransportChannel()).isNull(); Truth.assertThat(context.withChannel(channel2).getTransportChannel()).isNull(); - // Merging a cleared context into default context preserves default context's transportChannel + // Merging a context with a cleared channel into the original context keeps the original + // channel and transportChannel HttpJsonCallContext mergedWithNullChannel = context.merge(nullChannelContext); Truth.assertThat(mergedWithNullChannel.getChannel()).isSameInstanceAs(channel1); Truth.assertThat(mergedWithNullChannel.getTransportChannel()) diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProviderTest.java b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProviderTest.java index b63ce6ba04be..ee9c476adf18 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProviderTest.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/InstantiatingHttpJsonChannelProviderTest.java @@ -246,7 +246,16 @@ void channelCreation_withWorkloadCertPath_wrapsWithRefreshingHttpJsonChannel() @Test void channelCreation_withCustomHttpTransport_ignoresWorkloadCertPathAndDoesNotWrap() - throws IOException { + throws IOException, GeneralSecurityException { + // mTLS is otherwise fully configured, so only the custom transport prevents wrapping. Lenient + // because a custom transport is expected to skip these lookups. + Mockito.lenient().when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); + Mockito.lenient() + .when(certificateBasedAccess.getWorkloadCertPath()) + .thenReturn("fake/cert/path.json"); + com.google.auth.mtls.MtlsProvider mtlsProvider = + new com.google.api.gax.rpc.testing.FakeMtlsProvider( + com.google.api.gax.rpc.testing.FakeMtlsProvider.createTestMtlsKeyStore(), "", false); com.google.api.client.http.HttpTransport mockHttpTransport = org.mockito.Mockito.mock(com.google.api.client.http.HttpTransport.class); @@ -254,6 +263,7 @@ void channelCreation_withCustomHttpTransport_ignoresWorkloadCertPathAndDoesNotWr InstantiatingHttpJsonChannelProvider.newBuilder() .setEndpoint(DEFAULT_ENDPOINT) .setHttpTransport(mockHttpTransport) + .setMtlsProvider(mtlsProvider) .setCertificateBasedAccess(certificateBasedAccess) .build(); provider = (InstantiatingHttpJsonChannelProvider) provider.withHeaders(DEFAULT_HEADER_MAP); @@ -429,7 +439,8 @@ void getTransportChannel_whenRotationEnabledAndKeyStoreNull_throwsIOException() } @Test - void getTransportChannel_whenKeyStoreUninitialized_causeIsSecurityException() throws Exception { + void getTransportChannel_whenKeyStoreUninitialized_wrapsGeneralSecurityException() + throws Exception { Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); Mockito.when(certificateBasedAccess.getWorkloadCertPath()).thenReturn("fake/cert/path.json"); // A KeyStore that was never loaded makes transport creation fail with a KeyStoreException. @@ -480,7 +491,7 @@ void getTransportChannel_whenMtlsKeyStoreThrowsRuntimeException_propagatesUnwrap } @Test - void createHttpTransport_withMtlsAndConscrypt_configuresSecurityProvider() + void createHttpTransport_withMtlsKeyStore_returnsMtlsTransport() throws IOException, GeneralSecurityException { Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true); com.google.auth.mtls.MtlsProvider provider = @@ -497,6 +508,7 @@ void createHttpTransport_withMtlsAndConscrypt_configuresSecurityProvider() com.google.api.client.http.HttpTransport transport = channelProvider.createHttpTransport(); assertThat(transport).isNotNull(); assertThat(transport).isInstanceOf(com.google.api.client.http.javanet.NetHttpTransport.class); + assertThat(((com.google.api.client.http.javanet.NetHttpTransport) transport).isMtls()).isTrue(); } @Test diff --git a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java index bf7b33da84ba..de274b5fbdc7 100644 --- a/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java +++ b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java @@ -217,7 +217,7 @@ private void rotateCertificate(RefreshingHttpJsonChannel channel) { } @Test - void testShouldRefreshNullCertPath() { + void testShouldRefreshFalseWhenCertPathNull() { testCertPath = null; RefreshingHttpJsonChannel channel = createTestChannel(); assertFalse(channel.shouldRefresh()); @@ -503,6 +503,7 @@ void testConcurrentNewCallDuringRefresh() throws InterruptedException { int threadCount = 10; java.util.concurrent.ExecutorService executorService = java.util.concurrent.Executors.newFixedThreadPool(threadCount); + java.util.concurrent.CountDownLatch start = new java.util.concurrent.CountDownLatch(1); java.util.concurrent.CountDownLatch latch = new java.util.concurrent.CountDownLatch(threadCount); AtomicInteger successCount = new AtomicInteger(0); @@ -511,8 +512,11 @@ void testConcurrentNewCallDuringRefresh() throws InterruptedException { executorService.submit( () -> { try { + start.await(); channel.newCall(null, null); successCount.incrementAndGet(); + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); } finally { latch.countDown(); } @@ -520,12 +524,16 @@ void testConcurrentNewCallDuringRefresh() throws InterruptedException { } rotateCertificate(channel); + // Release the workers just before refreshing so their calls overlap the transport swap. + start.countDown(); channel.refresh(); - latch.await(5, TimeUnit.SECONDS); + assertTrue(latch.await(5, TimeUnit.SECONDS)); executorService.shutdown(); + assertTrue(executorService.awaitTermination(5, TimeUnit.SECONDS)); assertEquals(threadCount, successCount.get()); + assertEquals(1, channel.getGeneration()); } @Test diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ApiResultRetryAlgorithmTest.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ApiResultRetryAlgorithmTest.java index 200aaddf7730..1e730d2db1f3 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ApiResultRetryAlgorithmTest.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ApiResultRetryAlgorithmTest.java @@ -282,7 +282,7 @@ void testSecondRotationFailureStopsWithTotalTimeoutAndNoMaxAttempts() { } @Test - void testSecondRotationFailureDoesNotConsumeNormalRetryBudget() { + void testSecondRotationFailureStopsDespiteRemainingMaxAttempts() { ApiCallContext context = mock(ApiCallContext.class, Mockito.withSettings().withoutAnnotations()); when(context.getRetryableCodes()).thenReturn(Sets.newHashSet(Code.UNAVAILABLE)); @@ -317,7 +317,7 @@ void testSecondRotationFailureDoesNotConsumeNormalRetryBudget() { } @Test - void testStreamRotationRetryIsAvailableAgainAfterProgress() { + void testStreamRotationRetryAfterProgressNotBlockedByEarlierRetry() { ApiCallContext context = mock(ApiCallContext.class, Mockito.withSettings().withoutAnnotations()); when(context.getRetryableCodes()).thenReturn(Sets.newHashSet(Code.UNAVAILABLE)); diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/AttemptCallableTest.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/AttemptCallableTest.java index 6014f4e4ff7e..a1e148c241b6 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/AttemptCallableTest.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/AttemptCallableTest.java @@ -142,7 +142,7 @@ void testRpcTimeoutIsNotErased() { } @Test - void testUnauthenticatedExceptionReThrowPreservesContext() { + void testRefreshedUnauthenticated_flaggedAndPreservesContext() { FakeTransportChannel transportChannel = FakeTransportChannel.create(new FakeChannel()).setShouldRefresh(true); ApiCallContext callContext = @@ -276,7 +276,7 @@ void testPermanentUnauthenticatedFailure_sameGeneration_notFlagged() { assertThat(thrown).isInstanceOf(UnauthenticatedException.class); UnauthenticatedException rethrown = (UnauthenticatedException) thrown; - // Genuine permanent error on same generation is NOT retryable + // Genuine permanent error on same generation is not flagged for a rotation retry assertThat(rethrown.isRetryable()).isFalse(); assertThat(rethrown.isChannelRefreshed()).isFalse(); assertThat(innerChannel.getRefreshCount()).isEqualTo(0); @@ -284,6 +284,7 @@ void testPermanentUnauthenticatedFailure_sameGeneration_notFlagged() { @Test void testRefreshThrowsException_notFlaggedWhenGenerationUnchanged() { + AtomicInteger refreshCalls = new AtomicInteger(); FakeChannel fakeChannel = new FakeChannel() { @Override @@ -293,6 +294,7 @@ public boolean shouldRefresh() { @Override public void refresh() { + refreshCalls.incrementAndGet(); throw new RuntimeException("Refresh error"); } }; @@ -329,6 +331,7 @@ public void refresh() { UnauthenticatedException rethrown = (UnauthenticatedException) thrown; assertThat(rethrown.isRetryable()).isFalse(); assertThat(rethrown.isChannelRefreshed()).isFalse(); + assertThat(refreshCalls.get()).isEqualTo(1); } @Test @@ -432,7 +435,7 @@ public void refresh() { } @Test - void testGenerationAlreadyAdvanced_skipsShouldRefresh() { + void testGenerationAlreadyAdvanced_flaggedWithoutCheckingShouldRefresh() { AtomicInteger shouldRefreshCalls = new AtomicInteger(); FakeChannel fakeChannel = new FakeChannel() { diff --git a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ServerStreamingAttemptCallableTest.java b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ServerStreamingAttemptCallableTest.java index fd9eca2a4c05..5d93f3894874 100644 --- a/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ServerStreamingAttemptCallableTest.java +++ b/sdk-platform-java/gax-java/gax/src/test/java/com/google/api/gax/rpc/ServerStreamingAttemptCallableTest.java @@ -254,7 +254,7 @@ void testInitialRetry() { @Test @SuppressWarnings("ConstantConditions") - void testUnauthenticatedRefresh() { + void testUnauthenticatedRefreshWithoutGenerationAdvance_notFlagged() { TransportChannel transportChannel = Mockito.mock(TransportChannel.class); Mockito.when(transportChannel.shouldRefresh()).thenReturn(true); @@ -296,7 +296,7 @@ void testUnauthenticatedRefresh() { @Test @SuppressWarnings("ConstantConditions") - void testUnauthenticatedRefreshWithGenerationAdvanceRetries() { + void testUnauthenticatedRefreshWithGenerationAdvance_flagsChannelRefreshed() { TransportChannel transportChannel = Mockito.mock(TransportChannel.class); java.util.concurrent.atomic.AtomicLong generation = new java.util.concurrent.atomic.AtomicLong(0); @@ -342,6 +342,17 @@ void testUnauthenticatedRefreshWithGenerationAdvanceRetries() { .isTrue(); Truth.assertThat(((UnauthenticatedException) outerError.getCause()).isRetryable()).isFalse(); Truth.assertThat(outerError.getCause().getStackTrace()).isEqualTo(initialError.getStackTrace()); + // A fixed clock keeps the fake attempt (started at t=0) within its total timeout. + com.google.api.gax.retrying.StreamingRetryAlgorithm retryAlgorithm = + new com.google.api.gax.retrying.StreamingRetryAlgorithm<>( + new ApiResultRetryAlgorithm<>(), + new com.google.api.gax.retrying.ExponentialRetryAlgorithm( + RetrySettings.newBuilder().build(), new com.google.api.gax.core.FakeApiClock(0))); + // Mirror the retry executor: compute the next attempt, then ask whether to run it. + TimedAttemptSettings nextAttempt = + retryAlgorithm.createNextAttempt(outerError, null, fakeRetryingFuture.getAttemptSettings()); + Truth.assertThat(nextAttempt).isNotNull(); + Truth.assertThat(retryAlgorithm.shouldRetry(outerError, null, nextAttempt)).isTrue(); // Verify retry call resumes stream callable.call(); @@ -628,10 +639,13 @@ void testUnauthenticatedException_whenChannelRefreshes_flagsChannelRefreshed() t @Test void testUnauthenticatedException_whenChannelRefreshFails_notFlagged() throws Exception { + java.util.concurrent.atomic.AtomicInteger refreshCalls = + new java.util.concurrent.atomic.AtomicInteger(); FakeChannel fakeChannel = new FakeChannel() { @Override public void refresh() { + refreshCalls.incrementAndGet(); throw new RuntimeException("Refresh failed"); } }; @@ -653,6 +667,7 @@ public void refresh() { assertThrows( ExecutionException.class, () -> fakeRetryingFuture.getAttemptResult().get(1, TimeUnit.SECONDS)); + Truth.assertThat(refreshCalls.get()).isEqualTo(1); Truth.assertThat(ex.getCause()).isInstanceOf(ServerStreamingAttemptException.class); Throwable cause = ex.getCause().getCause(); Truth.assertThat(cause).isInstanceOf(UnauthenticatedException.class); @@ -661,7 +676,8 @@ public void refresh() { } @Test - void testUnauthenticated_generationAlreadyAdvanced_skipsShouldRefresh() throws Exception { + void testUnauthenticated_generationAlreadyAdvanced_flaggedWithoutCheckingShouldRefresh() + throws Exception { TransportChannel transportChannel = Mockito.mock(TransportChannel.class); java.util.concurrent.atomic.AtomicLong generation = new java.util.concurrent.atomic.AtomicLong(0);